"""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()