init
This commit is contained in:
@@ -0,0 +1,242 @@
|
||||
"""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()
|
||||
Reference in New Issue
Block a user