Ajustar Olmo3 7B con Axolotl en informática sin servidor con varias GPU

Ajustar el modelo Olmo3 7B Instruct en AI Runtime con Axolotl. Axolotl proporciona un marco de alto rendimiento para la fase posterior al entrenamiento del LLM con QLoRA (Adaptación cuantificada de bajo rango), lo que permite un ajuste eficiente en la infraestructura con varias GPU. El modelo entrenado se registra en MLflow y se registra en el catálogo de Unity para su implementación.

Nota

Este ejemplo requiere el entorno de IA Databricks versión 5 o superior.

Conectar al cómputo de GPU sin servidor

Este notebook requiere computación GPU sin servidor. Para conectarse:

  1. Haga clic en el selector de proceso del cuaderno en la parte superior derecha y seleccione GPU sin servidor.
  2. En el lado derecho, haga clic en el botón de entorno.
  3. Seleccione 8xH100 como acelerador.
  4. Seleccione AI v5 como entorno y haga clic en Aplicar.

Instalación de dependencias necesarias

Instala Axolotl con soporte para Flash Attention y versiones compatibles de trl y bibliotecas de optimización. El paquete cut-cross-entropy proporciona cálculos de pérdida eficientes en memoria para modelos de lenguaje grandes.

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

Recuperación del token de HuggingFace

Obtiene el token de autenticación HuggingFace de los secretos de Databricks. Este token es necesario para descargar el modelo base olmo3 7B desde huggingFace Hub.

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

Configuración de parámetros de entrenamiento

Configura la configuración de entrenamiento de Axolotl basada en el ejemplo olmo3-7b-qlora.yaml . Entre las modificaciones clave se incluyen las siguientes:

  • Integración de MLflow para el seguimiento de experimentos
  • Ruta de acceso al volumen de Unity Catalog para el almacenamiento de puntos de control
  • SDPA (Atención por producto escalar y escalado) en lugar de Atención flash para una mayor compatibilidad con GPU

Definir rutas del catálogo Unity

Crea widgets para especificar la ubicación del catálogo de Unity para almacenar puntos de control del modelo. El directorio de salida combina el catálogo, el esquema, el volumen y el nombre del modelo en una ruta de acceso completa.

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

Deshabilitar la telemetría

Deshabilita el seguimiento de uso de Axolotl estableciendo la variable de entorno.

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

Creación de una configuración de Axolotl

Define la configuración de entrenamiento completa utilizando el formato DictDefault de Axolotl. Esto incluye la configuración del modelo (QLoRA con cuantificación de 4 bits), la configuración del conjunto de datos (formato Alpaca), los hiperparámetros de LoRA (rango 32, alfa 16), los parámetros de entrenamiento (1 época, tamaño de lote 2, acumulación de degradado 4) y la integración de MLflow para el seguimiento de experimentos.

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

Configuración de la asignación de memoria cuDA de PyTorch

Optimiza la administración de memoria de GPU para un entrenamiento eficaz en configuraciones de varias GPU.

from axolotl.utils import set_pytorch_cuda_alloc_conf

set_pytorch_cuda_alloc_conf()

Ejecutar el entrenamiento distribuido en computación GPU sin servidor

Usa el decorador @distributed de la API de GPUs sin servidor para distribuir el trabajo de entrenamiento de Axolotl en 8 GPU H100. El decorador controla la orquestación de varias GPU, lo que permite que la función de entrenamiento se ejecute en un entorno distribuido sin la configuración manual del clúster.

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)

Ejecución del trabajo de entrenamiento

Inicia el trabajo de entrenamiento distribuido. La función carga el conjunto de datos, valida la configuración, entrena el modelo para 16 pasos y devuelve el identificador de ejecución de MLflow para el seguimiento.

run_id = result[0]
print(run_id)

Extracción del identificador de ejecución de MLflow

Recupera el identificador de ejecución de MLflow de los resultados de entrenamiento para el registro del modelo y el seguimiento de experimentos.

Registro del modelo optimizado en el catálogo de Unity

Carga el adaptador loRA entrenado, lo combina con el modelo base y registra el modelo combinado en el catálogo de Unity a través de MLflow. Esto hace que el modelo esté disponible para la implementación y la inferencia.

Nota: Este paso requiere la capacidad de cómputo de la GPU H100 para cargar el checkpoint del modelo. La ejecución en GPU más pequeñas puede provocar errores de memoria insuficiente de CUDA.

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

Pasos siguientes

Cuaderno de ejemplo

Ajustar Olmo3 7B con Axolotl en informática sin servidor con varias GPU

Obtención del cuaderno