Ajuste completo de Qwen3-4B

Ajusta por completo el gran modelo de lenguaje Qwen3-4B en una sola GPU H100. En este tutorial se muestra cómo:

  • Ejecute el ajuste completo, que actualiza todos los parámetros del modelo para adaptarse al máximo a los datos.
  • Uso del entorno de Databricks AI v5 sin instalar bibliotecas adicionales
  • Utilizar TRL (Transformer Reinforcement Learning) para el ajuste fino supervisado
  • Registro del modelo optimizado en el Catálogo de Unity para la gobernanza y la implementación

Conceptos clave:

  • Ajuste fino completo: actualiza todos los pesos del modelo, lo que le otorga la máxima capacidad para aprender de su conjunto de datos, a costa de un mayor uso de memoria y capacidad de cálculo que los métodos eficientes en cuanto a parámetros.
  • TRL: una biblioteca para entrenar modelos de lenguaje con aprendizaje de refuerzo y ajuste fino supervisado
  • Entrenamiento eficiente en memoria: utiliza precisión mixta BF16 y puntos de control de gradiente (gradient checkpointing) para ajustar un modelo de 4B de parámetros (fine-tuning completo) en una única GPU H100.

Nota

Este ejemplo requiere el entorno de IA Databricks versión 5 o superior.

Ajuste completo vs. matriz de decisión LoRA

En este cuaderno usa ajuste completo, que actualiza todos los parámetros del modelo. La alternativa, LoRA (adaptación de rango bajo), congela el modelo base y entrena solo pequeñas capas adaptadoras.

Scenario Recommendation Reason
Cambio de comportamiento del modelo principal Ajuste completo Actualiza todos los parámetros para los cambios fundamentales en el comportamiento del modelo
Calidad más alta posible en una sola tarea Ajuste completo No hay aproximación de rango bajo, por lo que el modelo tiene capacidad completa para adaptarse
Memoria de GPU limitada LoRA Permite que modelos más grandes quepan en memoria entrenando solo ~1 % de los parámetros
Varios adaptadores específicos para cada tarea LoRA Intercambio de adaptadores diferentes en el mismo modelo base

El ajuste fino completo de un modelo de 4B de parámetros requiere significativamente más memoria de GPU que LoRA, debido a que se deben mantener los estados del optimizador y los gradientes para cada parámetro. Este cuaderno usa precisión mixta BF16 y puntos de control de gradiente para que el entrenamiento quepa en una sola GPU H100 (80 GB).

Conectar al cómputo de GPU sin servidor

Para conectarse al servicio de cómputo con GPU sin servidor:

  1. Haga clic en el menú desplegable Conectar del cuaderno y seleccione GPU sin servidor.
  2. Elija una GPU 1x H100 como acelerador.
  3. Abra el panel Entorno y elija AI v5 como entorno base.
  4. Haga clic en Aplicar.

Para más información, consulte la documentación de proceso de GPU.

Importar bibliotecas

El entorno de Databricks AI v5 ya incluye todas las bibliotecas necesarias para este ejemplo (como trl, transformers, datasetsy mlflow), por lo que no se necesita ninguna instalación adicional.

La celda siguiente importa las bibliotecas necesarias para el entrenamiento del modelo, el control de conjuntos de datos y el seguimiento de MLflow.

from datasets import load_dataset
from transformers import AutoModelForCausalLM, AutoTokenizer
from trl import (
    SFTConfig,
    SFTTrainer,
    setup_chat_format
)
import torch
import mlflow

Configuración de la instalación

Integración de Unity Catalog

La celda siguiente configura dónde se almacenará y registrará el modelo ajustado:

  • Catálogo y esquema: organice los modelos dentro del espacio de nombres del catálogo de Unity (valor predeterminado: main.default)
  • Nombre del modelo: el nombre del modelo registrado en el catálogo de Unity para la gobernanza y la implementación
  • Volumen: Volumen del catálogo de Unity para almacenar checkpoints del modelo durante el entrenamiento

Estos widgets permiten personalizar la ubicación de almacenamiento sin editar código. El modelo se registrará como {catalog}.{schema}.{model_name} para facilitar el acceso y el control de versiones.

Hiperparámetros de entrenamiento

La celda también define los parámetros de entrenamiento clave:

  • Modelo y conjunto de datos: Qwen3-4B con el conjunto de datos conversacional de Capybara
  • Tamaño de lote (1): número de ejemplos por GPU y por paso de entrenamiento, que se mantiene pequeño para que un ajuste fino completo quepa en memoria
  • Acumulación de gradiente (8): acumula los gradientes durante 8 lotes para un tamaño de lote efectivo de 8
  • Tasa de aprendizaje (2e-5): tasa conservadora adecuada para el ajuste fino completo
  • Pasos máximos (50): limita el entrenamiento en 50 pasos para una ejecución rápida de demostración
  • Registro y punto de comprobación: guarda el progreso cada 25 pasos, registra las métricas cada 10 pasos.
dbutils.widgets.text("uc_catalog", "main")
dbutils.widgets.text("uc_schema", "default")
dbutils.widgets.text("uc_model_name", "qwen3_4b_assistant")
dbutils.widgets.text("uc_volume", "checkpoints")

UC_CATALOG = dbutils.widgets.get("uc_catalog")
UC_SCHEMA = dbutils.widgets.get("uc_schema")
UC_MODEL_NAME = dbutils.widgets.get("uc_model_name")
UC_VOLUME = dbutils.widgets.get("uc_volume")

print(f"UC_CATALOG: {UC_CATALOG}")
print(f"UC_SCHEMA: {UC_SCHEMA}")
print(f"UC_MODEL_NAME: {UC_MODEL_NAME}")
print(f"UC_VOLUME: {UC_VOLUME}")

# MLflow and Unity Catalog configuration

# Model selection
MODEL_NAME = "Qwen/Qwen3-4B"
DATASET_NAME = "trl-lib/Capybara"
OUTPUT_DIR = f"/Volumes/{UC_CATALOG}/{UC_SCHEMA}/{UC_VOLUME}/{UC_MODEL_NAME}"

# Training hyperparameters
BATCH_SIZE = 1
GRADIENT_ACCUMULATION_STEPS = 8
LEARNING_RATE = 2e-5
MAX_STEPS = 50
EVAL_STEPS = 25
LOGGING_STEPS = 10
SAVE_STEPS = 25

Carga y preparación del conjunto de datos

La celda siguiente carga el conjunto de datos de entrenamiento y lo prepara para ajustarlo:

  • Conjunto de datos: trl-lib/Capybara : datos conversacionales de alta calidad optimizados para instrucciones siguientes
  • División de entrenamiento y validación: crea una división de 90/10 si no existe ningún conjunto de pruebas.
  • Validación de datos: garantiza un formato adecuado para el ajuste preciso de la conversación
dataset = load_dataset(DATASET_NAME)
print(f"✓ Dataset loaded: {dataset}")

if "test" not in dataset:
    print("Creating validation split from training data...")
    dataset = dataset["train"].train_test_split(test_size=0.1, seed=42)
    print("✓ Data split: 90% train, 10% validation")

Inicializar el modelo y el tokenizador

La celda siguiente carga el modelo base y el tokenizador y, a continuación, los configura para el ajuste preciso de la conversación:

  • Carga del modelo: Descarga Qwen3-4B desde Hugging Face con precisión BF16
  • Configuración del tokenizador: Configura el tokenizador rápido con espaciado adecuado
  • Formato de chat: aplica una plantilla de chat para conversaciones estructuradas si el tokenizador aún no define uno.
  • Configuración de token: establece el token de relleno al token EOS para el manejo adecuado de secuencias
model = AutoModelForCausalLM.from_pretrained(
    MODEL_NAME,
    torch_dtype=torch.bfloat16,
    trust_remote_code=True,
)

tokenizer = AutoTokenizer.from_pretrained(
    MODEL_NAME,
    trust_remote_code=True,
    use_fast=True
)

# Chat template formatting for conversational fine-tuning
if tokenizer.chat_template is None:
    print("Adding chat template for proper conversation formatting...")
    model, tokenizer = setup_chat_format(model, tokenizer, format="chatml")
    print("✓ ChatML format applied for structured conversations")

if tokenizer.pad_token is None:
    tokenizer.pad_token = tokenizer.eos_token
    print("✓ Padding token set to EOS token")

print("✓ Model and tokenizer loaded successfully")

Entrenamiento del modelo

La celda siguiente configura y ejecuta el proceso de ajuste completo:

Configuración de entrenamiento

  • Configuración del lote: 1 muestra por dispositivo con 8 pasos de acumulación de gradiente (tamaño de lote efectivo: 8)
  • Optimización: pasos de calentamiento, decaimiento de peso y selección del mejor modelo basada en la pérdida durante la evaluación
  • Registro: notifica métricas a MLflow para el seguimiento de experimentos

Optimizaciones clave habilitadas

  • Precisión mixta BF16: cálculo más rápido con una superficie de memoria inferior, adecuada para GPU H100
  • Puntos de control de gradiente: intercambia cómputo adicional por una gran reducción de la memoria de activaciones, lo que permite que un ajuste fino completo de 4B quepa en una sola H100
  • Acumulación de gradiente: simula tamaños de lote más grandes para un entrenamiento más estable
  • Puntos de comprobación: guarda el modelo cada 25 pasos con un límite de 2 puntos de comprobación.

El bucle de entrenamiento registra el progreso cada 10 pasos y evalúa cada 25 pasos.

with mlflow.start_run(run_name=f"{MODEL_NAME}_full-fine-tuning", log_system_metrics=True):
    try:
        print(f"Learning rate: {LEARNING_RATE}")

        training_args_dict = {
            "output_dir": OUTPUT_DIR,
            "per_device_train_batch_size": BATCH_SIZE,
            "per_device_eval_batch_size": BATCH_SIZE,
            "gradient_accumulation_steps": GRADIENT_ACCUMULATION_STEPS,
            "learning_rate": LEARNING_RATE,
            "max_steps": MAX_STEPS,
            "eval_steps": EVAL_STEPS,
            "logging_steps": LOGGING_STEPS,
            "save_steps": SAVE_STEPS,
            "save_total_limit": 2,
            "report_to": "mlflow",  # Log to MLflow
            "warmup_steps": 10,
            "weight_decay": 0.01,
            "metric_for_best_model": "eval_loss",
            "greater_is_better": False,
            "eval_strategy": "steps",  # Run evaluation every eval_steps
            "save_strategy": "steps",  # Checkpoint on the same cadence as eval
            "load_best_model_at_end": True,  # Register the best-eval checkpoint, not the last
            "dataloader_pin_memory": False,
            "remove_unused_columns": False,
            "bf16": True,  # Mixed precision training
            "gradient_checkpointing": True,  # Reduce activation memory for full fine-tuning
            "gradient_checkpointing_kwargs": {"use_reentrant": False},
        }

        training_args = SFTConfig(**training_args_dict)

        trainer = SFTTrainer(
            model=model,
            args=training_args,
            train_dataset=dataset["train"],
            eval_dataset=dataset["test"],
            processing_class=tokenizer,
        )

        print("\n" + "="*50)
        print("STARTING TRAINING")
        print("="*50)

        print("🚀 Full fine-tuning Qwen3-4B on a single H100 GPU")

        trainer.train()
        print("\n✓ Training completed successfully!")

    except Exception as e:
        print(f"✗ Training failed: {e}")
        raise

Guardar elementos del modelo

La siguiente celda guarda el modelo entrenado y el tokenizador en el volumen del Unity Catalog:

  • Pesos completos del modelo: guarda el modelo completo optimizado, listo para cargarse directamente para la inferencia.
  • Tokenizer: guarda la configuración del tokenizador para la inferencia.
  • Ubicación de almacenamiento: guarda en /Volumes/{catalog}/{schema}/{volume}/{model_name}
try:
    print("\nSaving trained model...")

    trainer.save_model(training_args.output_dir)
    print("✓ Full model weights saved")

    tokenizer.save_pretrained(training_args.output_dir)
    print("✓ Tokenizer saved with model")
    print(f"\n🎉 All artifacts saved to: {training_args.output_dir}")

except Exception as e:
    print(f"✗ Model saving failed: {e}")
    raise

Registro del modelo en el catálogo de Unity

La celda siguiente registra el modelo optimizado en el Catálogo de Unity para la gobernanza y la implementación:

Flujo de trabajo de registro de modelos

  1. Cargar modelo entrenado: carga el modelo completo guardado con todos sus pesos y el tokenizador
  2. Preparación para el registro: crea un diccionario de modelos de transformadores con el modelo y el tokenizador
  3. Registro en Unity Catalog: registros en MLflow y registros en Unity Catalog
  4. Agregar metadatos: incluye el tipo de tarea, la familia de modelos y la información de tamaño.

Ventajas del registro del catálogo de Unity

  • Gobernanza: registro de modelos centralizado con control de acceso y seguimiento de linaje
  • Control de versiones: administración automática de versiones para el ciclo de vida del modelo
  • Implementación: implementación sencilla para modelar puntos de conexión
  • Detectabilidad: los modelos se pueden buscar y documentar en el catálogo de Unity
mlflow_run_id = mlflow.last_active_run().info.run_id
print("\nRegistering model with MLflow and Unity Catalog...")

with mlflow.start_run(run_id=mlflow_run_id):
    try:
        # Load the trained full-weight model for registration
        print("Loading fine-tuned model for registration...")
        trained_model = AutoModelForCausalLM.from_pretrained(
            training_args.output_dir,
            torch_dtype=torch.bfloat16,
            trust_remote_code=True
        )
        tokenizer = AutoTokenizer.from_pretrained(training_args.output_dir)
        model_type = "Full fine-tuning"
        size_params = "4b"

        # Prepare transformers model dictionary
        transformers_model = {
            "model": trained_model,
            "tokenizer": tokenizer
        }

        # Create Unity Catalog model name
        full_model_name = f"{UC_CATALOG}.{UC_SCHEMA}.{UC_MODEL_NAME}"

        print(f"Registering model as: {full_model_name}")

        # Start MLflow run and log model
        task = "llm/v1/chat"
        model_info = mlflow.transformers.log_model(
            transformers_model=transformers_model,
            task=task,
            registered_model_name=full_model_name,
            metadata={
                "task": task,
                "pretrained_model_name": MODEL_NAME,
                "databricks_model_family": "Qwen3ForCausalLM",
                "databricks_model_size_parameters": size_params,
            },
            repo_type="local",  # Fix: specify repo_type for local path
        )

        print(f"✓ Model successfully registered in Unity Catalog: {full_model_name}")
        print(f"✓ MLflow model URI: {model_info.model_uri}")

        # Print deployment information
        print(f"\n📦 Model Registration Complete!")
        print(f"Unity Catalog Path: {full_model_name}")
        print(f"Model Type: {model_type}")

    except Exception as e:
        print(f"✗ Model registration failed: {e}")
        print("Model is still saved locally and can be registered manually")
        print(f"Local model path: {training_args.output_dir}")
        raise

Pasos siguientes

Su modelo Qwen3-4B se ha perfeccionado correctamente mediante ajuste fino de todos los pesos y se ha registrado en Unity Catalog. A continuación, puede hacer lo siguiente:

Cuaderno de ejemplo

Ajuste completo de Qwen3-4B

Obtener el portátil