Nota:
El acceso a esta página requiere autorización. Puede intentar iniciar sesión o cambiar directorios.
El acceso a esta página requiere autorización. Puede intentar cambiar los directorios.
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:
- Haga clic en el selector de proceso del cuaderno en la parte superior derecha y seleccione GPU sin servidor.
- En el lado derecho, haga clic en el botón de entorno.
- Seleccione 8xH100 como acelerador.
- 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
- Procedimientos recomendados para el proceso de GPU sin servidor
- Solución de problemas en el proceso de GPU sin servidor
- Entrenamiento distribuido con varias GPU y varios nodos