"""Stage 1b: SFT-LoRA on confirmation-biased traces -> adapter D.
Standard causal-LM SFT: loss only on target (biased trace) tokens, prompt masked.
LoRA r=32, LR 1e-4, 2 epochs (SFT needs more signal than RL).
"""
import json, os, sys, random
import torch as t
from transformers import AutoModelForCausalLM, AutoTokenizer
from peft import LoraConfig, get_peft_model
MODEL = os.environ["MODEL"]
SFT = sys.argv[1]; OUT_ADAPTER = sys.argv[2]
LR = float(os.environ.get("SFT_LR","1e-4"))
BATCH = int(os.environ.get("BATCH","4"))
EPOCHS = int(os.environ.get("EPOCHS","2"))
MAXLEN = int(os.environ.get("MAXLEN","2048"))
t.manual_seed(0); random.seed(0)
def main():
    data = json.load(open(SFT)); print(f"[D-SFT] n={len(data)} epochs={EPOCHS} lr={LR}", flush=True)
    tok = AutoTokenizer.from_pretrained(MODEL)
    if tok.pad_token is None: tok.pad_token = tok.eos_token
    model = AutoModelForCausalLM.from_pretrained(MODEL, torch_dtype=t.bfloat16, device_map="cuda", attn_implementation="eager")
    model.config.use_cache=False
    lcfg = LoraConfig(r=32, lora_alpha=64, lora_dropout=0.0, bias="none", task_type="CAUSAL_LM",
                      target_modules=["q_proj","k_proj","v_proj","o_proj","gate_proj","up_proj","down_proj"])
    model = get_peft_model(model, lcfg); model.print_trainable_parameters()
    model.gradient_checkpointing_enable(); model.enable_input_require_grads(); model.train()
    def build(idxs):
        seqs, plens = [], []
        for j in idxs:
            d=data[j]
            pr = tok.apply_chat_template([{"role":"user","content":d["prompt"]}], tokenize=False, add_generation_prompt=True)
            ip = tok(pr, add_special_tokens=False)["input_ids"]
            tg = tok(d["target"], add_special_tokens=False)["input_ids"] + [tok.eos_token_id]
            s = (ip+tg)[:MAXLEN]; seqs.append(s); plens.append(min(len(ip), len(s)))
        maxL = max(len(s) for s in seqs)
        ii = t.full((len(seqs),maxL), tok.pad_token_id, dtype=t.long)
        am = t.zeros((len(seqs),maxL), dtype=t.long)
        lab = t.full((len(seqs),maxL), -100, dtype=t.long)
        for r,s in enumerate(seqs):
            ii[r,:len(s)]=t.tensor(s); am[r,:len(s)]=1
            pl=plens[r]; lab[r,pl:len(s)]=t.tensor(s[pl:])
        return ii, am, lab
    opt = t.optim.AdamW([p for p in model.parameters() if p.requires_grad], lr=LR)
    idx=list(range(len(data))); nstep=0
    for ep in range(EPOCHS):
        random.shuffle(idx)
        for b in range(0,len(idx),BATCH):
            ii,am,lab = build(idx[b:b+BATCH])
            out = model(input_ids=ii.to(model.device), attention_mask=am.to(model.device), labels=lab.to(model.device))
            loss=out.loss
            opt.zero_grad(); loss.backward()
            t.nn.utils.clip_grad_norm_([p for p in model.parameters() if p.requires_grad],1.0)
            opt.step(); nstep+=1
            if nstep<=5 or nstep%100==0: print(f"[D-SFT] ep{ep} step{nstep} loss={loss.item():.4f}", flush=True)
    model.save_pretrained(OUT_ADAPTER); tok.save_pretrained(OUT_ADAPTER)
    print(f"[D-SFT] SAVED {OUT_ADAPTER} steps={nstep}", flush=True); print("TRAINDONE", flush=True)
if __name__=="__main__": main()
