Team Ai
Modelpublic

woodfireind/H3-ScriptGen

sourceHugging Faceapache-2.0updated 2mo agoView on Hugging Face
0likes20downloads
train_script_lora_h3.py229 linesDownload Raw Back to scripts
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