"""函数形式前向(functional_forward):不还原稠密权重,直接按生成物前向。

演示两件事:
1. 正确性:函数前向 == 先 reconstruct() 再前向(浮点误差内一致)。
2. 提速:函数前向 vs 稠密前向,分「单样本延迟」与「批量吞吐」两种场景,
   并对照理论提速。

这是「为提速」设计的推理路径:共享的区域函数只算一次,K≪n 时省掉大量
重复乘;0 阶(质心)提速最猛,1 阶(低秩)精度更高但提速稍弱。
"""
import numpy as np

import _common  # noqa: F401  # 先导入:完成 src 路径引导
import nrfunc


def bench(rng, n, D, K, r, batches, iters_batch=50, iters_single=2000):
    W = rng.standard_normal((n, D))
    print(f'规模:n={n} 单元, D={D} 权重维, K={K} 区, r={r} 主成分\n')
    for order in (0, 1):
        kw = {'r': r} if order == 1 else {}
        res = nrfunc.regionify(W, signal='G', K=K, order=order, **kw)

        # 1) 正确性:函数前向 vs 还原后前向
        Xc = rng.standard_normal((16, D))
        y_ref = Xc @ nrfunc.reconstruct(res).T
        y_fwd = nrfunc.functional_forward(res, Xc)
        err = float(np.abs(y_fwd - y_ref).max())
        print(f'  order={order} 正确性:最大绝对误差 {err:.2e}'
              f'({"一致" if err < 1e-8 else "不一致!"})')

        # 2) 提速:稠密路径的 reconstruct 一次性,不计入每次推理
        recon = nrfunc.reconstruct(res)
        theory = n / K if order == 0 else n * D / (K * (1 + r) * D + n * r)
        for B in batches:
            X = rng.standard_normal((B, D))
            iters = iters_batch if B > 1 else iters_single
            dense_ns = nrfunc.cpu_time(lambda x: x @ recon.T, X,
                                       iters=iters, warmup=10)
            func_ns = nrfunc.cpu_time(lambda x: nrfunc.functional_forward(res, x), X,
                                      iters=iters, warmup=10)
            label = '单样本延迟' if B == 1 else f'batch={B} 吞吐'
            print(f'    {label:12s} 稠密 {dense_ns:10.0f} ns | '
                  f'函数 {func_ns:10.0f} ns | 提速 {dense_ns / func_ns:6.2f}x')

        # 3) 体积对比
        raw_b = nrfunc.size_bytes_raw(n, D, bits=32)
        reg_b = nrfunc.size_bytes_regionalized(res, bits=8)
        print(f'    体积 {raw_b / 1024:9.1f} KB -> {reg_b / 1024:7.1f} KB'
              f'(省 {(1 - reg_b / raw_b) * 100:.1f}%); 理论提速 {theory:.1f}x\n')


def main():
    rng = np.random.default_rng(0)
    print('=== 函数形式前向:正确性 + 提速 ===\n')
    bench(rng, n=1024, D=512, K=16, r=8, batches=(1, 1024))


if __name__ == '__main__':
    main()