Feinabstimmung von Olmo3 7B mit Axolotl auf einer serverlosen Multi-GPU-Berechnungsplattform

Optimieren Sie das Olmo3 7B Instruct-Modell auf AI Runtime mit Axolotl. Axolotl bietet ein hochleistungsfähiges Framework für LLM nach dem Training mit QLoRA (quantized Low-Rank Adaption), wodurch eine effiziente Feinabstimmung in der Multi-GPU-Infrastruktur ermöglicht wird. Das trainierte Modell wird bei MLflow protokolliert und für die Bereitstellung im Unity-Katalog registriert.

Verbindung zu Serverless GPU-Compute herstellen

Für dieses Notebook ist ein serverloses GPU-Compute erforderlich. So stellen Sie eine Verbindung her:

  1. Klicken Sie oben rechts auf die Compute-Auswahl des Notebooks und wählen Sie Serverlose GPU aus.
  2. Klicken Sie auf der rechten Seite auf die Schaltfläche "Umgebung".
  3. Wählen Sie 8xH100 als Beschleuniger aus.
  4. Wählen Sie AI v5 als Umgebung aus und klicken Sie dann auf Übernehmen.

Installieren erforderlicher Abhängigkeiten

Installiert Axolotl mit Flash Attention-Unterstützung und kompatiblen Versionen von trl- und Optimierungsbibliotheken. Das cut-cross-entropy Paket bietet speichereffiziente Verlustberechnungen für große Sprachmodelle.

%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()

HuggingFace-Token abrufen

Ruft das HuggingFace-Authentifizierungstoken aus Databricks-Geheimnissen ab. Dieses Token ist erforderlich, um das Olmo3 7B-Basismodell aus dem HuggingFace Hub herunterzuladen.

HF_TOKEN = dbutils.secrets.get(scope="sgc-nightly-notebook", key="hf_token")

Konfigurieren von Schulungsparametern

Richtet die Axolotl-Schulungskonfiguration basierend auf dem Olmo3-7b-qlora.yaml-Beispiel ein. Zu den wichtigsten Änderungen gehören:

  • MLflow-Integration für die Experimentverfolgung
  • Unity-Katalog-Volumepfad für Checkpoint-Speicher
  • SDPA (Skalierte Punktproduktaufmerksamkeit) anstelle von Flash-Aufmerksamkeit für eine breitere GPU-Kompatibilität

Definieren von Unity-Katalogpfaden

Erstellt Widgets zum Angeben des Unity-Katalogspeicherorts zum Speichern von Modellprüfpunkten. Das Ausgabeverzeichnis kombiniert den Katalog-, Schema-, Volume- und Modellnamen in einem vollqualifizierten Pfad.

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}")

Telemetrie deaktivieren

Deaktiviert die Verwendungsnachverfolgung von Axolotl durch Festlegen der Umgebungsvariable.

import os
os.environ['AXOLOTL_DO_NOT_TRACK'] = '1'

Axolotl-Konfiguration erstellen

Definiert die vollständige Schulungskonfiguration mithilfe des Axolotl-Formats DictDefault . Dazu gehören Modelleinstellungen (QLoRA mit 4-Bit-Quantisierung), Datasetkonfiguration (Alpaca-Format), LoRA-Hyperparameter (Rank 32, Alpha 16), Trainingsparameter (1 Epoche, Batchgröße 2, Gradientenakkumulation 4) und MLflow-Integration für die Verfolgung von Experimenten.

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,
)

Konfigurieren der PyTorch-CUDA-Speicherzuweisung

Optimiert die GPU-Speicherverwaltung für eine effiziente Schulung in Multi-GPU-Setups.

from axolotl.utils import set_pytorch_cuda_alloc_conf

set_pytorch_cuda_alloc_conf()

Ausführen einer verteilten Schulung auf serverlosem GPU-Compute

Verwendet den Dekorator @distributed aus der serverlosen GPU-API, um den Axolotl-Trainingsauftrag auf 8 H100 GPUs zu verteilen. Der Dekorator behandelt die Multi-GPU-Orchestrierung, sodass die Trainingsfunktion in einer verteilten Umgebung ohne manuelles Clustersetup ausgeführt werden kann.

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)

Ausführen des Schulungsauftrags

Startet den verteilten Trainingsprozess. Die Funktion lädt das Dataset, überprüft die Konfiguration, trainiert das Modell für 16 Schritte und gibt die MLflow-Ausführungs-ID zur Nachverfolgung zurück.

run_id = result[0]
print(run_id)

Extrahieren der MLflow-Ausführungs-ID

Ruft die MLflow-Ausführungs-ID aus den Schulungsergebnissen für die Modellregistrierung und Experimentverfolgung ab.

Registrieren des fein abgestimmten Modells im Unity-Katalog

Lädt den trainierten LoRA-Adapter, führt ihn mit dem Basismodell zusammen und registriert das kombinierte Modell über MLflow im Unity-Katalog. Dadurch wird das Modell für die Bereitstellung und Ableitung verfügbar.

Hinweis: Für diesen Schritt ist eine H100 GPU-Berechnung erforderlich, um den Modellprüfpunkt zu laden. Die Ausführung auf kleineren GPUs kann zu CUDA-Out-of-Memory-Fehlern führen.

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")

Nächste Schritte

Beispiel-Notebook

Feinabstimmung von Olmo3 7B mit Axolotl auf einer serverlosen Multi-GPU-Berechnungsplattform

Notebook abrufen