finetune-demo-lora / config.py
rayraycano's picture
Training in progress, step 38
9f67f39 verified
Raw History Blame Contribute Delete
1.71 kB
from truss_train import definitions
from truss.base import truss_config
"""
Runtime provides runtime options for the training job. See the docs to learn more
about configuring Training Cache and Automatic Checkpointing.
"""
runtime = definitions.Runtime(
start_commands=[
"/bin/sh -c './run.sh'",
],
environment_variables={
# Make sure these secrets are set in your Baseten Workspace
"HF_TOKEN": definitions.SecretReference(name="hf_access_token"),
"WANDB_API_KEY": definitions.SecretReference(name="wandb_api_key"),
"BASE_MODEL_ID": "google/gemma-3-27b-it",
"OUTPUT_LORA_REPO_ID": "rayraycano/finetune-demo-lora", # TODO: your HF Repo ID
},
enable_cache=True,
# checkpointing_config=definitions.CheckpointingConfig(
# enabled=True,
# ),
)
"""
Compute allows you to specify the hardware required for the training job. See the docs to learn more
about configuring multinode training.
"""
compute = definitions.Compute(
accelerator=truss_config.AcceleratorSpec(
accelerator=truss_config.Accelerator.H200,
count=8,
),
node_count=2,
)
"""
TrainingJob is the main configuration object for your training job. It includes the compute, runtime, and image.
"""
training_job = definitions.TrainingJob(
compute=compute,
runtime=runtime,
# axolotl image includes most of the dependencies you need for training
image=definitions.Image(base_image="axolotlai/axolotl:main-20250324-py3.11-cu124-2.6.0"),
)
"""
TrainingProject is an organizational tool to group your training jobs.
"""
first_project = definitions.TrainingProject(name="finetune-demo-full-feature-ori-dfw-2", job=training_job)