"""小号 System One：BERT 式编码器 + Laya 那套 marker 决策头。

序列格式照搬 laya.common.build_sequence：
  [CLS] <type> 指令 [SEP] [MASK]选项0 [MASK]选项1 ... [SEP] 状态 [SEP]
选项的得分 = 对应 [MASK] 位置的隐状态过一个线性层。所有 tokenizer 臂共用这套结构，
唯一的变量就是 tokenizer。
"""
import math, time
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F

DEV = torch.device("cuda")


class Block(nn.Module):
    def __init__(self, d, h, ff, drop=0.1):
        super().__init__()
        self.h, self.dh = h, d // h
        self.n1, self.n2 = nn.LayerNorm(d), nn.LayerNorm(d)
        self.qkv = nn.Linear(d, 3 * d)
        self.o = nn.Linear(d, d)
        self.f1, self.f2 = nn.Linear(d, ff), nn.Linear(ff, d)
        self.drop = nn.Dropout(drop)

    def forward(self, x, pad):
        B, T, D = x.shape
        y = self.n1(x)
        q, k, v = self.qkv(y).view(B, T, 3, self.h, self.dh).permute(2, 0, 3, 1, 4)
        a = F.scaled_dot_product_attention(q, k, v, attn_mask=pad)
        x = x + self.drop(self.o(a.transpose(1, 2).reshape(B, T, D)))
        y = self.n2(x)
        return x + self.drop(self.f2(F.gelu(self.f1(y))))


class Enc(nn.Module):
    def __init__(self, V, d=384, L=6, h=6, ff=1536, maxpos=512, drop=0.1):
        super().__init__()
        self.V, self.d = V, d
        self.tok = nn.Embedding(V, d)
        self.pos = nn.Embedding(maxpos, d)
        self.ln0 = nn.LayerNorm(d)
        self.drop = nn.Dropout(drop)
        self.blocks = nn.ModuleList([Block(d, h, ff, drop) for _ in range(L)])
        self.lnf = nn.LayerNorm(d)
        self.apply(self._init)

    @staticmethod
    def _init(m):
        if isinstance(m, nn.Linear):
            nn.init.normal_(m.weight, std=0.02)
            if m.bias is not None:
                nn.init.zeros_(m.bias)
        elif isinstance(m, nn.Embedding):
            nn.init.normal_(m.weight, std=0.02)

    def forward(self, ids, att):
        B, T = ids.shape
        x = self.ln0(self.tok(ids) + self.pos(torch.arange(T, device=ids.device))[None])
        x = self.drop(x)
        pad = att[:, None, None, :].bool()
        for b in self.blocks:
            x = b(x, pad)
        return self.lnf(x)

    def n_params(self):
        emb = self.tok.weight.numel() + self.pos.weight.numel()
        tot = sum(p.numel() for p in self.parameters())
        return tot, tot - emb


class MLM(nn.Module):
    def __init__(self, enc):
        super().__init__()
        self.enc = enc
        self.proj = nn.Linear(enc.d, enc.d)
        self.ln = nn.LayerNorm(enc.d)
        self.bias = nn.Parameter(torch.zeros(enc.V))

    def forward(self, ids, att):
        h = self.enc(ids, att)
        h = self.ln(F.gelu(self.proj(h)))
        return h @ self.enc.tok.weight.T + self.bias


class Decider(nn.Module):
    """marker 位置 -> 每个选项一个 logit。"""
    def __init__(self, enc, nq=3):
        super().__init__()
        self.enc = enc
        self.qemb = nn.Embedding(nq, enc.d)
        self.head = nn.Sequential(nn.Linear(enc.d, enc.d), nn.GELU(), nn.LayerNorm(enc.d), nn.Linear(enc.d, 1))

    def forward(self, ids, att, mpos, mmask, qtype):
        h = self.enc(ids, att)
        g = torch.gather(h, 1, mpos[:, :, None].expand(-1, -1, h.size(-1)))
        g = g + self.qemb(qtype)[:, None, :]
        lg = self.head(g).squeeze(-1)
        return lg.masked_fill(~mmask, -1e4)


# ---------------- MLM 预训练 ----------------
def pretrain(model, blocks, tok, steps, bs=256, lr=3e-4, warm=0.03, log=200, seed=0, mask_p=0.15):
    model.to(DEV).train()
    opt = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=0.01, betas=(0.9, 0.98), eps=1e-6)
    sched = torch.optim.lr_scheduler.LambdaLR(
        opt, lambda s: min(1.0, (s + 1) / max(1, int(steps * warm))) *
                       (0.5 * (1 + math.cos(math.pi * min(1.0, s / steps))) * 0.99 + 0.01))
    rng = np.random.default_rng(seed)
    N = blocks.shape[0]
    t0 = time.time(); run = 0.0
    special = torch.tensor([tok.pad_id, tok.cls_id, tok.sep_id, tok.mask_id], device=DEV)
    for s in range(steps):
        sel = rng.integers(0, N, bs)
        ids = torch.from_numpy(blocks[sel].astype(np.int64)).to(DEV)
        att = torch.ones_like(ids)
        lab = ids.clone()
        prob = torch.rand(ids.shape, device=DEV)
        m = (prob < mask_p) & ~torch.isin(ids, special)
        lab[~m] = -100
        r = torch.rand(ids.shape, device=DEV)
        ids = torch.where(m & (r < 0.8), tok.mask_id, ids)
        rnd = torch.randint(5, tok.size, ids.shape, device=DEV)
        ids = torch.where(m & (r >= 0.8) & (r < 0.9), rnd, ids)
        with torch.autocast("cuda", dtype=torch.bfloat16):
            out = model(ids, att)
        loss = F.cross_entropy(out.float().view(-1, tok.size), lab.view(-1), ignore_index=-100)
        loss.backward()
        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
        opt.step(); sched.step(); opt.zero_grad(set_to_none=True)
        run += loss.item()
        if (s + 1) % log == 0:
            print(f"  step {s+1:6d}/{steps}  loss {run/log:.4f}  ppl {math.exp(min(20,run/log)):8.1f}"
                  f"  {(s+1)*bs*blocks.shape[1]/1e6:.0f}M tok  {time.time()-t0:6.0f}s", flush=True)
            run = 0.0
    return model


# ---------------- 决策序列 ----------------
def build_seq(tok, state, ins, opts, qtype, max_len=192, head_max_len=64, opt_cap=48):
    head = tok.encode(f"{qtype} question: {ins}")
    oids = [[tok.mask_id] + tok.encode(o)[:opt_cap] for o in opts]
    budget = head_max_len - sum(len(o) for o in oids)
    if budget < 16:
        per = max(4, (head_max_len - 16) // max(1, len(oids)))
        oids = [o[:per] for o in oids]
        budget = head_max_len - sum(len(o) for o in oids)
    head = head[:max(8, budget)]
    ids = [tok.cls_id] + head + [tok.sep_id]
    mk = []
    for o in oids:
        mk.append(len(ids)); ids.extend(o)
    ids.append(tok.sep_id)
    room = max(0, max_len - len(ids) - 1)
    ids = ids + tok.encode(state)[:room] + [tok.sep_id]
    return ids[:max_len], [m for m in mk if m < max_len]


def collate(items, pad_id):
    n = len(items); L = max(len(i["ids"]) for i in items); K = max(len(i["mk"]) for i in items)
    ids = torch.full((n, L), pad_id, dtype=torch.long)
    att = torch.zeros((n, L), dtype=torch.long)
    mp = torch.zeros((n, K), dtype=torch.long); mm = torch.zeros((n, K), dtype=torch.bool)
    for i, it in enumerate(items):
        ids[i, :len(it["ids"])] = torch.tensor(it["ids"]); att[i, :len(it["ids"])] = 1
        k = len(it["mk"]); mp[i, :k] = torch.tensor(it["mk"]); mm[i, :k] = True
    return dict(ids=ids, att=att, mp=mp, mm=mm, qt=torch.tensor([it["qt"] for it in items]))


def finetune(model, items, y, pad_id, epochs=3, bs=64, lr_enc=5e-5, lr_head=3e-4, seed=0,
             eval_fn=None, soft=False):
    model.to(DEV).train()
    ep_ = [p for n, p in model.named_parameters() if n.startswith("enc.")]
    hp_ = [p for n, p in model.named_parameters() if not n.startswith("enc.")]
    opt = torch.optim.AdamW([{"params": ep_, "lr": lr_enc}, {"params": hp_, "lr": lr_head}], weight_decay=0.01)
    steps = max(1, len(items) // bs) * epochs
    sch = torch.optim.lr_scheduler.OneCycleLR(opt, max_lr=[lr_enc, lr_head], total_steps=steps, pct_start=0.1)
    idx = np.arange(len(items)); done = 0
    for ep in range(epochs):
        rng = np.random.default_rng(seed + ep); rng.shuffle(idx)
        idx = np.concatenate([b[np.argsort([len(items[j]["ids"]) for j in b])]
                              for b in np.array_split(idx, max(1, len(idx) // 2048))])
        bl = [idx[i:i + bs] for i in range(0, len(idx), bs)]
        rng.shuffle(bl)
        model.train()
        for sel in bl:
            if done >= steps: break
            b = collate([items[j] for j in sel], pad_id)
            with torch.autocast("cuda", dtype=torch.bfloat16):
                lg = model(b["ids"].to(DEV), b["att"].to(DEV), b["mp"].to(DEV), b["mm"].to(DEV), b["qt"].to(DEV))
            lg = lg.float()
            if soft:
                t = torch.tensor(np.stack([[1 - y[j], y[j]] for j in sel]), dtype=torch.float32, device=DEV)
                loss = -(t * F.log_softmax(lg, -1)).sum(-1).mean()
            else:
                loss = F.cross_entropy(lg, torch.tensor([y[j] for j in sel], device=DEV))
            loss.backward()
            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
            opt.step(); sch.step(); opt.zero_grad(set_to_none=True); done += 1
        if eval_fn:
            eval_fn(ep, model)
    return model


@torch.no_grad()
def predict(model, items, pad_id, bs=256):
    model.eval()
    order = np.argsort([len(i["ids"]) for i in items])
    K = max(len(i["mk"]) for i in items)
    out = np.full((len(items), K), -1e4, dtype=np.float64)
    for i in range(0, len(order), bs):
        sel = order[i:i + bs]
        b = collate([items[j] for j in sel], pad_id)
        with torch.autocast("cuda", dtype=torch.bfloat16):
            lg = model(b["ids"].to(DEV), b["att"].to(DEV), b["mp"].to(DEV), b["mm"].to(DEV), b["qt"].to(DEV))
        v = lg.float().cpu().numpy()
        out[sel[:, None], np.arange(v.shape[1])[None, :]] = v
    return out
