Affinez Olmo3 7B avec Axolotl sur un compute serverless multi-GPU.
Affiner le modèle Olmo3 7B Instruct sur AI Runtime à l’aide d’ Axolotl. Axolotl fournit un framework haute performance pour le post-entraînement de LLM avec QLoRA (Quantized Low-Rank Adaptation), permettant un affinement efficace sur une infrastructure multi-GPU. Le modèle entraîné est enregistré dans MLflow et dans Unity Catalog pour le déploiement.
Se connecter au compute GPU serverless
Ce Notebook nécessite un compute GPU serverless. Connecter :
- Cliquez sur le sélecteur de compute du notebook en haut à droite et sélectionnez Serverless GPU .
- Sur le côté droit, cliquez sur le bouton d'environnement.
- Sélectionnez 8xH100 comme Accélérateur
- Sélectionnez AI v5 comme environnement, puis cliquez sur Appliquer
Installer les dépendances requises
Installe Axolotl avec la prise en charge de Flash Attention et des versions compatibles de trl et des bibliothèques d’optimisation. Le package cut-cross-entropy fournit un calcul de perte économe en mémoire pour les grands modèles de langage.
%pip install --no-build-isolation "axolotl[flash-attn]==0.13.1"
%pip install "trl==0.27.1"
%pip install "torchao==0.16.0"
%pip install "cut-cross-entropy[transformers] @ git+https://github.com/axolotl-ai-cloud/ml-cross-entropy.git@f4b5712"
dbutils.library.restartPython()
Récupérer le jeton HuggingFace
Récupère le jeton d'authentification HuggingFace des secrets Databricks. Ce jeton est requis pour download le modèle de base Olmo3 7B depuis le Hub HuggingFace.
HF_TOKEN = dbutils.secrets.get(scope="sgc-nightly-notebook", key="hf_token")
Configurer les paramètres d'entraînement
Configure la configuration d'entraînement Axolotl basée sur l'exemple olmo3-7b-qlora.yaml. Les modifications clés incluent :
- Intégration MLflow pour le suivi des expérimentations
- Chemin d’accès au volume Unity Catalog pour le stockage des points de contrôle
- SDPA (Scaled Dot Product Attention) plutôt que Flash Attention pour une compatibilité GPU plus large
Définir les chemins Unity Catalog
Crée des widgets pour spécifier l'emplacement Unity Catalog pour le stockage des points de contrôle du modèle. Le répertoire de sortie combine le nom du catalogue, du schéma, du volume et du modèle en un chemin d'accès entièrement qualifié.
dbutils.widgets.text("uc_catalog", "main")
dbutils.widgets.text("uc_schema", "default")
dbutils.widgets.text("uc_volume", "checkpoints")
dbutils.widgets.text("model", "openai/gpt-oss-20b")
UC_CATALOG = dbutils.widgets.get("uc_catalog")
UC_SCHEMA = dbutils.widgets.get("uc_schema")
UC_VOLUME = dbutils.widgets.get("uc_volume")
UC_MODEL_NAME = dbutils.widgets.get("model")
print(f"UC_CATALOG: {UC_CATALOG}")
print(f"UC_SCHEMA: {UC_SCHEMA}")
print(f"UC_VOLUME: {UC_VOLUME}")
print(f"UC_MODEL_NAME: {UC_MODEL_NAME}")
OUTPUT_DIR = f"/Volumes/{UC_CATALOG}/{UC_SCHEMA}/{UC_VOLUME}/{UC_MODEL_NAME}"
print(f"OUTPUT_DIR: {OUTPUT_DIR}")
Désactiver la télémétrie
Désactive le suivi d'utilisation d'Axolotl en définissant la variable d'environnement.
import os
os.environ['AXOLOTL_DO_NOT_TRACK'] = '1'
Créer une configuration Axolotl
Définit la configuration d'entraînement complète à l'aide du format DictDefault d'Axolotl. Cela inclut les paramètres du modèle (QLoRA avec quantification 4 bits), la configuration du dataset (format Alpaca), les hyper-parameters LoRA (rang 32, alpha 16), les parameters d'entraînement (1 epoch, taille de batch 2, accumulation de gradient 4), et l'intégration de MLflow pour le suivi d'Experimentation.
from axolotl.cli.config import load_cfg
from axolotl.utils.dict import DictDefault
# Config is based on with some changes to fit GPU types
# https://raw.githubusercontent.com/axolotl-ai-cloud/axolotl/main/examples/olmo3/olmo3-7b-qlora.yaml
# Axolotl provides full control and transparency over model and training configuration
config = DictDefault(
base_model="allenai/Olmo-3-7B-Instruct-SFT",
plugins=[
"axolotl.integrations.cut_cross_entropy.CutCrossEntropyPlugin"
],
load_in_8bit=False,
load_in_4bit=True,
datasets=[
{
"path": "fozziethebeat/alpaca_messages_2k_test",
"type": "chat_template"
}
],
dataset_prepared_path="last_run_prepared",
val_set_size=0.1,
output_dir=OUTPUT_DIR,
adapter="qlora",
lora_model_dir=None,
sequence_len=2048,
sample_packing=True,
lora_r=32,
lora_alpha=16,
lora_dropout=0.05,
lora_target_linear=True,
lora_target_modules=[
"gate_proj",
"down_proj",
"up_proj",
"q_proj",
"v_proj",
"k_proj",
"o_proj"
],
wandb_project=None,
wandb_entity=None,
wandb_watch=None,
wandb_name=None,
wandb_log_model=None,
gradient_accumulation_steps=4,
micro_batch_size=2,
num_epochs=1,
optimizer="adamw_bnb_8bit",
lr_scheduler="cosine",
learning_rate=0.0002,
bf16="auto",
tf32=False,
gradient_checkpointing=True,
resume_from_checkpoint=None,
logging_steps=1,
flash_attention=False,
warmup_ratio=0.1,
evals_per_epoch=1,
saves_per_epoch=1,
# Eval dataset is too small
eval_sample_packing=False,
# Write metrics to MLflow
use_mlflow=True,
mlflow_tracking_uri="databricks",
mlflow_run_name="olmo3-7b-qlora-axolotl",
hf_mlflow_log_artifacts=False,
wandb_mode="disabled",
attn_implementation="sdpa",
sdpa_attention=True,
save_first_step=True,
device_map=None,
)
Configurer l'allocation de mémoire CUDA PyTorch
Optimise la gestion de la mémoire GPU pour un entraînement efficace sur des configurations multi-GPU.
from axolotl.utils import set_pytorch_cuda_alloc_conf
set_pytorch_cuda_alloc_conf()
Exécuter l'entraînement distribué sur le compute GPU serverless
Utilise le décorateur @distributed de l'API GPU serverless pour distribuer le job d'entraînement Axolotl sur 8 GPU H100. Le décorateur gère l'orchestration multi-GPU, permettant à la fonction d'entraînement de s'exécuter dans un environnement distribué sans configuration manuelle de cluster.
from serverless_gpu.launcher import distributed
from serverless_gpu.compute import GPUType
@distributed(gpus=8, gpu_type=GPUType.H100)
def run_train(cfg: DictDefault):
import os
os.environ['HF_TOKEN'] = HF_TOKEN
from axolotl.common.datasets import load_datasets
# Load, parse and tokenize the datasets to be formatted with qwen3 chat template
# Drop long samples from the dataset that overflow the max sequence length
# validates the configuration
cfg = load_cfg(cfg)
dataset_meta = load_datasets(cfg=cfg)
from axolotl.train import train
# just train the first 16 steps for demo.
# This is sufficient to align the model as we've used packing to maximize the trainable samples per step.
cfg.max_steps = 16
model, tokenizer, trainer = train(cfg=cfg, dataset_meta=dataset_meta)
import mlflow
mlflow_run_id = None
if mlflow.last_active_run() is not None:
mlflow_run_id = mlflow.last_active_run().info.run_id
return mlflow_run_id
result = run_train.distributed(config)
Exécuter le Job de formation
Lance le Job de formation distribuée. La fonction charge le dataset, valide la configuration, entraîne le modèle en 16 étapes et renvoie l'ID d'exécution MLflow pour le suivi.
run_id = result[0]
print(run_id)
Extraire l'ID d'exécution MLflow
Récupère l'ID d'exécution MLflow des résultats de l'entraînement pour l'enregistrement de modèle et le suivi des expérimentations.
Enregistrez le modèle affiné dans Unity Catalog
Charge l'adaptateur LoRA entraîné, le merge avec le modèle de base et enregistre le modèle combiné dans Unity Catalog via MLflow. Cela rend le modèle disponible pour le déploiement et l'inférence.
Remarque : cette étape nécessite un compute GPU H100 pour charger le checkpoint du modèle. L'exécution sur des GPU plus petits peut entraîner des erreurs CUDA de mémoire insuffisante.
from transformers import AutoTokenizer, AutoModelForCausalLM, pipeline
from peft import PeftModel
import mlflow
import torch
HF_MODEL_NAME = "allenai/Olmo-3-7B-Instruct-SFT"
torch.cuda.empty_cache()
# Load the trained model for registration
print("Loading LoRA model for registration...")
# For LoRA models, we need both base model and adapter
base_model = AutoModelForCausalLM.from_pretrained(
HF_MODEL_NAME,
trust_remote_code=True
)
# Load tokenizer
tokenizer = AutoTokenizer.from_pretrained(HF_MODEL_NAME)
adapter_dir = OUTPUT_DIR
peft_model = PeftModel.from_pretrained(base_model, adapter_dir)
# Merge LoRA into base and drop PEFT wrappers
merged_model = peft_model.merge_and_unload()
merged_model.generation_config.temperature = None
merged_model.generation_config.top_p = None
# Create Unity Catalog model name
full_model_name = f"{UC_CATALOG}.{UC_SCHEMA}.{UC_MODEL_NAME}"
print(f"Registering model as: {full_model_name}")
text_gen_pipe = pipeline(
task="text-generation",
model=merged_model,
tokenizer=tokenizer,
)
input_example = ["Hello, world!"]
with mlflow.start_run(run_id=run_id):
model_info = mlflow.transformers.log_model(
transformers_model=text_gen_pipe,
name="model",
input_example=input_example,
registered_model_name=full_model_name,
)
print(f"✓ Model successfully registered in Unity Catalog: {full_model_name}")
print(f"✓ MLflow model URI: {model_info.model_uri}")
print(f"✓ Model version: {model_info.registered_model_version}")
print(f"\n📦 Model Registration Complete!")
print(f"Unity Catalog Path: {full_model_name}")
print(f"Optimization: Cut Cross Entropy + QLoRA")
Étapes suivantes
- Meilleures pratiques pour le compute GPU Serverless
- Dépanner les problèmes sur le compute GPU serverless
- Formation distribuée multi-GPU et multi-nœuds