243 lines
9.8 KiB
Python
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()
|