"""函数形式前向(functional_forward):不还原稠密权重,直接按生成物前向。
演示两件事:
1. 正确性:函数前向 == 先 reconstruct() 再前向(浮点误差内一致)。
2. 提速:函数前向 vs 稠密前向,分「单样本延迟」与「批量吞吐」两种场景,
并对照理论提速。
这是「为提速」设计的推理路径:共享的区域函数只算一次,K≪n 时省掉大量
重复乘;0 阶(质心)提速最猛,1 阶(低秩)精度更高但提速稍弱。
"""
import numpy as np
import _common
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)
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 "不一致!"})')
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')
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()