Aller au contenu principal

Affinez Olmo3 7B avec Axolotl sur un compute serverless multi-GPU.

Ouvrir dans Databricks

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 :

  1. Cliquez sur le sélecteur de compute du notebook en haut à droite et sélectionnez Serverless GPU .
  2. Sur le côté droit, cliquez sur le bouton d'environnement.
  3. Sélectionnez 8xH100 comme Accélérateur
  4. 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.

Python
%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.

Python
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é.

Python
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.

Python
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.

Python
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.

Python
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.

Python
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
Python
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.

Python
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.

Python
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

Exemple de Notebook

Affiner Olmo3 7B avec Axolotl sur un compute serverless multi-GPU