"""端到端演示:成熟模型权重 → 函数化 → 生成物(价值对比 / 量化打包 / 分片存储)。

运行:python examples/functionalize_demo.py
"""
import numpy as np

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


def main():
    # 1) 上游:你已训练好的成熟模型,抽出一层权重(这里用随机矩阵占位)
    rng = np.random.default_rng(0)
    W = rng.standard_normal((256, 128)).astype(np.float32)

    # 2) 函数化:256 个输出单元 → 16 个区域函数(1 阶低秩,每区 8 个主成分)
    res = nrfunc.regionify(W, signal='G', K=16, order=1, r=8)
    print(f"[函数化] 区域数 K={res['K']},1 阶解释方差={res['explained']:.3f}")

    # 2.5) 生成物 → 重建权重:还原近似权重,可当模型权重用(下游再接入推理)
    recon = nrfunc.reconstruct(res)
    recon_ev = nrfunc.explained_variance(W, res['assign'], res['K'], lambda: recon)
    print(f"[重建权重] shape={recon.shape},还原保真度={recon_ev:.3f}(=函数化口径,一致)")

    # 3) 生成物 · 价值对比(函数化前 vs 后:体积 / 内存 / CPU / GPU)
    tbl = nrfunc.compare_value(W, res, bits_before=32, bits_after=8)
    print("[价值对比] 函数化前 vs 后")
    print(tbl['_table'])

    # 4) 生成物 · 量化打包(把区域函数参数量化成可部署字节)
    packed, scale, n_out, n_in = nrfunc.quantize_weights(res['means'], bits=8)
    print(f"[量化打包] means 打包后 {len(packed)} 字节")

    # 5) 生成物 · 分散存储到 4 个节点,并行读取
    store = nrfunc.build_store(res, n_shards=4)
    print(f"[分片存储] 分布={store.placement()},负载={store.load_balance().tolist()}")
    params = store.region_read_many([0, 1, 2, 3])
    print(f"[并行读取] 读回 {len(params)} 个区域函数")


if __name__ == '__main__':
    main()