"""
最小 DPO 玩具示例（阶段 3 配套）
--------------------------------
目标：用偏好对 (x, y_w, y_l) 训练一个极小策略，使
  log π(y_w|x) - log π_ref(y_w|x)  相对  log π(y_l|x) - log π_ref(y_l|x)  变大。

不依赖大模型：把「回复」编码成固定词表上的 one-hot 袋，仅演示损失形态。
运行：python dpo_toy.py
"""

from __future__ import annotations

import math
from dataclasses import dataclass

import torch
import torch.nn as nn
import torch.nn.functional as F


# ---------- 玩具词表与数据 ----------
# 每个「回复」是词 id 列表；这里刻意让 win 含 good，lose 含 bad。
VOCAB = {"<pad>": 0, "hello": 1, "thanks": 2, "good": 3, "bad": 4, "help": 5, "spam": 6}
INV = {v: k for k, v in VOCAB.items()}
V = len(VOCAB)
MAX_LEN = 4


def encode(tokens: list[str]) -> torch.Tensor:
    ids = [VOCAB[t] for t in tokens][:MAX_LEN]
    ids += [0] * (MAX_LEN - len(ids))
    return torch.tensor(ids, dtype=torch.long)


@dataclass
class Pref:
    prompt: str
    win: list[str]
    lose: list[str]


PREFS = [
    Pref("q1", ["hello", "good", "help"], ["hello", "bad", "spam"]),
    Pref("q1", ["thanks", "good"], ["spam", "bad"]),
    Pref("q2", ["help", "good"], ["bad", "spam"]),
    Pref("q2", ["hello", "thanks", "good"], ["hello", "bad"]),
    Pref("q3", ["good", "help"], ["spam", "bad", "spam"]),
    Pref("q3", ["thanks", "help", "good"], ["bad", "bad"]),
]


def bag_of_words(ids: torch.Tensor) -> torch.Tensor:
    """把 token id 序列变成词袋向量（演示用，非真实 LM）。"""
    x = torch.zeros(V)
    for i in ids.tolist():
        if i != 0:
            x[i] += 1.0
    return x


class TinyPolicy(nn.Module):
    """
    输出每个词的 logit；序列 logprob ≈ 词袋 · log_softmax(logits)
    仅用于展示 DPO 数学，不是可部署语言模型。
    """

    def __init__(self):
        super().__init__()
        self.emb = nn.Linear(V, 16)
        self.out = nn.Linear(16, V)

    def seq_logprob(self, bow: torch.Tensor) -> torch.Tensor:
        h = torch.tanh(self.emb(bow))
        logits = self.out(h)
        log_probs = F.log_softmax(logits, dim=-1)
        # 词袋计数加权的序列伪 logprob
        return (bow * log_probs).sum()


def dpo_loss(
    policy: TinyPolicy,
    ref: TinyPolicy,
    bow_w: torch.Tensor,
    bow_l: torch.Tensor,
    beta: float,
) -> torch.Tensor:
    """
    L = -log σ( β[log πθ(yw)/πref(yw) - log πθ(yl)/πref(yl)] )

    其中 log π(y)/πref(y) = logπ(y) - logπref(y)
    """
    with torch.no_grad():
        log_ref_w = ref.seq_logprob(bow_w)
        log_ref_l = ref.seq_logprob(bow_l)

    log_pi_w = policy.seq_logprob(bow_w)
    log_pi_l = policy.seq_logprob(bow_l)

    diff = beta * ((log_pi_w - log_ref_w) - (log_pi_l - log_ref_l))
    # -log sigmoid(diff)
    return -F.logsigmoid(diff)


def preference_margin(policy: TinyPolicy, ref: TinyPolicy, beta: float = 0.1) -> float:
    """平均隐式奖励差：越大说明越偏向 win。"""
    margins = []
    for p in PREFS:
        bw = bag_of_words(encode(p.win))
        bl = bag_of_words(encode(p.lose))
        with torch.no_grad():
            m = beta * (
                (policy.seq_logprob(bw) - ref.seq_logprob(bw))
                - (policy.seq_logprob(bl) - ref.seq_logprob(bl))
            )
        margins.append(m.item())
    return sum(margins) / len(margins)


def main():
    torch.manual_seed(0)
    beta = 0.1
    lr = 5e-2
    steps = 200

    # ref = 初始 SFT 快照（冻结）
    policy = TinyPolicy()
    ref = TinyPolicy()
    ref.load_state_dict(policy.state_dict())
    for p in ref.parameters():
        p.requires_grad_(False)

    opt = torch.optim.Adam(policy.parameters(), lr=lr)

    print("=== 最小 DPO 玩具 ===")
    print(f"beta={beta}, steps={steps}")
    print(f"初始 margin: {preference_margin(policy, ref, beta):.4f}  (≈0 合理)")

    for t in range(1, steps + 1):
        total = 0.0
        for pref in PREFS:
            bw = bag_of_words(encode(pref.win))
            bl = bag_of_words(encode(pref.lose))
            loss = dpo_loss(policy, ref, bw, bl, beta)
            opt.zero_grad()
            loss.backward()
            opt.step()
            total += loss.item()
        if t % 40 == 0 or t == 1:
            m = preference_margin(policy, ref, beta)
            print(f"step {t:3d} | loss {total/len(PREFS):.4f} | margin {m:.4f}")

    m_final = preference_margin(policy, ref, beta)
    print("---")
    print(f"最终 margin: {m_final:.4f}")
    print("解读：margin 上升 = 相对 ref，模型更抬高 win、压低 lose。")
    print("调参：beta 过大 → 不敢离开 ref；过小 → 过拟合噪声偏好。")
    if m_final <= 0.05:
        print("警告：margin 几乎没升，检查学习率或数据。")
    else:
        print("通过：玩具 DPO 方向正确。")


if __name__ == "__main__":
    main()
