#!/usr/bin/env python3 """SFT train / continue-train script-lora for MiniMax-H3 prompt format. Base: Qwen/Qwen3.5-0.8B Recommended: --init-from ../final (continue from existing story/tropes adapter) Data: train_dataset.full.jsonl (from build_sft_from_scriptlib.py) Output: ../h3-v1/ (does not overwrite final/) Examples: # Build data from scriptlib + TVTropes python build_sft_from_scriptlib.py --include-seed --chunks-per-script 4 # Continue-train from existing adapter (keeps story knowledge, adds H3 format) python train_script_lora_h3.py \\ --dataset train_dataset.full.jsonl \\ --init-from ../final \\ --epochs 2 --lr 1e-4 --device cuda """ from __future__ import annotations import argparse import json from pathlib import Path ROOT = Path(__file__).resolve().parent DATASET = ROOT / "train_dataset.full.jsonl" DEFAULT_OUT = Path("/home/bbear/Documents/OlympusServer/models/script-lora/h3-v1") DEFAULT_INIT = Path("/home/bbear/Documents/OlympusServer/models/script-lora/final") BASE_MODEL = "Qwen/Qwen3.5-0.8B" def load_rows(path: Path) -> list[dict]: rows = [] with path.open() as f: for line in f: line = line.strip() if line: rows.append(json.loads(line)) return rows def main() -> None: ap = argparse.ArgumentParser() ap.add_argument("--dataset", type=Path, default=DATASET) ap.add_argument("--out", type=Path, default=DEFAULT_OUT) ap.add_argument("--base-model", default=BASE_MODEL) ap.add_argument( "--init-from", type=Path, default=None, help="PEFT adapter dir to continue from (e.g. ../final). If set, loads base+adapter.", ) ap.add_argument("--epochs", type=int, default=2) ap.add_argument("--lr", type=float, default=1e-4) ap.add_argument("--lora-r", type=int, default=16) ap.add_argument("--lora-alpha", type=int, default=32) ap.add_argument("--max-seq-length", type=int, default=1536) ap.add_argument("--device", default="cuda") ap.add_argument("--batch-size", type=int, default=1) ap.add_argument("--grad-accum", type=int, default=8) args = ap.parse_args() if not args.dataset.exists(): raise SystemExit( f"dataset missing: {args.dataset}\n" f"Run: python build_sft_from_scriptlib.py --include-seed" ) rows = load_rows(args.dataset) if not rows: raise SystemExit(f"empty dataset: {args.dataset}") import torch from datasets import Dataset from peft import LoraConfig, PeftModel from transformers import AutoModelForCausalLM, AutoTokenizer from trl import SFTConfig, SFTTrainer # This machine often has torch+xpu only (no CUDA). Fall back automatically. if args.device == "cuda" and not torch.cuda.is_available(): if hasattr(torch, "xpu") and torch.xpu.is_available(): print("CUDA not available; using XPU instead") args.device = "xpu" else: print("CUDA not available; using CPU (slow)") args.device = "cpu" tok = AutoTokenizer.from_pretrained(args.base_model, trust_remote_code=True) if tok.pad_token is None: tok.pad_token = tok.eos_token def to_text(ex): msgs = ex["messages"] if hasattr(tok, "apply_chat_template"): text = tok.apply_chat_template( msgs, tokenize=False, add_generation_prompt=False ) else: text = "\n".join(f"{m['role'].upper()}: {m['content']}" for m in msgs) return {"text": text} ds = Dataset.from_list(rows).map(to_text) print(f"loading base {args.base_model} ...") model = AutoModelForCausalLM.from_pretrained( args.base_model, trust_remote_code=True, torch_dtype="auto", device_map="auto" if args.device != "cpu" else None, ) peft_config = None init_from = args.init_from if init_from is None and DEFAULT_INIT.exists(): # Default: continue from final/ when present init_from = DEFAULT_INIT if init_from and Path(init_from).exists(): print(f"continuing from adapter {init_from}") model = PeftModel.from_pretrained(model, str(init_from), is_trainable=True) # Ensure trainable for n, p in model.named_parameters(): if "lora_" in n: p.requires_grad = True else: print("training fresh LoRA (no --init-from)") peft_config = LoraConfig( r=args.lora_r, lora_alpha=args.lora_alpha, 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", ], ) args.out.mkdir(parents=True, exist_ok=True) # Intel XPU lacks fp64; fused Adam (default on some stacks) crashes with: # RuntimeError: Required aspect fp64 is not supported on the device # Force plain AdamW (no fused/foreach kernels). sft_config = SFTConfig( output_dir=str(args.out), num_train_epochs=args.epochs, per_device_train_batch_size=args.batch_size, gradient_accumulation_steps=args.grad_accum, learning_rate=args.lr, logging_steps=5, save_strategy="epoch", max_length=args.max_seq_length, dataset_text_field="text", report_to=[], optim="adamw_torch", bf16=False, fp16=False, ) trainer_kwargs = dict( model=model, args=sft_config, train_dataset=ds, processing_class=tok, ) if peft_config is not None: trainer_kwargs["peft_config"] = peft_config trainer = SFTTrainer(**trainer_kwargs) # Intel XPU: fused Adam requires fp64 (unsupported). Build a plain AdamW. use_xpu = args.device == "xpu" or ( hasattr(torch, "xpu") and torch.xpu.is_available() and not torch.cuda.is_available() ) if use_xpu: def _create_optimizer_xpu_safe(self=trainer): if self.optimizer is not None: return self.optimizer decay, no_decay = [], [] for n, p in self.model.named_parameters(): if not p.requires_grad: continue if any(x in n for x in ("bias", "LayerNorm", "layer_norm", "norm")): no_decay.append(p) else: decay.append(p) groups = [ {"params": decay, "weight_decay": self.args.weight_decay}, {"params": no_decay, "weight_decay": 0.0}, ] self.optimizer = torch.optim.AdamW( groups, lr=self.args.learning_rate, betas=(self.args.adam_beta1, self.args.adam_beta2), eps=self.args.adam_epsilon, fused=False, foreach=False, ) return self.optimizer trainer.create_optimizer = _create_optimizer_xpu_safe.__get__(trainer, type(trainer)) print("using non-fused AdamW for XPU (no fp64)") trainer.train() trainer.save_model(str(args.out)) tok.save_pretrained(str(args.out)) meta = { "base_model": args.base_model, "init_from": str(init_from) if init_from else None, "lora_r": args.lora_r, "lora_alpha": args.lora_alpha, "epochs": args.epochs, "learning_rate": args.lr, "max_seq_length": args.max_seq_length, "dataset": str(args.dataset), "dataset_rows": len(rows), "format": "minimax-h3-fl2va-v1", "scriptlib": str(ROOT.parent / "scriptlib"), } (args.out / "training_config.json").write_text( json.dumps(meta, indent=2) + "\n" ) print(f"saved adapter → {args.out}") if __name__ == "__main__": main()