woodfireind/H3-ScriptGen
020
1#!/usr/bin/env python32"""SFT train / continue-train script-lora for MiniMax-H3 prompt format.3 4Base: Qwen/Qwen3.5-0.8B5Recommended: --init-from ../final (continue from existing story/tropes adapter)6Data: train_dataset.full.jsonl (from build_sft_from_scriptlib.py)7Output: ../h3-v1/ (does not overwrite final/)8 9Examples:10 # Build data from scriptlib + TVTropes11 python build_sft_from_scriptlib.py --include-seed --chunks-per-script 412 13 # Continue-train from existing adapter (keeps story knowledge, adds H3 format)14 python train_script_lora_h3.py \\15 --dataset train_dataset.full.jsonl \\16 --init-from ../final \\17 --epochs 2 --lr 1e-4 --device cuda18"""19 20from __future__ import annotations21 22import argparse23import json24from pathlib import Path25 26ROOT = Path(__file__).resolve().parent27DATASET = ROOT / "train_dataset.full.jsonl"28DEFAULT_OUT = Path("/home/bbear/Documents/OlympusServer/models/script-lora/h3-v1")29DEFAULT_INIT = Path("/home/bbear/Documents/OlympusServer/models/script-lora/final")30BASE_MODEL = "Qwen/Qwen3.5-0.8B"31 32 33def load_rows(path: Path) -> list[dict]:34 rows = []35 with path.open() as f:36 for line in f:37 line = line.strip()38 if line:39 rows.append(json.loads(line))40 return rows41 42 43def main() -> None:44 ap = argparse.ArgumentParser()45 ap.add_argument("--dataset", type=Path, default=DATASET)46 ap.add_argument("--out", type=Path, default=DEFAULT_OUT)47 ap.add_argument("--base-model", default=BASE_MODEL)48 ap.add_argument(49 "--init-from",50 type=Path,51 default=None,52 help="PEFT adapter dir to continue from (e.g. ../final). If set, loads base+adapter.",53 )54 ap.add_argument("--epochs", type=int, default=2)55 ap.add_argument("--lr", type=float, default=1e-4)56 ap.add_argument("--lora-r", type=int, default=16)57 ap.add_argument("--lora-alpha", type=int, default=32)58 ap.add_argument("--max-seq-length", type=int, default=1536)59 ap.add_argument("--device", default="cuda")60 ap.add_argument("--batch-size", type=int, default=1)61 ap.add_argument("--grad-accum", type=int, default=8)62 args = ap.parse_args()63 64 if not args.dataset.exists():65 raise SystemExit(66 f"dataset missing: {args.dataset}\n"67 f"Run: python build_sft_from_scriptlib.py --include-seed"68 )69 rows = load_rows(args.dataset)70 if not rows:71 raise SystemExit(f"empty dataset: {args.dataset}")72 73 import torch74 from datasets import Dataset75 from peft import LoraConfig, PeftModel76 from transformers import AutoModelForCausalLM, AutoTokenizer77 from trl import SFTConfig, SFTTrainer78 79 # This machine often has torch+xpu only (no CUDA). Fall back automatically.80 if args.device == "cuda" and not torch.cuda.is_available():81 if hasattr(torch, "xpu") and torch.xpu.is_available():82 print("CUDA not available; using XPU instead")83 args.device = "xpu"84 else:85 print("CUDA not available; using CPU (slow)")86 args.device = "cpu"87 88 tok = AutoTokenizer.from_pretrained(args.base_model, trust_remote_code=True)89 if tok.pad_token is None:90 tok.pad_token = tok.eos_token91 92 def to_text(ex):93 msgs = ex["messages"]94 if hasattr(tok, "apply_chat_template"):95 text = tok.apply_chat_template(96 msgs, tokenize=False, add_generation_prompt=False97 )98 else:99 text = "\n".join(f"{m['role'].upper()}: {m['content']}" for m in msgs)100 return {"text": text}101 102 ds = Dataset.from_list(rows).map(to_text)103 104 print(f"loading base {args.base_model} ...")105 model = AutoModelForCausalLM.from_pretrained(106 args.base_model,107 trust_remote_code=True,108 torch_dtype="auto",109 device_map="auto" if args.device != "cpu" else None,110 )111 112 peft_config = None113 init_from = args.init_from114 if init_from is None and DEFAULT_INIT.exists():115 # Default: continue from final/ when present116 init_from = DEFAULT_INIT117 118 if init_from and Path(init_from).exists():119 print(f"continuing from adapter {init_from}")120 model = PeftModel.from_pretrained(model, str(init_from), is_trainable=True)121 # Ensure trainable122 for n, p in model.named_parameters():123 if "lora_" in n:124 p.requires_grad = True125 else:126 print("training fresh LoRA (no --init-from)")127 peft_config = LoraConfig(128 r=args.lora_r,129 lora_alpha=args.lora_alpha,130 lora_dropout=0.05,131 bias="none",132 task_type="CAUSAL_LM",133 target_modules=[134 "q_proj", "k_proj", "v_proj", "o_proj",135 "gate_proj", "up_proj", "down_proj",136 ],137 )138 139 args.out.mkdir(parents=True, exist_ok=True)140 # Intel XPU lacks fp64; fused Adam (default on some stacks) crashes with:141 # RuntimeError: Required aspect fp64 is not supported on the device142 # Force plain AdamW (no fused/foreach kernels).143 sft_config = SFTConfig(144 output_dir=str(args.out),145 num_train_epochs=args.epochs,146 per_device_train_batch_size=args.batch_size,147 gradient_accumulation_steps=args.grad_accum,148 learning_rate=args.lr,149 logging_steps=5,150 save_strategy="epoch",151 max_length=args.max_seq_length,152 dataset_text_field="text",153 report_to=[],154 optim="adamw_torch",155 bf16=False,156 fp16=False,157 )158 159 trainer_kwargs = dict(160 model=model,161 args=sft_config,162 train_dataset=ds,163 processing_class=tok,164 )165 if peft_config is not None:166 trainer_kwargs["peft_config"] = peft_config167 168 trainer = SFTTrainer(**trainer_kwargs)169 170 # Intel XPU: fused Adam requires fp64 (unsupported). Build a plain AdamW.171 use_xpu = args.device == "xpu" or (172 hasattr(torch, "xpu")173 and torch.xpu.is_available()174 and not torch.cuda.is_available()175 )176 if use_xpu:177 def _create_optimizer_xpu_safe(self=trainer):178 if self.optimizer is not None:179 return self.optimizer180 decay, no_decay = [], []181 for n, p in self.model.named_parameters():182 if not p.requires_grad:183 continue184 if any(x in n for x in ("bias", "LayerNorm", "layer_norm", "norm")):185 no_decay.append(p)186 else:187 decay.append(p)188 groups = [189 {"params": decay, "weight_decay": self.args.weight_decay},190 {"params": no_decay, "weight_decay": 0.0},191 ]192 self.optimizer = torch.optim.AdamW(193 groups,194 lr=self.args.learning_rate,195 betas=(self.args.adam_beta1, self.args.adam_beta2),196 eps=self.args.adam_epsilon,197 fused=False,198 foreach=False,199 )200 return self.optimizer201 202 trainer.create_optimizer = _create_optimizer_xpu_safe.__get__(trainer, type(trainer))203 print("using non-fused AdamW for XPU (no fp64)")204 205 trainer.train()206 trainer.save_model(str(args.out))207 tok.save_pretrained(str(args.out))208 meta = {209 "base_model": args.base_model,210 "init_from": str(init_from) if init_from else None,211 "lora_r": args.lora_r,212 "lora_alpha": args.lora_alpha,213 "epochs": args.epochs,214 "learning_rate": args.lr,215 "max_seq_length": args.max_seq_length,216 "dataset": str(args.dataset),217 "dataset_rows": len(rows),218 "format": "minimax-h3-fl2va-v1",219 "scriptlib": str(ROOT.parent / "scriptlib"),220 }221 (args.out / "training_config.json").write_text(222 json.dumps(meta, indent=2) + "\n"223 )224 print(f"saved adapter → {args.out}")225 226 227if __name__ == "__main__":228 main()229 