"""Optional local base-model SFT. Syntax checked only; runtime unverified.
No model download or remote custom code. CPU float32, batch 1, no packing.
Use a locally prepared base checkpoint, e.g. SmolLM2-135M, not an Instruct model.
"""
import argparse
import json
from pathlib import Path
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer, set_seed


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument('--model', required=True, help='Local base checkpoint directory')
    ap.add_argument('--output', required=True)
    ap.add_argument('--variant', choices=['copy_only','balanced','full_loss'], default='balanced')
    ap.add_argument('--seed', type=int, default=67)
    args = ap.parse_args()
    out = Path(args.output)
    out.mkdir(parents=True, exist_ok=False)
    set_seed(args.seed)
    torch.set_num_threads(2)
    tok = AutoTokenizer.from_pretrained(args.model, local_files_only=True, use_fast=True)
    assert tok.is_fast and tok.eos_token_id is not None
    model = AutoModelForCausalLM.from_pretrained(args.model, local_files_only=True).float().cpu()
    model.config.use_cache = False
    data = json.loads((Path(__file__).parent/'data.json').read_text())

    def prompt(r, shifted=False):
        if shifted:
            return f"Instruction: {r['task']}\nData: {r['group']}\nResponse:"
        return f"Task: {r['task']}\nInput: {r['group']}\nAnswer:"

    def encode(r, full=False):
        prefix = prompt(r)
        answer = ' ' + ''.join(r['answer'][:-1])
        enc = tok(prefix + answer, add_special_tokens=False, return_offsets_mapping=True)
        ids = enc['input_ids'] + [tok.eos_token_id]
        # Mask boundary-crossing token to avoid supervising prompt characters.
        labels = [t if full or (s >= len(prefix) and e > s) else -100
                  for t,(s,e) in zip(enc['input_ids'],enc['offset_mapping'])] + [tok.eos_token_id]
        assert len(ids) <= 256, 'Refuse truncation; inspect data before training.'
        assert sum(x != -100 for x in labels[:-1]) > 0, 'Answer vanished at token boundary.'
        return {'input_ids':torch.tensor([ids]), 'labels':torch.tensor([labels])}

    train = ([r for r in data['train'] if r['task']=='COPY']*2 if args.variant=='copy_only' else data['train'])
    batches = [encode(r, args.variant=='full_loss') for r in train]
    assert not set(r['group'] for r in train) & set(data['split']['dev'] + data['split']['test'])
    print('first input/labels:',batches[0]['input_ids'].shape,batches[0]['labels'].tolist())
    print('supervised:',tok.decode([x for x in batches[0]['labels'][0].tolist() if x != -100]))

    @torch.no_grad()
    def nll(rows):
        model.eval()
        total, count = 0., 0
        for r in rows:
            b = encode(r)
            n = (b['labels'][:,1:] != -100).sum().item()
            total += model(**b).loss.item()*n
            count += n
        return total/count

    @torch.no_grad()
    def evaluate(rows, shifted=False):
        model.eval()
        records=[]
        for r in rows:
            inp = tok(prompt(r,shifted),add_special_tokens=False,return_tensors='pt')
            generated = model.generate(**inp,max_new_tokens=16,do_sample=False,
                                       pad_token_id=tok.eos_token_id,eos_token_id=tok.eos_token_id)
            text = tok.decode(generated[0,inp['input_ids'].shape[1]:],skip_special_tokens=True).strip()
            expected = ''.join(r['answer'][:-1])
            records.append({'group':r['group'],'task':r['task'],'output':text,'expected':expected,'correct':text==expected})
        return {'exact_match':sum(r['correct'] for r in records)/len(records),'predictions':records}

    baseline = evaluate(data['test'])
    opt = torch.optim.AdamW(model.parameters(),lr=2e-5)
    best, best_step, trace = float('inf'), None, []
    counts = 0
    for step in range(1,81):
        model.train()
        b = batches[(step-1)%len(batches)]
        counts += (b['labels'][:,1:] != -100).sum().item()
        loss = model(**b).loss  # AutoModelForCausalLM handles the single causal shift.
        assert torch.isfinite(loss)
        opt.zero_grad()
        loss.backward()
        torch.nn.utils.clip_grad_norm_(model.parameters(),1.)
        opt.step()
        if step in [16,32,64,80]:
            dev = nll(data['dev'])
            trace.append({'step':step,'dev_response_nll':dev,'supervised_tokens':counts})
            print(trace[-1],flush=True)
            if dev < best:
                best,best_step = dev,step
                model.save_pretrained(out/'selected')
                tok.save_pretrained(out/'selected')
    model = AutoModelForCausalLM.from_pretrained(out/'selected',local_files_only=True).float().cpu()
    report = {'seed':args.seed,'variant':args.variant,'model_path':str(Path(args.model).resolve()),
              'model_commit':getattr(model.config,'_commit_hash',None),'selected_step':best_step,
              'trace':trace,'baseline':baseline,'test':evaluate(data['test']),
              'template_shift':evaluate(data['test'],True),'test_response_nll':nll(data['test']),
              'torch':torch.__version__,'template':'explicit plain-text base-model template; no chat template'}
    (out/'results.json').write_text(json.dumps(report,indent=2))


if __name__ == '__main__':
    main()
