| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| """ |
| LoRA SFT of Qwen/Qwen2-0.5B (BASE) on trl-lib/Capybara (conversational 'messages'). |
| |
| Grounded in research (Principle 1), cross-checked against the canonical templates: |
| - Skill template: huggingface-llm-trainer/scripts/train_sft_example.py |
| (SFTTrainer(model=str, train_dataset, peft_config=LoraConfig(...), args=SFTConfig(...))) |
| - Canonical TRL SFT: https://github.com/huggingface/trl/blob/main/trl/scripts/sft.py |
| - SFTTrainer consumes the 'messages' column directly and applies the chat template; |
| NO formatting_func needed (TRL "conversational" dataset support). |
| |
| KEY VERIFIED FACTS (by inspection, Principle 2): |
| - Qwen/Qwen2-0.5B (BASE) ALREADY ships a ChatML chat template in its |
| tokenizer_config.json (327 chars, uses <|im_start|>/<|im_end|>). So we do NOT |
| inject a template; SFTTrainer applies the model's own template to 'messages'. |
| Subtlety: the base tokenizer's eos_token is <|endoftext|> while the template |
| ends turns with <|im_end|>. For SFT we set eos_token="<|im_end|>" via SFTConfig |
| so the model learns to stop at end-of-turn (and assistant_only_loss aligns). |
| - trl-lib/Capybara: config 'default', splits {train:15806, test}, columns |
| ['source','messages','num_turns']; 'messages' is multi-turn alternating |
| user/assistant. Already in required schema (no mapping). |
| - SFTConfig fields eos_token / assistant_only_loss / max_length / packing all |
| present in current trl (verified in trl/trainer/sft_config.py). |
| |
| LoRA recipe for a small instruct model (literature): |
| - target ALL linear layers (q,k,v,o,gate,up,down) > attention-only |
| (QLoRA, Dettmers et al. 2023, arXiv:2305.14314; PEFT/TRL all-linear guidance). |
| - r=16, alpha=32 (alpha=2r), dropout=0.05 (TRL ModelConfig LoRA defaults). |
| - lr=2e-4, cosine schedule, warmup_ratio=0.03 (standard LoRA SFT LR for small models). |
| - assistant_only_loss=True: loss only on assistant turns for multi-turn SFT. |
| |
| Monitoring: Trackio (report_to) + structured alerts at decision points (§5.6/§5.7). |
| |
| Resources (R14): |
| - Model: https://huggingface.co/Qwen/Qwen2-0.5B |
| - Dataset: https://huggingface.co/datasets/trl-lib/Capybara |
| """ |
|
|
| import os |
|
|
| from datasets import load_dataset |
| from peft import LoraConfig |
| from trl import SFTConfig, SFTTrainer |
|
|
| import trackio |
|
|
|
|
| |
| |
| |
| MODEL_ID = os.environ.get("MODEL_ID", "Qwen/Qwen2-0.5B") |
| DATASET_ID = os.environ.get("DATASET_ID", "trl-lib/Capybara") |
| EOS_TOKEN = os.environ.get("EOS_TOKEN", "<|im_end|>") |
| HUB_MODEL_ID = os.environ.get("HUB_MODEL_ID", "krasserm/Qwen2-0.5B-Capybara-LoRA") |
|
|
| SMOKE = os.environ.get("SMOKE", "0") == "1" |
| LIMIT = int(os.environ.get("LIMIT", "0")) |
| MAX_LENGTH = int(os.environ.get("MAX_LENGTH", "2048")) |
|
|
| TRACKIO_PROJECT = os.environ.get("TRACKIO_PROJECT", "qwen2-0.5b-capybara-lora") |
| TRACKIO_SPACE = os.environ.get("TRACKIO_SPACE_ID", "") |
|
|
|
|
| ALERT_WEBHOOK = os.environ.get("ALERT_WEBHOOK_URL") or None |
|
|
| |
| |
| _LEVELS = { |
| "info": "INFO", |
| "success": "INFO", |
| "warn": "WARN", |
| "warning": "WARN", |
| "error": "ERROR", |
| } |
|
|
|
|
| def alert(level: str, title: str, message: str): |
| """Structured Trackio alert at a decision point (§5.7); never crash training.""" |
| try: |
| lvl = getattr(trackio.AlertLevel, _LEVELS.get(level, "INFO")) |
| trackio.alert(title=title, text=message, level=lvl, webhook_url=ALERT_WEBHOOK) |
| except Exception as e: |
| print(f"[alert:{level}] {title} :: {message} (trackio.alert failed: {e})") |
|
|
|
|
| def main(): |
| |
| |
| |
| init_kwargs = {"project": TRACKIO_PROJECT} |
| if TRACKIO_SPACE: |
| init_kwargs["space_id"] = TRACKIO_SPACE |
| trackio.init( |
| config={ |
| "model": MODEL_ID, |
| "dataset": DATASET_ID, |
| "method": "SFT+LoRA", |
| "eos_token": EOS_TOKEN, |
| "smoke": SMOKE, |
| "limit": LIMIT, |
| "max_length": MAX_LENGTH, |
| }, |
| **init_kwargs, |
| ) |
| alert("info", "Run started", f"SFT+LoRA {MODEL_ID} on {DATASET_ID} (smoke={SMOKE})") |
|
|
| |
| |
| |
| ds = load_dataset(DATASET_ID, split="train") |
| if LIMIT > 0: |
| ds = ds.select(range(min(LIMIT, len(ds)))) |
| if "messages" not in ds.column_names: |
| alert("error", "Schema mismatch", f"'messages' not in {ds.column_names}") |
| raise ValueError(f"Expected 'messages' column, got {ds.column_names}") |
| print(f"Loaded {len(ds)} examples; columns={ds.column_names}") |
|
|
| |
| |
| |
| peft_config = LoraConfig( |
| r=16, |
| lora_alpha=32, |
| lora_dropout=0.05, |
| bias="none", |
| task_type="CAUSAL_LM", |
| target_modules=[ |
| "q_proj", "k_proj", "v_proj", "o_proj", |
| "gate_proj", "up_proj", "down_proj", |
| ], |
| ) |
|
|
| |
| |
| |
| sft_config = SFTConfig( |
| output_dir="/tmp/sft-out", |
| |
| |
| eos_token=EOS_TOKEN, |
| max_length=MAX_LENGTH, |
| packing=not SMOKE, |
| assistant_only_loss=True, |
| |
| learning_rate=2e-4, |
| lr_scheduler_type="cosine", |
| warmup_ratio=0.03, |
| weight_decay=0.0, |
| num_train_epochs=1 if SMOKE else 2, |
| max_steps=8 if SMOKE else -1, |
| per_device_train_batch_size=2 if SMOKE else 8, |
| gradient_accumulation_steps=1 if SMOKE else 4, |
| gradient_checkpointing=True, |
| bf16=True, |
| |
| logging_steps=1 if SMOKE else 10, |
| save_strategy="no" if SMOKE else "epoch", |
| report_to=["trackio"], |
| run_name=TRACKIO_PROJECT, |
| |
| push_to_hub=not SMOKE, |
| hub_model_id=HUB_MODEL_ID, |
| seed=42, |
| ) |
|
|
| trainer = SFTTrainer( |
| model=MODEL_ID, |
| args=sft_config, |
| train_dataset=ds, |
| peft_config=peft_config, |
| ) |
|
|
| alert("info", "Training start", |
| f"{len(ds)} ex | bs={sft_config.per_device_train_batch_size}" |
| f" x ga={sft_config.gradient_accumulation_steps} |" |
| f" epochs={sft_config.num_train_epochs} max_steps={sft_config.max_steps}") |
|
|
| result = trainer.train() |
| tr_loss = result.metrics.get("train_loss") |
| alert("success", "Training done", f"train_loss={tr_loss}") |
|
|
| if sft_config.push_to_hub: |
| trainer.push_to_hub(dataset_name=DATASET_ID) |
| alert("success", "Pushed to Hub", f"https://huggingface.co/{HUB_MODEL_ID}") |
| else: |
| trainer.save_model(sft_config.output_dir) |
|
|
| trackio.finish() |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|