Ottimizzazione completa di Qwen3-4B

Ottimizzare completamente il modello linguistico di grandi dimensioni Qwen3-4B su una singola GPU H100. Questa procedura dettagliata illustra come:

  • Eseguire l'ottimizzazione completa, che aggiorna ogni parametro del modello per il massimo adattamento ai dati
  • Usare l'ambiente di Intelligenza artificiale di Databricks v5 senza installare librerie aggiuntive
  • Sfruttare il TRL (Transformer Reinforcement Learning) per l'ottimizzazione con supervisione
  • Registrare il modello ottimizzato in Unity Catalog per la governance e la distribuzione

Concetti chiave:

  • Ottimizzazione completa: aggiorna tutti i pesi del modello, offrendo al modello la massima capacità di apprendere dal set di dati a un costo maggiore di memoria e calcolo rispetto ai metodi efficienti per i parametri
  • TRL: una libreria per l'addestramento di modelli linguistici con apprendimento per rinforzo e affinamento supervisionato
  • Addestramento efficiente in termini di memoria: usa la precisione mista BF16 e il gradient checkpointing per far rientrare un fine-tuning completo di un modello da 4 miliardi di parametri su una singola GPU H100

Note

Questo esempio richiede l'ambiente di IA Databricks versione 5 o superiore.

Ottimizzazione completa e matrice decisionale LoRA

Questo notebook usa l'ottimizzazione completa, che aggiorna tutti i parametri del modello. L'alternativa, LoRA (Low-Rank Adaptation), blocca il modello di base ed esegue il training solo di piccoli livelli adattatori.

Scenario Raccomandazione Ragione
Modifica principale del comportamento del modello Ottimizzazione completa Aggiorna tutti i parametri per le modifiche fondamentali al comportamento del modello
Massima qualità possibile in un'unica attività Ottimizzazione completa Nessuna approssimazione di basso rango, quindi il modello ha capacità completa per adattarsi
Memoria GPU limitata LoRA Consente di far stare in memoria modelli più grandi addestrando solo circa l’1% dei parametri
Adattatori multipli specifici per attività LoRA Scambiare adattatori diversi nello stesso modello di base

Il fine-tuning completo di un modello da 4 miliardi di parametri richiede una quantità di memoria GPU significativamente maggiore rispetto a LoRA, perché lo stato dell’ottimizzatore e i gradienti vengono mantenuti per ogni parametro. Questo notebook usa la precisione mista BF16 e il checkpointing del gradiente in modo che l'addestramento possa essere eseguito su una singola GPU H100 (80 GB).

Connessione alla computazione GPU senza server

Per connettersi al calcolo GPU serverless:

  1. Fare clic sul menu a discesa Connetti nel notebook e selezionare GPU serverless.
  2. Scegliere una GPU H100 1x come acceleratore.
  3. Aprire il pannello Ambiente e scegliere intelligenza artificiale v5 come ambiente di base.
  4. Fare clic su Applica.

Per altre informazioni, vedere la documentazione di calcolo della GPU.

Importare librerie

L'ambiente di Intelligenza artificiale di Databricks v5 include già tutte le librerie necessarie per questo esempio , ad esempio trl, transformersdatasets, e mlflow, quindi non è necessaria alcuna installazione aggiuntiva.

La cella successiva importa le librerie necessarie per il training del modello, la gestione dei set di dati e il rilevamento MLflow.

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

Impostazione della configurazione

Integrazione del catalogo Unity

La cella successiva configura la posizione in cui verrà archiviato e registrato il modello ottimizzato:

  • Catalogo e schema: organizzare i modelli all'interno dello spazio dei nomi del catalogo Unity (impostazione predefinita: main.default)
  • Nome modello: nome del modello registrato nel catalogo Unity per la governance e la distribuzione
  • Volume: Volume di Unity Catalog per l'archiviazione dei checkpoint del modello durante l'addestramento

Questi widget consentono di personalizzare la posizione di archiviazione senza modificare il codice. Il modello verrà registrato come {catalog}.{schema}.{model_name} per semplificare l'accesso e il controllo della versione.

Addestramento degli iperparametri

La cella definisce anche i parametri di training chiave:

  • Modello e set di dati: Qwen3-4B con il set di dati conversazionale Capybara
  • Dimensione del batch (1): numero di esempi per GPU per ogni fase di addestramento, mantenuto basso per far rientrare in memoria un fine-tuning completo
  • Accumulo dei gradienti (8): accumula i gradienti su 8 batch per una dimensione del batch effettiva di 8
  • Tasso di apprendimento (2e-5): tasso conservativo appropriato per l'ottimizzazione completa
  • Max steps (50): limita l’addestramento a 50 passaggi per una rapida esecuzione dimostrativa
  • Registrazione e checkpoint: salva lo stato di avanzamento ogni 25 passaggi, registra le metriche ogni 10 passaggi
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

Caricare e preparare il set di dati

La cella successiva carica il set di dati di training e lo prepara per l'ottimizzazione:

  • Set di dati: trl-lib/Capybara - Dati di dialogo di alta qualità ottimizzati per seguire le istruzioni
  • Divisione di addestramento/convalida: crea una divisione di 90/10 se non esiste alcun insieme di test
  • Convalida dei dati: garantisce una formattazione corretta per l'ottimizzazione della conversazione
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")

Inizializzare il modello e il tokenizzatore

La cella successiva carica il modello di base e il tokenizer, quindi li configura per l'ottimizzazione della conversazione:

  • Caricamento del modello: scarica Qwen3-4B da Hugging Face con precisione BF16
  • Configurazione del tokenizer: configura il tokenizzatore veloce con spaziatura interna corretta
  • Formattazione della chat: applica un modello per la chat alle conversazioni strutturate se il tokenizer non ne definisce già uno
  • Configurazione del token: imposta il token di riempimento sul token EOS per la gestione corretta della sequenza
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")

Addestra il modello

La cella successiva configura ed esegue il processo di ottimizzazione completa:

Configurazione della formazione

  • Configurazione del batch: 1 campione per dispositivo con 8 passaggi di accumulo del gradiente (dimensione effettiva del batch: 8)
  • Ottimizzazione: passaggi di riscaldamento, decadimento del peso e selezione del modello migliore in base alla perdita di valutazione
  • Registrazione: segnala le metriche a MLflow per il rilevamento dell'esperimento

Ottimizzazioni principali abilitate

  • Precisione mista BF16: calcolo più veloce con minore utilizzo di memoria, particolarmente adatto alle GPU H100
  • Checkpointing del gradiente: richiede calcolo aggiuntivo in cambio di una notevole riduzione della memoria di attivazione, il che consente di eseguire il fine-tuning completo di un modello da 4B su una singola H100
  • Accumulo dei gradienti: simula dimensioni batch maggiori per il training stabile
  • Checkpointing: salva il modello ogni 25 passi con un limite di 2 checkpoint

Il ciclo di training registra lo stato di avanzamento ogni 10 passaggi e valuta ogni 25 passaggi.

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

Salvare gli artefatti del modello

La cella successiva salva nel volume del catalogo Unity il modello addestrato e il tokenizer.

  • Pesi completi del modello: salva il modello completo sottoposto a fine-tuning, pronto per essere caricato direttamente per l'inferenza
  • Tokenizer: salva la configurazione del tokenizer per l'inferenza
  • Percorso di archiviazione: Salva in /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

Registrare il modello nel catalogo unity

La cella successiva registra il modello ottimizzato in Unity Catalog per la governance e la distribuzione:

Flusso di lavoro di registrazione del modello

  1. Carica il modello addestrato: carica il modello con pesi completi salvato e il tokenizer
  2. Preparazione della registrazione dei log: crea un dizionario dei modelli Transformers con il modello e il tokenizer
  3. Registra nel Catalogo Unity: registra in MLflow e registra nel Catalogo Unity
  4. Aggiungere metadati: include informazioni sul tipo di attività, sulla famiglia di modelli e sulle dimensioni

Vantaggi della registrazione del catalogo Unity

  • Governance: Registro dei modelli centralizzato con controllo di accesso e rilevamento della derivazione
  • Controllo delle versioni: gestione automatica delle versioni per il ciclo di vita del modello
  • Distribuzione: facile distribuzione per modellare gli endpoint di gestione
  • Scopribilità: i modelli sono ricercabili e documentati in Unity Catalog
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

Passaggi successivi

Il modello Qwen3-4B è stato ottimizzato correttamente con l'ottimizzazione completa e registrato in Unity Catalog. Successivamente, è possibile:

Notebook di esempio

Ottimizzazione completa di Qwen3-4B

Ottieni il notebook