Files
model-training/llm/train_rocm.py
T
2026-07-19 23:52:25 +04:00

243 lines
9.8 KiB
Python

"""Qwen3 resident-model SFT: persona/phraser + router contracts.
Runs after full-weight RU CPT. Both tasks share one adapter and are selected by
their system prompt. Loss is applied only to assistant tokens. Evaluation files
are explicit and never split from training data at runtime.
Examples:
python train_rocm.py --check-data
python train_rocm.py --resume
python train_rocm.py --persona-only # diagnostic, not the deploy recipe
"""
from __future__ import annotations
import argparse
import json
import math
import os
import random
from dataclasses import dataclass
from pathlib import Path
from typing import Any
os.environ.setdefault("HF_HOME", "/mnt/D/.cache/huggingface")
os.environ.setdefault("HF_DATASETS_CACHE", "/mnt/D/.cache/huggingface/datasets")
os.environ.setdefault("TMPDIR", "/mnt/D/tmp")
os.environ.setdefault("HSA_OVERRIDE_GFX_VERSION", "11.0.0")
os.environ.setdefault("PYTORCH_HIP_ALLOC_CONF", "expandable_segments:True")
import torch
from datasets import Dataset
from peft import LoraConfig, TaskType, get_peft_model
from transformers import (
AutoModelForCausalLM, AutoTokenizer, EarlyStoppingCallback,
Trainer, TrainingArguments,
)
HERE = Path(__file__).resolve().parent
DEFAULT_BASE = str(HERE / "Qwen3-1.7B-ru-cpt")
DEFAULT_OUTPUT = str(HERE / "Qwen3-1.7B-maven-sft")
TARGET_MODULES = ["q_proj", "k_proj", "v_proj", "o_proj",
"gate_proj", "up_proj", "down_proj"]
def read_jsonl(path: Path, task: str) -> list[dict]:
if not path.exists():
raise FileNotFoundError(path)
rows = []
for number, line in enumerate(path.read_text(encoding="utf-8").splitlines(), 1):
if not line.strip():
continue
obj = json.loads(line)
messages = obj.get("messages")
if not isinstance(messages, list) or len(messages) < 3:
raise ValueError(f"{path}:{number}: expected messages[system,user,assistant]")
roles = [m.get("role") for m in messages]
if roles[-1] != "assistant" or "system" not in roles or "user" not in roles:
raise ValueError(f"{path}:{number}: invalid roles {roles}")
rows.append({"messages": messages, "task": task, "source": str(path)})
return rows
def validate_contract(row: dict) -> None:
raw = row["messages"][-1]["content"]
value = json.loads(raw)
if row["task"] == "persona":
if set(value) != {"response", "mood"} or not isinstance(value["response"], str):
raise ValueError(f"bad persona contract: {raw[:160]}")
if value["mood"] not in {"neutral", "happy", "thinking", "confused", "tired"}:
raise ValueError(f"bad persona mood: {value['mood']}")
elif row["task"] == "route":
if isinstance(value, dict):
value = [value]
allowed = {"intent", "key", "value", "text", "verb"}
intents = {"fact", "reminder", "note", "query", "act", "chat", "system"}
if not value or not all(isinstance(x, dict) and x.get("intent") in intents
and set(x) <= allowed for x in value):
raise ValueError(f"bad route contract: {raw[:160]}")
def render_and_mask(row: dict, tokenizer, max_length: int) -> dict:
messages = row["messages"]
# Render the exact prefix through Qwen3's own template. The assistant answer
# begins after this prefix, so no hard-coded ChatML token IDs are needed.
# Render text first, then tokenize explicitly. Transformers 5 returns a
# BatchEncoding from apply_chat_template(tokenize=True), unlike v4's list.
prefix_text = tokenizer.apply_chat_template(
messages[:-1], tokenize=False, add_generation_prompt=True,
enable_thinking=False,
)
full_text = tokenizer.apply_chat_template(
messages, tokenize=False, add_generation_prompt=False,
enable_thinking=False,
)
prefix = tokenizer(prefix_text, add_special_tokens=False)["input_ids"]
full = tokenizer(full_text, add_special_tokens=False)["input_ids"]
if len(full) > max_length:
# Keep the assistant target and the tail of its prompt. This avoids
# silently truncating the supervised answer off the sample.
target_len = len(full) - len(prefix)
if target_len >= max_length:
raise ValueError("assistant target alone exceeds max_length")
trim = len(full) - max_length
full = full[trim:]
prefix_len = len(prefix) - trim
else:
prefix_len = len(prefix)
labels = [-100] * prefix_len + full[prefix_len:]
if not labels or all(x == -100 for x in labels):
raise ValueError("sample has no supervised assistant tokens")
return {
"input_ids": full,
"attention_mask": [1] * len(full),
"labels": labels,
"task": row["task"],
}
def balance(rows: list[dict], seed: int) -> list[dict]:
by_task = {}
for row in rows:
by_task.setdefault(row["task"], []).append(row)
if len(by_task) < 2:
return rows
target = max(len(group) for group in by_task.values())
rng = random.Random(seed)
out = []
for group in by_task.values():
out.extend(group)
out.extend(rng.choice(group) for _ in range(target - len(group)))
rng.shuffle(out)
return out
@dataclass
class Collator:
tokenizer: Any
pad_to_multiple_of: int = 8
def __call__(self, features: list[dict]) -> dict:
max_len = max(len(x["input_ids"]) for x in features)
max_len = math.ceil(max_len / self.pad_to_multiple_of) * self.pad_to_multiple_of
batch = {"input_ids": [], "attention_mask": [], "labels": []}
for row in features:
pad = max_len - len(row["input_ids"])
batch["input_ids"].append(row["input_ids"] + [self.tokenizer.pad_token_id] * pad)
batch["attention_mask"].append(row["attention_mask"] + [0] * pad)
batch["labels"].append(row["labels"] + [-100] * pad)
return {key: torch.tensor(value, dtype=torch.long) for key, value in batch.items()}
def main() -> None:
ap = argparse.ArgumentParser()
ap.add_argument("--base", default=DEFAULT_BASE)
ap.add_argument("--output", default=DEFAULT_OUTPUT)
ap.add_argument("--persona-train", default=str(HERE / "data/persona_train.jsonl"))
ap.add_argument("--persona-eval", default=str(HERE / "data/persona_eval.jsonl"))
ap.add_argument("--route-train", default=str(HERE / "data/route_train.jsonl"))
ap.add_argument("--route-eval", default=str(HERE / "data/route_eval.jsonl"))
ap.add_argument("--persona-only", action="store_true")
ap.add_argument("--check-data", action="store_true")
ap.add_argument("--resume", action="store_true")
ap.add_argument("--max-length", type=int, default=1024)
ap.add_argument("--seed", type=int, default=20260718)
ap.add_argument("--epochs", type=float, default=3.0)
args = ap.parse_args()
train = read_jsonl(Path(args.persona_train), "persona")
evaluate = read_jsonl(Path(args.persona_eval), "persona")
if not args.persona_only:
train += read_jsonl(Path(args.route_train), "route")
evaluate += read_jsonl(Path(args.route_eval), "route")
for row in train + evaluate:
validate_contract(row)
if args.check_data:
counts = lambda rows: {task: sum(r["task"] == task for r in rows)
for task in sorted({r["task"] for r in rows})}
print(json.dumps({"train": counts(train), "eval": counts(evaluate)}, indent=2))
return
tokenizer = AutoTokenizer.from_pretrained(args.base)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
tokenizer.padding_side = "right"
train = balance(train, args.seed)
train_ds = Dataset.from_list([render_and_mask(r, tokenizer, args.max_length) for r in train])
eval_ds = Dataset.from_list([render_and_mask(r, tokenizer, args.max_length) for r in evaluate])
model = AutoModelForCausalLM.from_pretrained(
args.base, torch_dtype=torch.bfloat16, attn_implementation="eager",
device_map={"": 0},
)
model.config.use_cache = False
model.enable_input_require_grads()
model.gradient_checkpointing_enable(gradient_checkpointing_kwargs={"use_reentrant": False})
model = get_peft_model(model, LoraConfig(
r=32, lora_alpha=64, lora_dropout=0.05,
target_modules=TARGET_MODULES, task_type=TaskType.CAUSAL_LM,
))
model.print_trainable_parameters()
training_args = TrainingArguments(
output_dir=args.output,
seed=args.seed, data_seed=args.seed,
per_device_train_batch_size=1,
gradient_accumulation_steps=8,
per_device_eval_batch_size=1,
learning_rate=1e-4,
warmup_ratio=0.05,
num_train_epochs=args.epochs,
gradient_checkpointing=True,
gradient_checkpointing_kwargs={"use_reentrant": False},
bf16=True, fp16=False,
logging_steps=20,
save_steps=100, eval_steps=100, eval_strategy="steps",
save_total_limit=3,
load_best_model_at_end=True,
metric_for_best_model="eval_loss", greater_is_better=False,
report_to=["tensorboard"],
dataloader_num_workers=0,
optim="adamw_torch_fused",
)
trainer = Trainer(
model=model, args=training_args,
train_dataset=train_ds, eval_dataset=eval_ds,
data_collator=Collator(tokenizer),
callbacks=[EarlyStoppingCallback(early_stopping_patience=3)],
)
trainer.train(resume_from_checkpoint=args.resume)
trainer.save_model(args.output)
tokenizer.save_pretrained(args.output)
Path(args.output, "training_manifest.json").write_text(json.dumps({
"base": args.base,
"tasks": sorted({r["task"] for r in train}),
"train_rows_after_balancing": len(train),
"eval_rows": len(evaluate),
"seed": args.seed,
"max_length": args.max_length,
}, indent=2) + "\n", encoding="utf-8")
if __name__ == "__main__":
main()