Finjustera Llama 3.1 8B med mosaic LLM Foundry på Databricks Serverless GPU

Finjustera en Llama 3.1 8B-modell på AI Runtime med mosaic LLM Foundry, en kodbas för träning, finjustering, utvärdering och distribution av stora språkmodeller med stöd för distribuerade träningsstrategier.

Datorn använder:

  • Mosaic LLM Foundry: Ett ramverk för träning och finjustering av LLM med inbyggt stöd för FSDP, effektiv datainläsning och MLflow-integrering
  • FSDP (Fully Sharded Data Parallel): Distribuerar modellparametrar, gradienter och optimerartillstånd mellan GPU:er
  • Databricks Serverless GPU: Kör distribuerad träning på ansluten serverlös GPU-beräkning
  • Unity Catalog: Lagrar modellkontrollpunkter och registrerar tränade modeller
  • MLflow: Spårar experiment och loggar träningsmått

Note

Detta exempel kräver Standardmiljö version 5 eller högre.

Ansluta till serverlös GPU-beräkning

Den här notebooken kräver serverlös GPU-beräkning. Så här ansluter du:

  1. Klicka på notebook-filens beräkningsväljare längst upp till höger och välj Serverlös GPU
  2. Öppna sidopanelen Miljö på höger sida av anteckningsboken
  3. Ställ in Accelerator8xH100
  4. Välj standardbasmiljön och ange Miljöversion till 5, som innehåller de bibliotek som behövs för att köra det här exemplet
  5. Välj Använd och klicka på Bekräfta för att använda den här miljön för din notebook

Installera nödvändiga bibliotek

Installera Mosaic LLM Foundry och dess beroenden för distribuerad träning. Det förbyggda flash-attn wheel-paketet installeras först så att pip kan återanvända det i stället för att kompilera flash-attention från källkod (vilket är långsamt) när pip löser upp :

  • flash-attn: Optimerad uppmärksamhetsimplementering, installerad från ett fördefinierat hjul
  • llm-foundry: Kärnramverk för LLM-utbildning och finjustering
  • hf_transfer: Snabbare modellnedladdningar från Hugging Face
  • yamlmagic: Aktiverar YAML-konfiguration i notebook-celler
%pip install --no-deps "https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.6cxx11abiFALSE-cp312-cp312-linux_x86_64.whl"
%pip install llm-foundry[gpu]==0.20.0
%pip install hf_transfer
%pip install git+https://github.com/josejg/yamlmagic.git

Starta om Python-miljön

Starta om Python-kerneln för att säkerställa att de nyligen installerade paketen är tillgängliga.

dbutils.library.restartPython()

Konfigurera Unity-katalogsökvägar för modelllagring

Konfigurera Unity Catalog-platser för att lagra modellkontrollpunkter och registrera den tränade modellen. Konfigurationen använder frågeparametrar som kan anpassas utan att redigera koden.

dbutils.widgets.text("uc_catalog", "main")
dbutils.widgets.text("uc_schema", "default")
dbutils.widgets.text("uc_model_name", "llama3_1-8b")
dbutils.widgets.text("uc_volume", "checkpoints")

UC_CATALOG = dbutils.widgets.get("uc_catalog")
UC_SCHEMA = dbutils.widgets.get("uc_schema")
UC_MODEL_NAME = dbutils.widgets.get("uc_model_name")
UC_VOLUME = dbutils.widgets.get("uc_volume")

MLFLOW_EXPERIMENT_NAME = '/Workspace/Shared/llm-foundry-sgc' # TODO: update this name

print(f"UC_CATALOG: {UC_CATALOG}")
print(f"UC_SCHEMA: {UC_SCHEMA}")
print(f"UC_MODEL_NAME: {UC_MODEL_NAME}")
print(f"UC_VOLUME: {UC_VOLUME}")
print(f"EXPERIMENT_NAME: {MLFLOW_EXPERIMENT_NAME}")

# Model selection - Choose based on your compute constraints
OUTPUT_DIR = f"/Volumes/{UC_CATALOG}/{UC_SCHEMA}/{UC_VOLUME}/{UC_MODEL_NAME}" # Save checkpoint to UC Volume

print(f"OUTPUT_DIR: {OUTPUT_DIR}")

Definiera träningskonfiguration med YAML

Läs in finjusteringskonfigurationen från YAML-format. Konfigurationen anger:

  • Modellarkitektur och förtränad vikt (Llama 3.1 8B)
  • FSDP-inställningar för distribuerad träning
  • Träna hyperparametrar (inlärningshastighet, batchstorlek, optimerare)
  • Konfiguration av datasamling (mosaicml/dolly_hhrlhf)
  • MLflow-loggning och modellkontrollpunkter
  • Återanrop för övervakning och optimering
%load_ext yamlmagic
%%yaml config
seed: 17
model:
  name: hf_causal_lm
  pretrained: true
  init_device: mixed
  use_auth_token: true
  use_flash_attention_2: true
  pretrained_model_name_or_path: meta-llama/Llama-3.1-8B
loggers:
  mlflow:
    resume: true
    tracking_uri: databricks
    rename_metrics:
      time/token: time/num_tokens
      lr-DecoupledLionW/group0: learning_rate
    log_system_metrics: true
    experiment_name: "mlflow_experiment_name"
    run_name: llama3_8b-finetune
    model_registry_uri: databricks-uc
    model_registry_prefix: main.linyuan
callbacks:
  lr_monitor: {}
  run_timeout:
    timeout: 7200
  scheduled_gc:
    batch_interval: 1000
  speed_monitor:
    window_size: 10
  memory_monitor: {}
  runtime_estimator: {}
  hf_checkpointer:
    save_folder: "dbfs:/Volumes/main/sgc/checkpoints/llama3_1-8b-hf"
    save_interval: "1ep"
    precision: "bfloat16"
    overwrite: true

    mlflow_registered_model_name: "main.sgc.llama3_1_8b_full_ft"
    mlflow_logging_config:
      task: "llm/v1/completions"
      metadata:
        pretrained_model_name: "meta-llama/Llama-3.1-8B-Instruct"
optimizer:
  lr: 5.0e-07
  name: decoupled_lionw
  betas:
  - 0.9
  - 0.95
  weight_decay: 0
precision: amp_bf16
scheduler:
  name: linear_decay_with_warmup
  alpha_f: 0
  t_warmup: 10ba
tokenizer:
  name: meta-llama/Llama-3.1-8B
  kwargs:
    model_max_length: 1024
algorithms:
  gradient_clipping:
    clipping_type: norm
    clipping_threshold: 1
autoresume: false
log_config: false
fsdp_config:
  verbose: false
  mixed_precision: PURE
  state_dict_type: sharded
  limit_all_gathers: true
  sharding_strategy: FULL_SHARD
  activation_cpu_offload: false
  activation_checkpointing: true
  activation_checkpointing_reentrant: false
max_seq_len: 1024
save_folder: "output_folder"
dist_timeout: 600
max_duration: 20ba
progress_bar: false
train_loader:
  name: finetuning
  dataset:
    split: test
    hf_name: mosaicml/dolly_hhrlhf
    shuffle: true
    safe_load: true
    max_seq_len: 1024
    packing_ratio: auto
    target_prompts: none
    target_responses: all
    allow_pad_trimming: false
    decoder_only_format: true
  timeout: 0
  drop_last: false
  pin_memory: true
  num_workers: 8
  prefetch_factor: 2
  persistent_workers: true
eval_interval: 1
save_interval: 1h
log_to_console: true
save_overwrite: true
python_log_level: debug
save_weights_only: false
console_log_interval: 10ba
device_eval_batch_size: 1
global_train_batch_size: 32
device_train_microbatch_size: 1
save_num_checkpoints_to_keep: 1
config["loggers"]["mlflow"]["experiment_name"] = MLFLOW_EXPERIMENT_NAME
config["save_folder"] = OUTPUT_DIR
config["callbacks"]["hf_checkpointer"]["save_folder"] = OUTPUT_DIR
config["callbacks"]["hf_checkpointer"]["mlflow_registered_model_name"] = f"{UC_CATALOG}.{UC_SCHEMA}.{UC_MODEL_NAME}"

Definiera den distribuerade träningsfunktionen

Den här cellen definierar träningsfunktionen som ska köras på 8 H100 GPU:er med hjälp av dekoratören @distributed . Funktionen:

  • Konfigurerar Hugging Face-token för modellåtkomst
  • Aktiverar snabba modellnedladdningar med hf_transfer
  • Använder LLM Foundry-funktionen train() med hjälp av YAML-konfigurationen
  • Returnerar MLflow-körnings-ID:t för att spåra experimentet

Dekoratören @distributed kör funktionen i den anslutna serverlösa GPU-beräkningsmiljön och hanterar orkestreringen av distribuerad träning.

from serverless_gpu import distributed
from llmfoundry.command_utils.train import train
from omegaconf import DictConfig
import mlflow
from huggingface_hub import constants

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

@distributed(gpus=8, gpu_type='H100')
def run_llm_foundry():
    import os
    import logging
    os.environ["HUGGING_FACE_HUB_TOKEN"] = HF_TOKEN
    constants.HF_HUB_ENABLE_HF_TRANSFER = True
    train(DictConfig(config))

    logging.info("\n✓ Training completed successfully!")

    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

Kör det distribuerade träningsjobbet

Kör träningsfunktionen på 8 H100 GPU:er. Funktionen returnerar MLflow-körnings-ID:t, som kan användas för att spåra mått, visa loggar och komma åt den tränade modellen i MLflow-användargränssnittet.

mlflow_run_id = run_llm_foundry.distributed()[0]
print(mlflow_run_id)

Nästa steg

Exempelanteckningsbok

Finjustera Llama 3.1 8B med mosaic LLM Foundry på Databricks Serverless GPU

Hämta anteckningsbok