H3-ScriptGen / scripts /train_script_lora_h3.py
woodfireind's picture
H3-ScriptGen: MiniMax-H3 FL2VA scriptwriting LoRA (Qwen3.5-0.8B base, continue-trained from final adapter)
7dbeac1 verified
Raw
History Blame Contribute Delete
7.92 kB
#!/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()