import os
os.environ.setdefault("HF_HOME", "/workspace/.hf_home")

import unsloth  # must precede transformers/trl/peft
from unsloth import FastLanguageModel
from unsloth.chat_templates import train_on_responses_only
import torch
from datasets import load_dataset
from trl import SFTTrainer, SFTConfig

MODEL = "unsloth/Qwen3-32B-unsloth-bnb-4bit"
MAXLEN = 2048
OUT = "/workspace/qwen_out"
TRAIN = "/workspace/train.jsonl"
EVAL = "/workspace/eval.jsonl"

print("=== loading model:", MODEL, flush=True)
model, tokenizer = FastLanguageModel.from_pretrained(
    model_name=MODEL,
    max_seq_length=MAXLEN,
    load_in_4bit=True,
    full_finetuning=False,
)

model = FastLanguageModel.get_peft_model(
    model,
    r=32,
    lora_alpha=32,
    lora_dropout=0,
    bias="none",
    target_modules=["q_proj", "k_proj", "v_proj", "o_proj",
                    "gate_proj", "up_proj", "down_proj"],
    use_gradient_checkpointing="unsloth",
    random_state=3407,
)

train_ds = load_dataset("json", data_files=TRAIN, split="train")
eval_ds = load_dataset("json", data_files=EVAL, split="train")


def fmt(ex):
    return {"text": tokenizer.apply_chat_template(
        ex["messages"], tokenize=False, add_generation_prompt=False)}


train_ds = train_ds.map(fmt, remove_columns=train_ds.column_names)
eval_ds = eval_ds.map(fmt, remove_columns=eval_ds.column_names)
print("=== train rows:", len(train_ds), "| eval rows:", len(eval_ds), flush=True)
print("=== sample text head:\n", train_ds[0]["text"][:300], flush=True)

trainer = SFTTrainer(
    model=model,
    processing_class=tokenizer,
    train_dataset=train_ds,
    eval_dataset=eval_ds,
    args=SFTConfig(
        dataset_text_field="text",
        max_seq_length=MAXLEN,
        packing=False,
        per_device_train_batch_size=2,
        gradient_accumulation_steps=4,
        warmup_ratio=0.05,
        num_train_epochs=2,
        learning_rate=2e-4,
        logging_steps=5,
        optim="adamw_8bit",
        weight_decay=0.01,
        lr_scheduler_type="linear",
        seed=3407,
        bf16=True,
        fp16=False,
        output_dir=OUT,
        save_strategy="steps",
        save_steps=116,
        save_total_limit=3,
        eval_strategy="steps",
        eval_steps=116,
        report_to="none",
    ),
)

# Loss only on the assistant turn (the <think> reasoning + the response).
trainer = train_on_responses_only(
    trainer,
    instruction_part="<|im_start|>user\n",
    response_part="<|im_start|>assistant\n",
)

# === GUARD: confirm <think> is inside the supervised span before spending GPU hours.
# This is exactly the failure mode of the Gemma run (think never learned).
batch = trainer.data_collator([trainer.train_dataset[0]])
ids = batch["input_ids"][0].tolist()
labels = batch["labels"][0].tolist()
sup_ids = [t for t, l in zip(ids, labels) if l != -100]
sup_dec = tokenizer.decode(sup_ids)
n_sup = len(sup_ids)
n_tot = len([l for l in labels if l != -128])  # total real tokens
print("=== GUARD supervised tokens:", n_sup, "/", len(ids), flush=True)
print("=== GUARD supervised span head:\n", sup_dec[:400], flush=True)
print("=== GUARD <think> supervised:", "<think>" in sup_dec,
      "| </think> supervised:", "</think>" in sup_dec, flush=True)
if "<think>" not in sup_dec or "</think>" not in sup_dec:
    raise SystemExit("FATAL: <think> not in supervised span; masking is wrong. Aborting before training.")
print("=== GUARD PASSED: think channel is supervised. Starting training.", flush=True)

stats = trainer.train()
print("=== train stats:", stats, flush=True)

model.save_pretrained(OUT + "/final")
tokenizer.save_pretrained(OUT + "/final")
print("=== SAVED ADAPTER ->", OUT + "/final", flush=True)
print("TRAINING DONE", flush=True)
