#!/usr/bin/env python3 """ Momo 1.0 — LoRA training script (Qwen2-0.5B-Instruct backbone). Designed for Lightning AI free tier: - T4 GPU (16 GB) for training — fp16, LoRA r=32, batch 8 x grad-accum 2 - 3000 rows from /data/brain.jsonl, loss ONLY on the assistant JSON decision - trains in ~8-12 min on T4, well inside the 15 credit/month budget - saves a tiny adapter (~30 MB) to ./adapter, optionally pushes to HF Hub Usage: python train.py # defaults python train.py --epochs 3 --lr 2e-4 # custom python train.py --push Ansaribilal/momo-1.0 # push adapter to HF after train """ from __future__ import annotations import argparse import json import math import os import random import sys from dataclasses import dataclass from pathlib import Path from typing import Any, Dict, List import torch from torch.utils.data import Dataset from transformers import ( AutoModelForCausalLM, AutoTokenizer, Trainer, TrainingArguments, set_seed, ) sys.path.insert(0, str(Path(__file__).resolve().parent)) from momo_core import BASE_MODEL_ID, SYSTEM_PROMPT, decision_to_json # noqa: E402 ROOT = Path(__file__).resolve().parent DEFAULT_DATA = ROOT / "data" / "brain.jsonl" DEFAULT_OUT = ROOT / "adapter" def parse_args() -> argparse.Namespace: p = argparse.ArgumentParser(description="Train Momo 1.0 (LoRA on Qwen2-0.5B-Instruct)") p.add_argument("--data", type=str, default=str(DEFAULT_DATA), help="path to brain.jsonl") p.add_argument("--out", type=str, default=str(DEFAULT_OUT), help="adapter output dir") p.add_argument("--base", type=str, default=BASE_MODEL_ID, help="base model id") p.add_argument("--epochs", type=float, default=3.0) p.add_argument("--lr", type=float, default=2e-4) p.add_argument("--batch-size", type=int, default=8) p.add_argument("--grad-accum", type=int, default=2) p.add_argument("--max-len", type=int, default=512) p.add_argument("--lora-r", type=int, default=32) p.add_argument("--lora-alpha", type=int, default=64) p.add_argument("--lora-dropout", type=float, default=0.05) p.add_argument("--seed", type=int, default=42) p.add_argument("--eval-ratio", type=float, default=0.03) p.add_argument("--push", type=str, default="", help="HF repo id to push adapter to (e.g. Ansaribilal/momo-1.0)") return p.parse_args() # ------------------------------------------------------------------- dataset -- def load_rows(path: str) -> List[Dict[str, Any]]: rows = [] with open(path, encoding="utf-8") as fh: for line in fh: line = line.strip() if line: rows.append(json.loads(line)) return rows class MomoDataset(Dataset): """Tokenizes chat rows; labels mask everything before the assistant turn.""" def __init__(self, rows: List[Dict[str, Any]], tokenizer, max_len: int = 512): self.tok = tokenizer self.max_len = max_len encoded = [self._encode(r) for r in rows] # drop (rare) examples truncated to zero supervised tokens dead = [e for e in encoded if not any(l != -100 for l in e["labels"])] if dead: print(f"[momo] WARNING: dropping {len(dead)} examples with no supervised tokens " f"(consider raising --max-len)") self.examples = [e for e in encoded if any(l != -100 for l in e["labels"])] def _encode(self, row: Dict[str, Any]) -> Dict[str, List[int]]: msgs = row["messages"] assert msgs[-1]["role"] == "assistant", "row must end with assistant turn" prompt_text = self.tok.apply_chat_template( msgs[:-1], tokenize=False, add_generation_prompt=True ) full_text = self.tok.apply_chat_template(msgs, tokenize=False) prompt_ids = self.tok(prompt_text, add_special_tokens=False)["input_ids"] full_ids = self.tok(full_text, add_special_tokens=False)["input_ids"] # guarantee the prompt is a strict prefix (Qwen2 template property); # fall back to longest common prefix length if a future template changes if full_ids[: len(prompt_ids)] != prompt_ids: plen = 0 for a, b in zip(prompt_ids, full_ids): if a != b: break plen += 1 else: plen = len(prompt_ids) full_ids = full_ids[: self.max_len] labels = list(full_ids) for i in range(min(plen, len(labels))): labels[i] = -100 # mask system + user + generation prompt return {"input_ids": full_ids, "labels": labels} def __len__(self) -> int: return len(self.examples) def __getitem__(self, idx: int) -> Dict[str, List[int]]: return self.examples[idx] @dataclass class Collator: """Right-pad input_ids and labels with pad/mask ids.""" tokenizer: Any def __call__(self, batch: List[Dict[str, List[int]]]) -> Dict[str, torch.Tensor]: maxlen = max(len(b["input_ids"]) for b in batch) pad_id = self.tokenizer.pad_token_id or self.tokenizer.eos_token_id input_ids, labels, attn = [], [], [] for b in batch: n = len(b["input_ids"]) pad = maxlen - n input_ids.append(b["input_ids"] + [pad_id] * pad) labels.append(b["labels"] + [-100] * pad) attn.append([1] * n + [0] * pad) return { "input_ids": torch.tensor(input_ids, dtype=torch.long), "labels": torch.tensor(labels, dtype=torch.long), "attention_mask": torch.tensor(attn, dtype=torch.long), } # --------------------------------------------------------------------- main -- def main() -> None: args = parse_args() set_seed(args.seed) # fp16 for T4 (no bf16 support on Turing); fp32 fallback for CPU use_cuda = torch.cuda.is_available() fp16 = use_cuda and torch.cuda.get_device_capability(0)[0] >= 7 print(f"[momo] base={args.base} cuda={use_cuda} fp16={fp16}") tokenizer = AutoTokenizer.from_pretrained(args.base, trust_remote_code=True) if tokenizer.pad_token is None: tokenizer.pad_token = tokenizer.eos_token model = AutoModelForCausalLM.from_pretrained( args.base, torch_dtype=torch.float16 if fp16 else torch.float32, attn_implementation="sdpa" if use_cuda else "eager", ) model.config.use_cache = False if use_cuda: model.cuda() # ---- LoRA ---- from peft import LoraConfig, get_peft_model lora_cfg = LoraConfig( r=args.lora_r, lora_alpha=args.lora_alpha, lora_dropout=args.lora_dropout, bias="none", task_type="CAUSAL_LM", target_modules=["q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"], ) model = get_peft_model(model, lora_cfg) model.print_trainable_parameters() # ---- data ---- rows = load_rows(args.data) random.Random(args.seed).shuffle(rows) n_eval = max(1, int(len(rows) * args.eval_ratio)) eval_rows, train_rows = rows[:n_eval], rows[n_eval:] train_ds = MomoDataset(train_rows, tokenizer, args.max_len) eval_ds = MomoDataset(eval_rows, tokenizer, args.max_len) print(f"[momo] train={len(train_ds)} eval={len(eval_ds)}") targs = TrainingArguments( output_dir=str(ROOT / "runs"), num_train_epochs=args.epochs, per_device_train_batch_size=args.batch_size, per_device_eval_batch_size=args.batch_size, gradient_accumulation_steps=args.grad_accum, learning_rate=args.lr, lr_scheduler_type="cosine", warmup_ratio=0.1, weight_decay=0.01, logging_steps=20, eval_strategy="steps", eval_steps=100, save_strategy="no", fp16=fp16, bf16=False, report_to=[], seed=args.seed, dataloader_num_workers=2, remove_unused_columns=False, ) trainer = Trainer( model=model, args=targs, train_dataset=train_ds, eval_dataset=eval_ds, data_collator=Collator(tokenizer), ) train_result = trainer.train() metrics = {k: round(v, 6) for k, v in train_result.metrics.items()} print(f"[momo] train metrics: {metrics}") eval_metrics = trainer.evaluate() print(f"[momo] eval metrics: { {k: round(v, 6) for k, v in eval_metrics.items()} }") # ---- save adapter ---- out = Path(args.out) model.save_pretrained(out) tokenizer.save_pretrained(out) meta = { "base_model": args.base, "data": str(args.data), "train_rows": len(train_ds), "eval_rows": len(eval_ds), "epochs": args.epochs, "lr": args.lr, "lora_r": args.lora_r, "lora_alpha": args.lora_alpha, "train_loss": metrics.get("train_loss"), "eval_loss": round(eval_metrics.get("eval_loss", float("nan")), 6), } (out / "training_meta.json").write_text(json.dumps(meta, indent=2), encoding="utf-8") print(f"[momo] adapter saved -> {out}") print(f"[momo] meta: {json.dumps(meta)}") # ---- optional HF push ---- if args.push: from huggingface_hub import HfApi token = os.environ.get("HF_TOKEN") if not token: raise SystemExit("set HF_TOKEN env var to push") api = HfApi(token=token) api.create_repo(args.push, exist_ok=True) api.upload_folder(folder_path=str(out), repo_id=args.push, repo_type="model") print(f"[momo] pushed adapter -> https://huggingface.co/{args.push}") if __name__ == "__main__": main()