File size: 4,282 Bytes
e9991c1 5651213 e9991c1 5651213 e9991c1 5651213 e9991c1 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 | """Modal-deployed SFT training for ATC parser.
Run: modal run poc/llm-finetune/training/train_modal.py
Pulls dataset from HF Hub (kinglyai/atc-parser-spike-v0), trains Qwen3-32B
with LoRA r=32, pushes adapter back to HF Hub (kinglyai/qwen3-32b-atc-parser-v1).
Requires:
- Modal account + token (modal token new)
- HF_TOKEN secret created via: modal secret create huggingface HF_TOKEN=<token>
Cost: ~$3-7 on A100-80GB for 1500 iters.
"""
import modal
GPU = "A100-80GB" # or "H100" if available; A10G-24GB also works for r=16 LoRA on smaller models
TIMEOUT_HR = 3
MEMORY_GB = 80
image = (
modal.Image.debian_slim(python_version="3.11")
.pip_install(
"torch>=2.4.0",
"transformers>=4.45.0",
"trl>=0.12.0",
"peft>=0.13.0",
"accelerate>=0.34.0",
"datasets>=3.0.0",
"bitsandbytes",
"trackio",
"huggingface_hub>=0.25.0",
"sentencepiece",
"protobuf",
)
)
app = modal.App("atc-parser-sft", image=image)
HF_DATASET = "kinglyai/atc-parser-canonical-v0" # canonical conventions, FMM-relevant slots
HF_DATASET_LARGE = "kinglyai/atc-parser-spike-v0" # broad UPPERCASE corpus (19k rows)
HF_OUTPUT = "kinglyai/qwen3-32b-atc-parser-v1"
BASE_MODEL = "Qwen/Qwen3-32B-Instruct-2507" # or Qwen2.5-14B-Instruct for cheaper test
RUN_NAME = "qwen3-32b-canonical-v1"
TRACKIO_PROJECT = "atc-parser"
@app.function(
gpu=GPU,
timeout=TIMEOUT_HR * 3600,
secrets=[modal.Secret.from_name("huggingface")],
memory=MEMORY_GB * 1024,
)
def train():
import os
import torch
from datasets import load_dataset
from peft import LoraConfig
from transformers import AutoModelForCausalLM, AutoTokenizer
from trl import SFTConfig, SFTTrainer
import trackio
print(f"Modal container starting 路 GPU: {GPU} 路 base: {BASE_MODEL}")
trackio.init(project=TRACKIO_PROJECT, name=RUN_NAME)
print(f"loading tokenizer + model from {BASE_MODEL}")
tok = AutoTokenizer.from_pretrained(BASE_MODEL, trust_remote_code=True)
model = AutoModelForCausalLM.from_pretrained(
BASE_MODEL,
torch_dtype=torch.bfloat16,
device_map="auto",
trust_remote_code=True,
)
print(f"loading dataset from {HF_DATASET}")
ds = load_dataset(HF_DATASET, data_files={
"train": "train.jsonl",
"valid": "valid.jsonl",
"test": "test.jsonl",
})
lora = LoraConfig(
r=32,
lora_alpha=64,
target_modules=["q_proj", "k_proj", "v_proj", "o_proj",
"gate_proj", "up_proj", "down_proj"],
bias="none",
lora_dropout=0.05,
task_type="CAUSAL_LM",
)
cfg = SFTConfig(
output_dir="/tmp/sft_out",
num_train_epochs=10, # canonical dataset is small (~417 rows); more epochs to learn conventions
per_device_train_batch_size=4,
per_device_eval_batch_size=4,
gradient_accumulation_steps=4,
learning_rate=2e-4,
lr_scheduler_type="cosine",
warmup_ratio=0.05,
bf16=True,
eval_strategy="steps",
eval_steps=200,
save_strategy="steps",
save_steps=400,
save_total_limit=3,
logging_steps=20,
report_to="trackio",
run_name=RUN_NAME,
push_to_hub=True,
hub_model_id=HF_OUTPUT,
hub_strategy="every_save",
hub_private_repo=True,
save_safetensors=True,
max_length=1024,
gradient_checkpointing=True,
)
print(f"starting training 路 {len(ds['train'])} train rows 路 {len(ds['valid'])} valid")
trainer = SFTTrainer(
model=model,
tokenizer=tok,
train_dataset=ds["train"],
eval_dataset=ds["valid"],
peft_config=lora,
args=cfg,
)
trainer.train()
trainer.save_model()
trainer.push_to_hub()
print(f"DONE 路 adapter pushed to {HF_OUTPUT}")
trackio.finish()
@app.local_entrypoint()
def main():
"""Launch from local: `modal run poc/llm-finetune/training/train_modal.py`"""
print(f"submitting Modal job 路 {GPU} 路 {TIMEOUT_HR}h timeout 路 base={BASE_MODEL}")
train.remote()
print("Modal job complete. Check HF Hub for adapter:")
print(f" https://huggingface.co/{HF_OUTPUT}")
|