#!/usr/bin/env python3
"""OFDM + Saleh 功放：DPD / DPoD / 无处理 的频谱效率—功率效率 Pareto 前沿。

判定规则写在 ../claims.json，跑之前定死。
零模型是"错参预失真"：用错误的 Saleh 参数做同样的信号操作。
"""
import json
import numpy as np

N_SC, CP, N_SYM, SEEDS = 1024, 64, 200, [1, 2, 3, 4, 5]
M = 16                       # 16-QAM
IBO_DB = np.arange(0.0, 12.5, 0.5)
SNR_DB = 25.0                # 收端 AWGN
ETA_MAX = 0.785              # B 类功放理论最高效率
# Saleh 模型参数（归一化，文献常用取值）
SALEH = dict(aa=2.0, ba=1.0, ap=4.0, bp=9.0)
SALEH_WRONG = dict(aa=1.2, ba=0.3, ap=1.0, bp=2.0)   # 零模型用的错误参数


def saleh(x, p):
    r = np.abs(x)
    amp = p["aa"] * r / (1.0 + p["ba"] * r ** 2)
    pha = p["ap"] * r ** 2 / (1.0 + p["bp"] * r ** 2)
    return amp * np.exp(1j * (np.angle(x) + pha))


def saleh_inverse(y, p):
    """逐点求 Saleh 的逆：给定输出幅度，解出输入幅度。饱和点以上截断。"""
    a, b = p["aa"], p["ba"]
    r_out = np.abs(y)
    r_sat = a / (2.0 * np.sqrt(b))                 # AM/AM 的最大值
    r_out = np.minimum(r_out, r_sat * (1 - 1e-9))
    # b*r^2*r_out - a*r + r_out = 0  →  取小根（工作在饱和点以下那支）
    disc = np.maximum(a ** 2 - 4.0 * b * r_out ** 2, 0.0)
    r_in = (a - np.sqrt(disc)) / (2.0 * b * np.maximum(r_out, 1e-12))
    pha = p["ap"] * r_in ** 2 / (1.0 + p["bp"] * r_in ** 2)
    return r_in * np.exp(1j * (np.angle(y) - pha))


def qam(rng, n, m=M):
    k = int(np.sqrt(m)); lv = np.arange(-(k - 1), k, 2)
    s = rng.choice(lv, n) + 1j * rng.choice(lv, n)
    return s / np.sqrt((np.abs(lv[:, None] + 1j * lv[None, :]) ** 2).mean())


def ofdm(rng, n_sym):
    S = qam(rng, n_sym * N_SC).reshape(n_sym, N_SC)
    x = np.fft.ifft(S, axis=1) * np.sqrt(N_SC)
    return S, np.hstack([x[:, -CP:], x]).ravel()


def sndr_to_se(S_tx, S_rx):
    """用发送/接收星座算带内 SNDR，再转频谱效率。"""
    a = np.vdot(S_tx, S_rx) / np.vdot(S_tx, S_tx)     # 最小二乘增益校正
    err = S_rx - a * S_tx
    sndr = (np.abs(a) ** 2 * np.mean(np.abs(S_tx) ** 2)) / max(np.mean(np.abs(err) ** 2), 1e-15)
    return float(np.log2(1.0 + sndr)), float(sndr)


def run_point(rng, ibo_db, arm):
    S, x = ofdm(rng, N_SYM)
    p_in = np.mean(np.abs(x) ** 2)
    r_sat = SALEH["aa"] / (2 * np.sqrt(SALEH["ba"]))
    p_sat = r_sat ** 2
    scale = np.sqrt(p_sat / (p_in * 10 ** (ibo_db / 10.0)))
    xs = x * scale

    if arm == "dpd":
        xs = saleh_inverse(xs, SALEH)
    elif arm == "mismatched_dpd_null":
        xs = saleh_inverse(xs, SALEH_WRONG)
    y = saleh(xs, SALEH)

    p_out = np.mean(np.abs(y) ** 2)
    noise = np.sqrt(p_out / (2 * 10 ** (SNR_DB / 10.0))) * (rng.standard_normal(y.shape)
                                                            + 1j * rng.standard_normal(y.shape))
    z = y + noise
    if arm == "dpod":
        z = saleh_inverse(z, SALEH)

    z = z.reshape(N_SYM, N_SC + CP)[:, CP:]
    S_rx = np.fft.fft(z, axis=1) / np.sqrt(N_SC)
    se, sndr = sndr_to_se(S.ravel(), S_rx.ravel())
    eta = ETA_MAX * np.sqrt(min(p_out / p_sat, 1.0))
    return se, float(eta), sndr


def pareto(points):
    """非支配集：功率效率与频谱效率都要越大越好。"""
    out = []
    for i, (e1, s1) in enumerate(points):
        if not any((e2 >= e1 and s2 >= s1) and (e2 > e1 or s2 > s1) for e2, s2 in points):
            out.append((e1, s1))
    return sorted(out)


def main():
    arms = ["no_processing", "dpd", "dpod", "mismatched_dpd_null"]
    res = {a: {"ibo": [], "se": [], "eta": [], "sndr": []} for a in arms}
    for a in arms:
        for ibo in IBO_DB:
            ses, etas, sn = [], [], []
            for sd in SEEDS:
                rng = np.random.default_rng(sd * 1000 + int(ibo * 10))
                se, eta, s = run_point(rng, ibo, a)
                ses.append(se); etas.append(eta); sn.append(s)
            res[a]["ibo"].append(float(ibo)); res[a]["se"].append(float(np.mean(ses)))
            res[a]["eta"].append(float(np.mean(etas))); res[a]["sndr"].append(float(np.mean(sn)))
        print(f"[run] {a:<22} 完成 {len(IBO_DB)} 个回退点")

    fronts = {a: pareto(list(zip(res[a]["eta"], res[a]["se"]))) for a in arms}

    # H1：DPD 是否在每个功率效率点上都不劣于无处理
    def se_at(a, eta_q):
        e = np.array(res[a]["eta"]); s = np.array(res[a]["se"])
        return float(np.interp(eta_q, e[np.argsort(e)], s[np.argsort(e)]))
    grid = np.linspace(max(min(res["dpd"]["eta"]), min(res["no_processing"]["eta"])),
                       min(max(res["dpd"]["eta"]), max(res["no_processing"]["eta"])), 40)
    h1 = all(se_at("dpd", q) > se_at("no_processing", q) for q in grid)
    h2_mid = all(se_at("dpd", q) >= se_at("dpod", q) for q in grid) and \
             all(se_at("dpod", q) >= se_at("no_processing", q) for q in grid)
    # 强非线性区（IBO<=3dB）DPD 与 DPoD 的差距是否更大
    lo = [i for i, v in enumerate(IBO_DB) if v <= 3.0]; hi = [i for i, v in enumerate(IBO_DB) if v >= 9.0]
    gap_lo = float(np.mean([res["dpd"]["se"][i] - res["dpod"]["se"][i] for i in lo]))
    gap_hi = float(np.mean([res["dpd"]["se"][i] - res["dpod"]["se"][i] for i in hi]))
    h2 = h2_mid and gap_lo > gap_hi
    null_gain = float(np.mean([res["mismatched_dpd_null"]["se"][i] - res["no_processing"]["se"][i]
                               for i in range(len(IBO_DB))]))
    dpd_gain = float(np.mean([res["dpd"]["se"][i] - res["no_processing"]["se"][i]
                              for i in range(len(IBO_DB))]))
    print(f"\n[H1] DPD 全程支配无处理 → {'✅成立' if h1 else '❌被推翻'}")
    print(f"[H2] DPD >= DPoD >= 无处理 且强非线性区差距更大 → {'✅成立' if h2 else '❌被推翻'}")
    print(f"     强非线性区(IBO<=3dB) DPD-DPoD 差 {gap_lo:.3f} bit/s/Hz｜弱非线性区(>=9dB) {gap_hi:.3f}")
    print(f"[零模型] 错参预失真平均增益 {null_gain:+.3f}｜正确 DPD 平均增益 {dpd_gain:+.3f} bit/s/Hz")
    print(f"     → {'✅ 增益来自真的逆了非线性' if dpd_gain > 3*max(null_gain,1e-6) else '⚠️ 错参也能拿到相当增益，需查实现'}")
    json.dump({"seeds": SEEDS, "n_sym": N_SYM, "n_sc": N_SC, "snr_db": SNR_DB,
               "saleh": SALEH, "saleh_wrong": SALEH_WRONG,
               "results": res, "fronts": {k: [list(p) for p in v] for k, v in fronts.items()},
               "H1_passes": bool(h1), "H2_passes": bool(h2),
               "gap_low_ibo": gap_lo, "gap_high_ibo": gap_hi,
               "null_mean_gain": null_gain, "dpd_mean_gain": dpd_gain},
              open("results.json", "w"), ensure_ascii=False, indent=1)


if __name__ == "__main__":
    main()
