# -*- coding: utf-8 -*-
"""功能口径(激活感知)函数化测试:`regionify_act` = 对 W Xᵀ 分解而非对 W 分解。

验收要点:
  1. **同秩同字节**下,功能 EV 显著优于现有权重口径(`regionify(order=1)`);
  2. 它确实逼近**理论最优**(对 W L 的秩-r 截断,L 为激活二阶矩因子)——不是"只好一点";
  3. **权重 EV 会下降**:这是有意的取舍(不为权重 MSE 花字节),必须作为断言写下来防止有人"改回去";
  4. **二进制格式零改动**:结果经 `to_bytes`/`to_bytes_v3` 往返后功能 EV 不变;
  5. 边界:激活维度不符 / 全零激活 报错,K>1 分区仍优于权重口径。
"""
import numpy as np
import pytest

import nrfunc as N
from nrfunc import reconstruct


def _discriminating(seed=1, n=96, D=64, r_latent=32, r=8, m=512):
    """构造"选哪些方向"会决定成败的数据:W 有 r_latent 个方向但只允许留 r 个;
    激活强各向异性且带大均值(真实激活的典型形态)。"""
    rng = np.random.default_rng(seed)
    Qn, _ = np.linalg.qr(rng.normal(size=(n, r_latent)))
    Qd, _ = np.linalg.qr(rng.normal(size=(D, r_latent)))
    s = 0.9 ** np.arange(r_latent)
    W = (Qn * s) @ Qd.T
    X = rng.normal(size=(m, D)) * np.exp(np.linspace(0, 4, D)) * 0.1 + 3.0
    return W, X, r


def _theoretical_best_ev_func(W, X, r):
    """对 `W L` 做秩-r 截断(Eckart–Young),映射回权重空间后的功能 EV = 上界。"""
    L, _ = N.activation_metric(X)
    Y = W @ L
    U, s, Vt = np.linalg.svd(Y, full_matrices=False)
    rr = min(r, U.shape[1])
    Wh = (U[:, :rr] * s[:rr]) @ Vt[:rr] @ np.linalg.pinv(L)
    return N.ev_func(W, Wh, X)


def test_functional_objective_beats_weight_objective_at_equal_rank():
    W, X, r = _discriminating()
    res_w = N.regionify(W, signal='G', K=1, order=1, r=r, iters=30, seed=1)
    res_a = N.regionify_act(W, X, K=1, r=r, iters=30, seed=1)
    fw = N.ev_func(W, reconstruct(res_w), X)
    fa = N.ev_func(W, reconstruct(res_a), X)
    assert fa > fw + 0.10, f'功能口径应显著更好({fw:.4f} → {fa:.4f})'
    # 有意取舍:权重 EV 反而更低
    assert res_a['ev_weight'] < res_w['explained'], \
        '功能口径的权重 EV 应更低(字节花在"模型真正用的方向"上)'
    # 逼近理论上界(允许 3% 余量)
    best = _theoretical_best_ev_func(W, X, r)
    assert fa > best - 0.03, f'应逼近理论最优({fa:.4f} vs 上界 {best:.4f})'


def test_roundtrip_through_container_preserves_functional_ev():
    """**零格式改动**:order=1 结果经两个载体往返后功能 EV 不变。"""
    W, X, r = _discriminating()
    res = N.regionify_act(W, X, K=1, r=r, iters=30, seed=1)
    base = N.ev_func(W, reconstruct(res), X)
    for name, dec in (('v1v2', lambda b: N.from_bytes(b)),
                      ('v3', lambda b: N.from_bytes_v3(b))):
        blob = N.to_bytes(res, bits=8) if name == 'v1v2' else N.to_bytes_v3(res, bits=8)
        back = dec(blob)
        got = N.ev_func(W, reconstruct(back), X)
        assert abs(got - base) < 1e-3, f'{name} 往返后功能 EV 变化过大({base:.4f} → {got:.4f})'
        assert int(back['order']) == 1 and int(back['n']) == res['n'] and int(back['D']) == res['D']


def test_regions_still_beat_weight_objective():
    """K>1(按白化方向分区)也应优于权重口径。"""
    W, X, r = _discriminating(n=128, D=48, r_latent=40, r=4, m=384)
    res_w = N.regionify(W, signal='G', K=4, order=1, r=r, iters=30, seed=2)
    res_a = N.regionify_act(W, X, K=4, r=r, iters=30, seed=2)
    fw = N.ev_func(W, reconstruct(res_w), X)
    fa = N.ev_func(W, reconstruct(res_a), X)
    assert fa > fw, f'K=4 时功能口径也应更好({fw:.4f} → {fa:.4f})'
    assert int(res_a['K']) == 4 and res_a['assign'].shape == (W.shape[0],)
    assert res_a['components'].shape == (4, r, W.shape[1])


def test_ev_func_sanity():
    rng = np.random.default_rng(3)
    W = rng.normal(size=(16, 8))
    X = rng.normal(size=(64, 8))
    assert N.ev_func(W, W, X) == pytest.approx(1.0, abs=1e-9), '完全相同应得 1'
    assert N.ev_func(W, np.zeros_like(W), X) == pytest.approx(0.0, abs=1e-9), '全零近似应得 0'


def test_activation_metric_shapes_and_psd():
    rng = np.random.default_rng(4)
    X = rng.normal(size=(128, 32))
    L, scale = N.activation_metric(X)
    assert L.shape[0] == 32, '因子第一维应为输入维'
    assert L.shape[1] <= 128, '有效秩不超过样本数'
    S = L @ L.T
    S_ref = X.T @ X / X.shape[0]
    assert np.allclose(S, S_ref, atol=1e-8), 'L Lᵀ 应等于未中心化二阶矩'
    Lc, _ = N.activation_metric(X, center=True)
    Sc = Lc @ Lc.T
    Xc = X - X.mean(axis=0, keepdims=True)
    assert np.allclose(Sc, Xc.T @ Xc / (X.shape[0] - 1), atol=1e-8), 'center=True 应为常规协方差'
    assert scale > 0


def test_input_validation():
    rng = np.random.default_rng(5)
    W = rng.normal(size=(8, 4))
    with pytest.raises(ValueError):
        N.regionify_act(W, rng.normal(size=(16, 5)))          # 激活维度不符
    with pytest.raises(ValueError):
        N.regionify_act(W, np.zeros((16, 4)))                  # 全零激活


def test_regionify_auto_picks_objective_by_activation_availability():
    """**目标自适应**:有标定输入→功能口径;没有→权重口径且**如实标注**;不得静默降级。"""
    W, X, r = _discriminating()
    a = N.regionify_auto(W, X, K=1, r=r, iters=30, seed=1)
    assert a['objective_used'] == 'func_act', '有标定输入且 order=1 应走功能口径'
    w = N.regionify_auto(W, None, K=1, r=r, iters=30, seed=1)
    assert w['objective_used'] == 'weight', '无标定输入应回退权重口径'
    assert 'ev_func' not in w, '无标定时不应凭空报功能 EV'
    assert N.ev_func(W, reconstruct(a), X) > N.ev_func(W, reconstruct(w), X), \
        '同一输入下功能口径应优于权重口径'
    # order=0 + 标定输入 → 功能分区(F);注意 F 要的是 (n,T) 每单元响应,由本函数代为换算
    z = N.regionify_auto(W, X, K=4, order=0, iters=30, seed=1)
    assert z['objective_used'] == 'func_F' and int(z['order']) == 0
    # 不拍绝对值阈值:只要求**显著优于"预测行均值"这个平凡基线**(order=0 本就弱于 order=1)
    mean_only = np.tile(W.mean(axis=0), (W.shape[0], 1))
    base = N.ev_func(W, mean_only, X)
    assert z['ev_func'] > base + 0.3, \
        f'order=0 功能分区应显著优于行均值基线({z["ev_func"]:.3f} vs {base:.3f})'
    # 显式要功能口径但没给标定输入 → 报错,**不许**静默退回权重口径
    with pytest.raises(ValueError):
        N.regionify_auto(W, None, K=1, r=r, objective='func')


def test_calib_inputs_convention_is_enforced():
    """**口径防混用**:`calib_inputs` 是 (m, D) 输入样本;传成 (n, T) 每单元响应必须报错。"""
    W, X, r = _discriminating()
    with pytest.raises(ValueError) as e:
        N.regionify_act(W, (X @ W.T).T)      # 误传成每单元响应 (n, m)
    assert 'calib_inputs' in str(e.value)