Nota
L'accesso a questa pagina richiede l'autorizzazione. È possibile provare ad accedere o modificare le directory.
L'accesso a questa pagina richiede l'autorizzazione. È possibile provare a modificare le directory.
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:
- Fare clic sul menu a discesa Connetti nel notebook e selezionare GPU serverless.
- Scegliere una GPU H100 1x come acceleratore.
- Aprire il pannello Ambiente e scegliere intelligenza artificiale v5 come ambiente di base.
- 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
- Carica il modello addestrato: carica il modello con pesi completi salvato e il tokenizer
- Preparazione della registrazione dei log: crea un dizionario dei modelli Transformers con il modello e il tokenizer
- Registra nel Catalogo Unity: registra in MLflow e registra nel Catalogo Unity
- 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:
- Distribuire il modello: Gestire i modelli con Model Serving
- Altre informazioni sul training distribuito: Training distribuito multi-GPU e multinodo
- Tenere traccia degli esperimenti e monitorare le GPU: rilevamento e osservabilità degli esperimenti
- Risolvere i problemi: risolvere i problemi relativi all'ambiente di calcolo GPU serverless