//! 读 Python 训练侧产出的两类二进制(禁用 JSON 存生成物),做「成熟模型稠密部署」
//! vs 「函数化生成物函数部署」的实质对比,输出比较日志 comparison.log:
//!   路线 A:model.bin      —— 全部层稠密权重,逐层稠密前向
//!   路线 B:functional.bin —— fc2 层用区域函数(函数前向),其余层稠密前向
//!
//! 二进制格式(小端,跨语言自描述):
//!   model.bin:n_layers(4B) + 每层 [n_out(4B) n_in(4B) f32[n_out*n_in]],层序 fc1/fc2/fc3/head
//!   functional.bin:fc1/fc3/head 三层的 [n_out n_in f32] + func_blob_len(8B) + NRFN 块
//!   NRFN 块(bits=32):magic'NRFN'(4B) version(1B) order(1B) bits(1B) K(4B) n(4B) D(4B) r(4B)
//!     + means(K×D f32) [+ components(K×r×D f32) + coeffs(n×r f32) if order=1] + assign(n×i32)
//!   模型:784 → 4096 → 2048 → 1024 → 网格头(60)。多对象检测(YOLO 式 2×2 网格)。
//!
//! 对比维度(同为 Rust 实现、同机同批):
//!   - 体积:两类二进制字节数 + 权重参数量(内存占用)
//!   - 精度:同一批输入的网格输出(objectness/分类/bbox)对比
//!   - CPU  :单样本延迟 + 批量吞吐

use std::collections::HashMap;
use std::fs;
use std::time::Instant;

mod gpu;

type Mat = Vec<Vec<f64>>;

/// int8 量化权重矩阵:值(有符号 int8,行主序拍平)+ 单个全局 scale(per-array)。
/// 供 int8 kernel 前向用(不反量化回 f64,直接 int8 域累加)。
struct QMat {
    q: Vec<i8>,        // n_out * d_in 个有符号 int8(行主序)
    scale: f64,        // 单全局 scale(per-array 对称量化)
    n_out: usize,
    d_in: usize,
}

/// int8 点积:Σ x[j] * w[j],AVX2 向量化。无 AVX2 时回退标量。
/// 把 int8 对称扩展成 int16,用 madd_epi16 做 int16×int16→int32 累加(避免 maddubs 的
/// unsigned/signed 符号陷阱,保证与标量逐位一致)。
#[inline]
fn dot_i8(x: &[i8], w: &[i8]) -> i32 {
    #[cfg(target_arch = "x86_64")]
    {
        if is_x86_feature_detected!("avx2") {
            return unsafe { dot_i8_avx2(x, w) };
        }
    }
    let mut acc: i32 = 0;
    for j in 0..x.len() {
        acc += x[j] as i32 * w[j] as i32;
    }
    acc
}

#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn dot_i8_avx2(x: &[i8], w: &[i8]) -> i32 {
    use std::arch::x86_64::*;
    let n = x.len();
    let mut acc0 = _mm256_setzero_si256();
    let mut acc1 = _mm256_setzero_si256();
    let mut i = 0usize;
    // 每次处理 16 个 int8:扩展到 16 个 int16,madd_epi16 两两乘加 → 8 个 int32;两路解 ILP
    while i + 32 <= n {
        let xa = _mm256_cvtepi8_epi16(_mm_loadu_si128(x.as_ptr().add(i) as *const __m128i));
        let wa = _mm256_cvtepi8_epi16(_mm_loadu_si128(w.as_ptr().add(i) as *const __m128i));
        let xb = _mm256_cvtepi8_epi16(_mm_loadu_si128(x.as_ptr().add(i + 16) as *const __m128i));
        let wb = _mm256_cvtepi8_epi16(_mm_loadu_si128(w.as_ptr().add(i + 16) as *const __m128i));
        acc0 = _mm256_add_epi32(acc0, _mm256_madd_epi16(xa, wa));
        acc1 = _mm256_add_epi32(acc1, _mm256_madd_epi16(xb, wb));
        i += 32;
    }
    while i + 16 <= n {
        let xa = _mm256_cvtepi8_epi16(_mm_loadu_si128(x.as_ptr().add(i) as *const __m128i));
        let wa = _mm256_cvtepi8_epi16(_mm_loadu_si128(w.as_ptr().add(i) as *const __m128i));
        acc0 = _mm256_add_epi32(acc0, _mm256_madd_epi16(xa, wa));
        i += 16;
    }
    acc0 = _mm256_add_epi32(acc0, acc1);
    // 尾部标量
    let mut tail: i32 = 0;
    while i < n {
        tail += x[i] as i32 * w[i] as i32;
        i += 1;
    }
    // 水平求和 8 个 i32
    let lo = _mm256_castsi256_si128(acc0);
    let hi = _mm256_extracti128_si256::<1>(acc0);
    let s = _mm_add_epi32(lo, hi);
    let s = _mm_hadd_epi32(s, s);
    let s = _mm_hadd_epi32(s, s);
    tail + _mm_cvtsi128_si32(s)
}

impl QMat {
    /// int8 前向:y = x @ W^T。x 亦量化为 int8(per-array),int32 累加,
    /// 最后反量化 y = acc * x_scale * w_scale。
    fn forward_int8(&self, x: &[f64], b: usize, y: &mut [f64]) {
        // 激活 per-array 量化
        let mut ax = 0.0f64;
        for &v in x {
            let a = v.abs();
            if a > ax { ax = a; }
        }
        let x_scale = if ax > 1e-12 { ax / 127.0 } else { 1.0 };
        let x_q: Vec<i8> = x.iter().map(|&v| (v / x_scale).round() as i8).collect();
        let n = self.n_out;
        let d = self.d_in;
        let wq = &self.q;
        let total_scale = x_scale * self.scale;
        for bi in 0..b {
            let xr = &x_q[bi * d..(bi + 1) * d];
            for i in 0..n {
                let wrow = &wq[i * d..(i + 1) * d];
                let acc = dot_i8(xr, wrow);
                y[bi * n + i] = acc as f64 * total_scale;
            }
        }
    }
}

// ---------- 二进制解析(零依赖,小端) ----------

/// 读整文件字节。
fn read_bytes(path: &str) -> Vec<u8> {
    fs::read(path).unwrap_or_else(|e| panic!("读 {path} 失败:{e}"))
}

fn u32_at(b: &[u8], off: usize) -> u32 {
    u32::from_le_bytes(b[off..off + 4].try_into().unwrap())
}
fn u64_at(b: &[u8], off: usize) -> u64 {
    u64::from_le_bytes(b[off..off + 8].try_into().unwrap())
}
fn f32_at(b: &[u8], off: usize) -> f32 {
    f32::from_le_bytes(b[off..off + 4].try_into().unwrap())
}

/// 解析「n_layers + 每层 [n_out n_in f32[]]」稠密权重块,返回 Mat 列表。
/// 支持量化:每层 n_out(4B) n_in(4B) bits(1B) + 权重;bits=32 直存 f32,
/// bits=8/4 为 scale(1×f32) + 打包 int8/int4 字节(per-array 对称量化,反量化回 f64)。
fn parse_dense_layers(b: &[u8], mut off: usize) -> (Vec<Mat>, usize) {
    let n_layers = u32_at(b, off) as usize;
    off += 4;
    let mut layers = Vec::with_capacity(n_layers);
    for _ in 0..n_layers {
        let n_out = u32_at(b, off) as usize;
        let n_in = u32_at(b, off + 4) as usize;
        off += 8;
        let bits = b[off];
        off += 1;
        let mut w = Vec::with_capacity(n_out);
        match bits {
            32 => {
                for i in 0..n_out {
                    let mut row = Vec::with_capacity(n_in);
                    for j in 0..n_in {
                        row.push(f32_at(b, off + (i * n_in + j) * 4) as f64);
                    }
                    w.push(row);
                }
                off += n_out * n_in * 4;
            }
            8 | 4 => {
                // packing 层:打包字节在前,单个全局 scale 在段尾(与 nrfunc.io.to_bytes 同口径)
                let packed_bytes = (n_out * n_in * bits as usize + 7) / 8;
                let scale = f32_at(b, off + packed_bytes) as f64;
                let qmax = (2f64.powi(bits as i32 - 1) - 1.0) as i64;
                for i in 0..n_out {
                    let mut row = Vec::with_capacity(n_in);
                    for j in 0..n_in {
                        let idx = i * n_in + j;
                        let v = if bits == 8 {
                            b[off + idx] as i64
                        } else {
                            // int4 两元素拼 1 字节:偶下标取低 nibble,奇下标取高 nibble
                            let byte = b[off + idx / 2];
                            if idx % 2 == 0 { (byte & 0x0F) as i64 } else { ((byte >> 4) & 0x0F) as i64 }
                        };
                        let q = v - qmax; // 去偏移,回到有符号
                        row.push(q as f64 * scale);
                    }
                    w.push(row);
                }
                off += packed_bytes + 4; // 打包字节 + scale
            }
            other => panic!("不支持的量化位宽 {other}(仅 32/8/4)"),
        }
        layers.push(w);
    }
    (layers, off)
}

/// 解析稠密权重块为 int8 量化矩阵列表(保留 int8 形式,供 int8 kernel 前向)。
/// 每层格式同 parse_dense_layers;bits=8/4 直接存 int8 值 + scale,bits=32 现场量化到 int8。
fn parse_dense_layers_q(b: &[u8], mut off: usize) -> (Vec<QMat>, usize) {
    let n_layers = u32_at(b, off) as usize;
    off += 4;
    let mut layers = Vec::with_capacity(n_layers);
    for _ in 0..n_layers {
        let n_out = u32_at(b, off) as usize;
        let n_in = u32_at(b, off + 4) as usize;
        off += 8;
        let bits = b[off];
        off += 1;
        match bits {
            32 => {
                // 现场量化到 int8(per-array 单 scale)
                let mut vals = Vec::with_capacity(n_out * n_in);
                let mut ax = 0.0f64;
                for i in 0..n_out * n_in {
                    let v = f32_at(b, off + i * 4) as f64;
                    vals.push(v);
                    if v.abs() > ax { ax = v.abs(); }
                }
                let scale = if ax > 1e-12 { ax / 127.0 } else { 1.0 };
                let q: Vec<i8> = vals.iter().map(|&v| (v / scale).round() as i8).collect();
                layers.push(QMat { q, scale, n_out, d_in: n_in });
                off += n_out * n_in * 4;
            }
            8 | 4 => {
                let packed_bytes = (n_out * n_in * bits as usize + 7) / 8;
                let scale = f32_at(b, off + packed_bytes) as f64;
                let qmax = (2f64.powi(bits as i32 - 1) - 1.0) as i64;
                let mut q = Vec::with_capacity(n_out * n_in);
                for i in 0..n_out * n_in {
                    let v = if bits == 8 {
                        b[off + i] as i64
                    } else {
                        let byte = b[off + i / 2];
                        if i % 2 == 0 { (byte & 0x0F) as i64 } else { ((byte >> 4) & 0x0F) as i64 }
                    };
                    q.push((v - qmax) as i8);
                }
                layers.push(QMat { q, scale, n_out, d_in: n_in });
                off += packed_bytes + 4;
            }
            other => panic!("不支持的量化位宽 {other}(仅 32/8/4)"),
        }
    }
    (layers, off)
}

/// 解析 NRFN 函数块,返回 (order, K, r, D, n, assign, centroids, means, comp2d, coeffs)。
/// 支持 v1(23B头)/v2(25B头) 与 bits=32/8/4(量化段为 scale(1×f32)+打包字节)。
#[allow(clippy::type_complexity)]
fn parse_nrfn(b: &[u8], mut off: usize)
    -> (i64, usize, usize, usize, usize, Vec<usize>,
        Option<Mat>, Option<Mat>, Option<Mat>, Option<Mat>) {
    let magic = &b[off..off + 4];
    assert_eq!(magic, b"NRFN", "NRFN magic 不匹配");
    let version = b[off + 4];
    let order = b[off + 5] as i64;
    let bits = b[off + 6];
    let k = u32_at(b, off + 7) as usize;
    let n = u32_at(b, off + 11) as usize;
    let d = u32_at(b, off + 15) as usize;
    let r = u32_at(b, off + 19) as usize;
    // v1=23B 头,v2=25B 头(追加 signal_index + op_version,本示例不读它们)
    let header = match version {
        1 => 23usize,
        2 => 25usize,
        v => panic!("NRFN 版本 {v} 不支持"),
    };
    off += header;

    let qmax = (2f64.powi(bits as i32 - 1) - 1.0) as i64;

    // 读一个「参数段」:bits=32 直存 f32[count];bits=8/4 为 打包字节 + scale(1×f32 在段尾)。
    // 返回 (值列表, 新 off)。
    let read_seg = |b: &[u8], off: usize, count: usize| -> (Vec<f64>, usize) {
        if bits == 32 {
            let mut v = Vec::with_capacity(count);
            for i in 0..count {
                v.push(f32_at(b, off + i * 4) as f64);
            }
            (v, off + count * 4)
        } else {
            let nbytes = (count * bits as usize + 7) / 8;
            let scale = f32_at(b, off + nbytes) as f64; // scale 在段尾
            let mut v = Vec::with_capacity(count);
            for i in 0..count {
                let q = if bits == 8 {
                    b[off + i] as i64
                } else {
                    let byte = b[off + i / 2];
                    if i % 2 == 0 { (byte & 0x0F) as i64 } else { ((byte >> 4) & 0x0F) as i64 }
                };
                v.push((q - qmax) as f64 * scale);
            }
            (v, off + nbytes + 4)
        }
    };

    // means (K×D)
    let (means_flat, mut off2) = read_seg(b, off, k * d);
    let mut means = Mat::with_capacity(k);
    for i in 0..k {
        means.push(means_flat[i * d..(i + 1) * d].to_vec());
    }

    let (centroids, comp2d, coeffs) = if order == 0 {
        (Some(means.clone()), None, None)
    } else {
        // components (K×r×D)
        let (comp_flat, off3) = read_seg(b, off2, k * r * d);
        let mut comp2d = Mat::with_capacity(k * r);
        for c in 0..k * r {
            comp2d.push(comp_flat[c * d..(c + 1) * d].to_vec());
        }
        // coeffs (n×r)
        let (co_flat, off4) = read_seg(b, off3, n * r);
        let mut coeffs = Mat::with_capacity(n);
        for i in 0..n {
            coeffs.push(co_flat[i * r..(i + 1) * r].to_vec());
        }
        off2 = off4;
        (None, Some(comp2d), Some(coeffs))
    };

    // assign (n × i32)
    let mut assign = Vec::with_capacity(n);
    for i in 0..n {
        assign.push(i32::from_le_bytes(b[off2 + i * 4..off2 + i * 4 + 4].try_into().unwrap()) as usize);
    }
    (order, k, r, d, n, assign, centroids, Some(means), comp2d, coeffs)
}

// ---------- 矩阵运算基元 ----------
#[inline(always)]
fn dot(x: &[f64], w: &[f64]) -> f64 {
    let n = x.len();
    let (mut a0, mut a1, mut a2, mut a3) = (0.0f64, 0.0f64, 0.0f64, 0.0f64);
    let (mut b0, mut b1, mut b2, mut b3) = (0.0f64, 0.0f64, 0.0f64, 0.0f64);
    let mut i = 0usize;
    while i + 8 <= n {
        a0 += x[i] * w[i];
        a1 += x[i + 1] * w[i + 1];
        a2 += x[i + 2] * w[i + 2];
        a3 += x[i + 3] * w[i + 3];
        b0 += x[i + 4] * w[i + 4];
        b1 += x[i + 5] * w[i + 5];
        b2 += x[i + 6] * w[i + 6];
        b3 += x[i + 7] * w[i + 7];
        i += 8;
    }
    while i + 4 <= n {
        a0 += x[i] * w[i];
        a1 += x[i + 1] * w[i + 1];
        a2 += x[i + 2] * w[i + 2];
        a3 += x[i + 3] * w[i + 3];
        i += 4;
    }
    let mut acc = (a0 + a1) + (a2 + a3) + (b0 + b1) + (b2 + b3);
    while i < n {
        acc += x[i] * w[i];
        i += 1;
    }
    acc
}

/// 稠密前向:y = x @ W^T。x 行主序 (b, d),W 行主序 (n, d),y 行主序 (b, n)。
fn dense_forward(x: &[f64], w: &Mat, b: usize, d: usize, n: usize, y: &mut [f64]) {
    for bi in 0..b {
        let xr = &x[bi * d..(bi + 1) * d];
        let mut i = 0usize;
        while i + 4 <= n {
            let (mut a0, mut a1, mut a2, mut a3) = (0.0f64, 0.0f64, 0.0f64, 0.0f64);
            let w0 = &w[i];
            let w1 = &w[i + 1];
            let w2 = &w[i + 2];
            let w3 = &w[i + 3];
            for dj in 0..d {
                let xv = xr[dj];
                a0 += xv * w0[dj];
                a1 += xv * w1[dj];
                a2 += xv * w2[dj];
                a3 += xv * w3[dj];
            }
            y[bi * n + i] = a0;
            y[bi * n + i + 1] = a1;
            y[bi * n + i + 2] = a2;
            y[bi * n + i + 3] = a3;
            i += 4;
        }
        while i < n {
            y[bi * n + i] = dot(xr, &w[i]);
            i += 1;
        }
    }
}

/// 0 阶函数前向:y[bi,i] = x[bi,:]·centroids[assign[i],:]
fn func0_forward(x: &[f64], centroids: &Mat, assign: &[usize],
                 b: usize, d: usize, n: usize, y: &mut [f64]) {
    let k = centroids.len();
    let mut z = vec![0.0f64; k];
    for bi in 0..b {
        let xr = &x[bi * d..(bi + 1) * d];
        for kk in 0..k {
            z[kk] = dot(xr, &centroids[kk]);
        }
        for i in 0..n {
            y[bi * n + i] = z[assign[i]];
        }
    }
}

/// 1 阶函数前向(共享均值 + 每单元低秩修正)
fn func1_forward(x: &[f64], means: &Mat, comp2d: &Mat, coeffs: &Mat, assign: &[usize],
                 b: usize, d: usize, n: usize, r: usize, y: &mut [f64]) {
    let k = means.len();
    let kr = k * r;
    let mut g = vec![0.0f64; k];
    let mut h = vec![0.0f64; kr];
    for bi in 0..b {
        let xr = &x[bi * d..(bi + 1) * d];
        for kk in 0..k {
            g[kk] = dot(xr, &means[kk]);
        }
        for c in 0..kr {
            h[c] = dot(xr, &comp2d[c]);
        }
        for i in 0..n {
            let a = assign[i];
            let base = a * r;
            let mut corr = 0.0f64;
            for j in 0..r {
                corr += h[base + j] * coeffs[i][j];
            }
            y[bi * n + i] = g[a] + corr;
        }
    }
}

/// 1 阶函数前向(int8):共享 GEMM(means/comp → g/h)走 int8 域累加(AVX2 dot_i8),
/// gather 的 coeffs 修正保持 f64(避免对动态 h 的二次量化误差叠加)。
/// 与 b1/`forward_int8_net` 同 per-array 对称量化口径。
fn func1_forward_int8(x: &[f64], means: &Mat, comp2d: &Mat, coeffs: &Mat, assign: &[usize],
                      b: usize, d: usize, n: usize, r: usize, y: &mut [f64]) {
    let k = means.len();
    let kr = k * r;
    // 量化 means / comp / 激活 x(各 per-array scale)
    let mut ax = 0.0f64;
    for &v in x { if v.abs() > ax { ax = v.abs(); } }
    let x_scale = if ax > 1e-12 { ax / 127.0 } else { 1.0 };
    let xq: Vec<i8> = x.iter().map(|&v| (v / x_scale).round() as i8).collect();

    let mut am = 0.0f64;
    for row in means { for &v in row { if v.abs() > am { am = v.abs(); } } }
    let m_scale = if am > 1e-12 { am / 127.0 } else { 1.0 };
    let means_q: Vec<Vec<i8>> = means.iter()
        .map(|row| row.iter().map(|&v| (v / m_scale).round() as i8).collect()).collect();

    let mut ac = 0.0f64;
    for row in comp2d { for &v in row { if v.abs() > ac { ac = v.abs(); } } }
    let c_scale = if ac > 1e-12 { ac / 127.0 } else { 1.0 };
    let comp_q: Vec<Vec<i8>> = comp2d.iter()
        .map(|row| row.iter().map(|&v| (v / c_scale).round() as i8).collect()).collect();

    let mut g = vec![0.0f64; k];
    let mut h = vec![0.0f64; kr];
    for bi in 0..b {
        let xr = &xq[bi * d..(bi + 1) * d];
        for kk in 0..k {
            g[kk] = dot_i8(xr, &means_q[kk]) as f64 * x_scale * m_scale;
        }
        for c in 0..kr {
            h[c] = dot_i8(xr, &comp_q[c]) as f64 * x_scale * c_scale;
        }
        for i in 0..n {
            let a = assign[i];
            let base = a * r;
            let mut corr = 0.0f64;
            for j in 0..r {
                corr += h[base + j] * coeffs[i][j];
            }
            y[bi * n + i] = g[a] + corr;
        }
    }
}

fn relu(x: &mut [f64]) {
    for v in x.iter_mut() {
        if *v < 0.0 {
            *v = 0.0;
        }
    }
}

/// 拼接 bias 尾列,输出 (b, d+1)
fn with_bias(x: &[f64], b: usize, d: usize) -> Vec<f64> {
    let mut out = vec![0.0f64; b * (d + 1)];
    for bi in 0..b {
        for j in 0..d {
            out[bi * (d + 1) + j] = x[bi * d + j];
        }
        out[bi * (d + 1) + d] = 1.0;
    }
    out
}

/// int8 全网络前向:fc1→relu→fc2→relu→fc3→relu→head,全链路 int8 累加。
/// layers 依序 fc1/fc2/fc3/head,每层含 bias 尾列(d = 输入维 + 1)。
/// 数值等价于「反量化后 f64 前向」的近似(int8 量化误差内)。
fn forward_int8_net(layers: &[QMat], x_in: &[f64], b: usize, y: &mut [f64]) {
    let d_in = layers[0].d_in - 1; // 首层输入维(不含 bias)
    let mut cur = with_bias(x_in, b, d_in); // (b, d+1)
    let n_layers = layers.len();
    for (li, layer) in layers.iter().enumerate() {
        let n = layer.n_out;
        let mut z = vec![0.0f64; b * n];
        layer.forward_int8(&cur, b, &mut z);
        if li < n_layers - 1 {
            relu(&mut z);
            cur = with_bias(&z, b, n);
        } else {
            // 最后一层 head 无 relu,直接写入输出
            y.copy_from_slice(&z);
        }
    }
}

/// 函数化 int8 全网络前向:fc1(int8 稠密)→relu→fc2(函数化 int8)→relu→fc3(int8 稠密)→relu→head。
/// 与 `forward_int8_net` 同口径,仅 fc2 层从稠密 int8 换成函数化 int8(共享 GEMM 走 dot_i8)。
fn forward_func_int8_net(q_fc1: &QMat, q_fc3: &QMat, q_head: &QMat,
                         means: &Mat, comp2d: &Mat, coeffs: &Mat, assign: &[usize],
                         r: usize, x_in: &[f64], b: usize, y: &mut [f64]) {
    let d_in = q_fc1.d_in - 1;
    // fc1(int8 稠密)+ relu
    let cur = with_bias(x_in, b, d_in);
    let n1 = q_fc1.n_out;
    let mut z1 = vec![0.0f64; b * n1];
    q_fc1.forward_int8(&cur, b, &mut z1);
    relu(&mut z1);
    // fc2(函数化 int8)+ relu
    let d2 = n1 + 1;
    let x2 = with_bias(&z1, b, n1);
    let n2 = assign.len();
    let mut z2 = vec![0.0f64; b * n2];
    func1_forward_int8(&x2, means, comp2d, coeffs, assign, b, d2, n2, r, &mut z2);
    relu(&mut z2);
    // fc3(int8 稠密)+ relu
    let x3 = with_bias(&z2, b, n2);
    let n3 = q_fc3.n_out;
    let mut z3 = vec![0.0f64; b * n3];
    q_fc3.forward_int8(&x3, b, &mut z3);
    relu(&mut z3);
    // head(int8 稠密,无 relu)
    let xh = with_bias(&z3, b, n3);
    q_head.forward_int8(&xh, b, y);
}

// ---------- 部署器 ----------
struct Deployer {
    fc2_dense: Option<Mat>,
    func_order: Option<i64>,
    func_r: Option<usize>,
    func_assign: Option<Vec<usize>>,
    func_centroids: Option<Mat>,
    func_means: Option<Mat>,
    func_comp2d: Option<Mat>,
    func_coeffs: Option<Mat>,
    fc1: Mat,
    fc3: Mat,
    head: Mat,
}

impl Deployer {
    /// 前向:x(b, 784) → y(b, 60)。use_func=true 时 fc2 走函数前向。
    fn forward(&self, x: &[f64], b: usize, d_in: usize, y: &mut [f64], use_func: bool) {
        // fc1 稠密前向
        let n1 = self.fc1.len();
        let d1 = self.fc1[0].len();
        let x1 = with_bias(x, b, d_in);
        let mut z1 = vec![0.0f64; b * n1];
        dense_forward(&x1, &self.fc1, b, d1, n1, &mut z1);
        relu(&mut z1);

        // fc2:稠密 or 函数前向
        let d2 = n1 + 1;
        let x2 = with_bias(&z1, b, n1);
        let n2_out = if use_func {
            // 函数化后 fc2 的输出单元数 = assign 长度(= n)
            self.func_assign.as_ref().unwrap().len()
        } else {
            self.fc2_dense.as_ref().unwrap().len()
        };
        let mut z2 = vec![0.0f64; b * n2_out];
        if use_func {
            match self.func_order.unwrap() {
                0 => func0_forward(&x2, self.func_centroids.as_ref().unwrap(),
                                   self.func_assign.as_ref().unwrap(),
                                   b, d2, n2_out, &mut z2),
                _ => {
                    let r = self.func_r.unwrap();
                    func1_forward(&x2, self.func_means.as_ref().unwrap(),
                                  self.func_comp2d.as_ref().unwrap(),
                                  self.func_coeffs.as_ref().unwrap(),
                                  self.func_assign.as_ref().unwrap(),
                                  b, d2, n2_out, r, &mut z2)
                }
            }
        } else {
            dense_forward(&x2, self.fc2_dense.as_ref().unwrap(), b, d2, n2_out, &mut z2);
        }
        relu(&mut z2);

        // fc3 稠密前向
        let n3 = self.fc3.len();
        let d3 = self.fc3[0].len();
        let x3 = with_bias(&z2, b, n2_out);
        let mut z3 = vec![0.0f64; b * n3];
        dense_forward(&x3, &self.fc3, b, d3, n3, &mut z3);
        relu(&mut z3);

        // head 稠密前向(无 relu,logits 输出)
        let nh = self.head.len();
        let dh = self.head[0].len();
        let xh = with_bias(&z3, b, n3);
        dense_forward(&xh, &self.head, b, dh, nh, y);
    }
}

// ---------- 测量 ----------
fn argmax(a: &[f64]) -> usize {
    let mut mi = 0usize;
    let mut mv = a[0];
    for (i, &v) in a.iter().enumerate() {
        if v > mv {
            mv = v;
            mi = i;
        }
    }
    mi
}

fn sigmoid(x: f64) -> f64 {
    1.0 / (1.0 + (-x).exp())
}

fn bench<F: FnMut()>(mut f: F, iters: usize) -> f64 {
    for _ in 0..3 {
        f();
    }
    let t0 = Instant::now();
    for _ in 0..iters {
        f();
    }
    t0.elapsed().as_secs_f64() / iters as f64 * 1e6
}

fn main() {
    let args: Vec<String> = std::env::args().collect();
    let art_dir = if args.len() > 1 { args[1].clone() } else { "../artifacts".to_string() };

    // 读二进制生成物(禁用 JSON 存生成物)
    let model_bytes = read_bytes(&format!("{}/model.bin", art_dir));
    let func_bytes = read_bytes(&format!("{}/functional.bin", art_dir));

    // 路线 A:model.bin → 4 层稠密(fc1/fc2/fc3/head)
    let (dense_layers, _) = parse_dense_layers(&model_bytes, 0);
    assert_eq!(dense_layers.len(), 4, "model.bin 应有 4 层");
    let fc2_dense = dense_layers[1].clone(); // 路线 A 的 f2 用 model 稠密;fc1/fc3/head 两线一致,取 functional.bin 侧

    // 路线 A·int8:model.bin → 4 层 int8 量化矩阵(直接 int8 域累加,不反量化)
    let (q_layers, _) = parse_dense_layers_q(&model_bytes, 0);
    assert_eq!(q_layers.len(), 4, "model.bin 应有 4 层(int8)");

    // 路线 B:functional.bin → 3 层稠密(fc1/fc3/head) + NRFN 函数块
    let (func_dense, off) = parse_dense_layers(&func_bytes, 0);
    assert_eq!(func_dense.len(), 3, "functional.bin 前部应有 3 层稠密");
    let fc1_f = func_dense[0].clone();
    let fc3_f = func_dense[1].clone();
    let head_f = func_dense[2].clone();
    let func_blob_len = u64_at(&func_bytes, off) as usize;
    let func_off = off + 8;
    // 从整个 func_bytes 里切出 NRFN 块,解析
    let (order, func_k, r, func_d, func_n, assign, centroids, means, comp2d, coeffs) =
        parse_nrfn(&func_bytes, func_off);
    assert_eq!(func_off + func_blob_len, func_bytes.len(), "functional.bin 长度不匹配");

    // stats.json 是配置/统计(非生成物,保留 JSON)
    let stats: HashMap<String, serde_json::Value> =
        serde_json::from_str(&fs::read_to_string(format!("{}/stats.json", art_dir)).unwrap())
            .unwrap();

    // 部署器:fc1/fc3/head 用 functional.bin 的三层(与 model.bin 一致),fc2 用 model 稠密或函数块
    let dep = Deployer {
        fc2_dense: Some(fc2_dense),
        func_order: Some(order),
        func_r: Some(r),
        func_assign: Some(assign),
        func_centroids: centroids,
        func_means: means,
        func_comp2d: comp2d,
        func_coeffs: coeffs,
        fc1: fc1_f,
        fc3: fc3_f,
        head: head_f,
    };

    let d_in = stats["n_input"].as_u64().unwrap() as usize;
    let out_dim = dep.head.len();
    let n_per_cell = stats["n_per_cell"].as_u64().unwrap() as usize;
    let n_class = stats["n_class"].as_u64().unwrap() as usize;
    let grid2 = out_dim / n_per_cell;

    // 测试输入(确定性随机,两路线共用)
    let mut rng = SimpleRng(12345u64);
    let b_single = 1usize;
    let b_batch = 64usize;
    let x_single: Vec<f64> = (0..b_single * d_in).map(|_| rng.next()).collect();
    let x_batch: Vec<f64> = (0..b_batch * d_in).map(|_| rng.next()).collect();

    // 精度对比
    let mut y_a = vec![0.0f64; b_batch * out_dim];
    let mut y_b = vec![0.0f64; b_batch * out_dim];
    dep.forward(&x_batch, b_batch, d_in, &mut y_a, false);
    dep.forward(&x_batch, b_batch, d_in, &mut y_b, true);

    let mut obj_agree = 0usize;
    let mut cls_agree = 0usize;
    let mut box_mad = 0.0f64;
    let total = b_batch * grid2;
    for bi in 0..b_batch {
        for g in 0..grid2 {
            let base = bi * out_dim + g * n_per_cell;
            let (oa, ob) = (sigmoid(y_a[base]), sigmoid(y_b[base]));
            if (oa > 0.5) == (ob > 0.5) {
                obj_agree += 1;
            }
            let ca = argmax(&y_a[base + 1..base + 1 + n_class]);
            let cb = argmax(&y_b[base + 1..base + 1 + n_class]);
            if ca == cb {
                cls_agree += 1;
            }
            for j in 0..4 {
                box_mad += (y_a[base + 1 + n_class + j]
                    - y_b[base + 1 + n_class + j]).abs();
            }
        }
    }
    let obj_ratio = obj_agree as f64 / total as f64;
    let cls_ratio = cls_agree as f64 / total as f64;
    let box_mad_avg = box_mad / (total * 4) as f64;

    // CPU 耗时
    let mut y1 = vec![0.0f64; out_dim];
    let mut yb = vec![0.0f64; b_batch * out_dim];
    let t_a_single = bench(|| dep.forward(&x_single, b_single, d_in, &mut y1, false), 200);
    let t_b_single = bench(|| dep.forward(&x_single, b_single, d_in, &mut y1, true), 200);
    let t_a_batch = bench(|| dep.forward(&x_batch, b_batch, d_in, &mut yb, false), 20);
    let t_b_batch = bench(|| dep.forward(&x_batch, b_batch, d_in, &mut yb, true), 20);

    // CPU int8 前向:int8 域累加(路线 A 稠密 int8),数值自检 + 耗时
    let mut yi = vec![0.0f64; b_batch * out_dim];
    forward_int8_net(&q_layers, &x_batch, b_batch, &mut yi);
    let mut int8_err = 0.0f64;
    for (i, &v) in yi.iter().enumerate() {
        let d = (v - y_a[i]).abs();
        if d > int8_err { int8_err = d; }
    }
    let mut yi1 = vec![0.0f64; out_dim];
    let mut yib = vec![0.0f64; b_batch * out_dim];
    let t_int8_single = bench(|| forward_int8_net(&q_layers, &x_single, b_single, &mut yi1), 200);
    let t_int8_batch = bench(|| forward_int8_net(&q_layers, &x_batch, b_batch, &mut yib), 20);

    // CPU 函数化 int8 前向:fc2 层函数化 int8(共享 GEMM 走 dot_i8),数值自检 + 耗时
    let func_means = dep.func_means.as_ref().unwrap();
    let func_comp = dep.func_comp2d.as_ref().unwrap();
    let func_coeffs = dep.func_coeffs.as_ref().unwrap();
    let func_assign = dep.func_assign.as_ref().unwrap();
    let mut yfi = vec![0.0f64; b_batch * out_dim];
    forward_func_int8_net(&q_layers[0], &q_layers[2], &q_layers[3],
                          func_means, func_comp, func_coeffs, func_assign,
                          r, &x_batch, b_batch, &mut yfi);
    let mut func_i8_err = 0.0f64;
    for (i, &v) in yfi.iter().enumerate() {
        let d = (v - y_b[i]).abs();
        if d > func_i8_err { func_i8_err = d; }
    }
    let mut yfi1 = vec![0.0f64; out_dim];
    let mut yfib = vec![0.0f64; b_batch * out_dim];
    let t_fi_single = bench(|| forward_func_int8_net(
        &q_layers[0], &q_layers[2], &q_layers[3],
        func_means, func_comp, func_coeffs, func_assign,
        r, &x_single, b_single, &mut yfi1), 200);
    let t_fi_batch = bench(|| forward_func_int8_net(
        &q_layers[0], &q_layers[2], &q_layers[3],
        func_means, func_comp, func_coeffs, func_assign,
        r, &x_batch, b_batch, &mut yfib), 20);

    // GPU 耗时(权重驻留 GPU,每次只传 x;cuBLAS + gather kernel)
    // 函数化参数(借用 dep 的字段,不 clone)
    let gpu_dep = gpu::GpuDeployer::new(
        &dep.fc1, dep.fc2_dense.as_ref().unwrap(), &dep.fc3, &dep.head,
        order, func_k, r, func_d, func_n,
        &dep.func_assign.as_ref().unwrap(),
        dep.func_centroids.as_ref(),
        dep.func_means.as_ref(),
        dep.func_comp2d.as_ref(),
        dep.func_coeffs.as_ref(),
    );
    let t_ga_single = bench(|| { gpu_dep.forward(&x_single, b_single, d_in, out_dim, false); }, 200);
    let t_gb_single = bench(|| { gpu_dep.forward(&x_single, b_single, d_in, out_dim, true); }, 200);
    let t_ga_batch = bench(|| { gpu_dep.forward(&x_batch, b_batch, d_in, out_dim, false); }, 20);
    let t_gb_batch = bench(|| { gpu_dep.forward(&x_batch, b_batch, d_in, out_dim, true); }, 20);

    // GPU 数值自检:GPU 稠密/函数化输出必须与 CPU 一致
    let gpu_ya = gpu_dep.forward(&x_batch, b_batch, d_in, out_dim, false);
    let gpu_yb = gpu_dep.forward(&x_batch, b_batch, d_in, out_dim, true);
    let mut gpu_err = 0.0f64;
    for (i, &v) in gpu_ya.iter().enumerate() {
        let d = (v - y_a[i]).abs();
        if d > gpu_err { gpu_err = d; }
    }
    for (i, &v) in gpu_yb.iter().enumerate() {
        let d = (v - y_b[i]).abs();
        if d > gpu_err { gpu_err = d; }
    }

    // GPU int8 稠密前向(Tensor Core gemm_ex(CUDA_R_8I),权重 int8 驻留)
    let gpu_i8 = gpu::GpuDeployerI8::new(
        &dep.fc1, dep.fc2_dense.as_ref().unwrap(), &dep.fc3, &dep.head,
    );
    let t_gi_single = bench(|| { gpu_i8.forward(&x_single, b_single, d_in); }, 200);
    let t_gi_batch = bench(|| { gpu_i8.forward(&x_batch, b_batch, d_in); }, 20);
    // GPU int8 数值自检:vs CPU f64(y_a)与 CPU int8(yi)
    let gpu_yi = gpu_i8.forward(&x_batch, b_batch, d_in);
    let mut gpu_i8_err = 0.0f64;
    for (i, &v) in gpu_yi.iter().enumerate() {
        let d = (v - y_a[i]).abs();
        if d > gpu_i8_err { gpu_i8_err = d; }
    }
    let mut gpu_i8_err_i = 0.0f64;
    for (i, &v) in gpu_yi.iter().enumerate() {
        let d = (v - yi[i]).abs();
        if d > gpu_i8_err_i { gpu_i8_err_i = d; }
    }

    // GPU int8 函数化前向(fc2 函数化 int8,shared GEMM 走 gemm_ex + gather f64)
    let gpu_fi = gpu::GpuDeployerI8Func::new(
        &dep.fc1, &dep.fc3, &dep.head,
        order, func_k, r, func_d, func_n,
        &dep.func_assign.as_ref().unwrap(),
        dep.func_centroids.as_ref(),
        dep.func_means.as_ref(),
        dep.func_comp2d.as_ref(),
        dep.func_coeffs.as_ref(),
    );
    let t_gfi_single = bench(|| { gpu_fi.forward(&x_single, b_single, d_in); }, 200);
    let t_gfi_batch = bench(|| { gpu_fi.forward(&x_batch, b_batch, d_in); }, 20);
    let gpu_yfi = gpu_fi.forward(&x_batch, b_batch, d_in);
    let mut gpu_fi_err = 0.0f64;
    for (i, &v) in gpu_yfi.iter().enumerate() {
        let d = (v - y_b[i]).abs();
        if d > gpu_fi_err { gpu_fi_err = d; }
    }

    // 体积 / 内存
    let vol_dense = fs::metadata(format!("{}/model.bin", art_dir)).unwrap().len();
    let vol_func = fs::metadata(format!("{}/functional.bin", art_dir)).unwrap().len();
    let n1 = dep.fc1.len();               // fc1 输出数(= 4096)
    let n2 = dep.fc2_dense.as_ref().unwrap().len();  // fc2 输出数(= 2048)
    let mem_fc2_dense = n2 * (n1 + 1) * 8;
    let (fk, fr, fd, fn_) = {
        // 从函数参数还原 fc2 函数内存:order/K/r/D/n
        let k = dep.func_centroids.as_ref().map(|c| c.len())
            .or_else(|| dep.func_means.as_ref().map(|m| m.len())).unwrap();
        let r = dep.func_r.unwrap();
        let d = dep.func_means.as_ref().map(|m| m[0].len())
            .unwrap_or_else(|| dep.func_centroids.as_ref().unwrap()[0].len());
        let n = dep.func_assign.as_ref().unwrap().len();
        (k, r, d, n)
    };
    let mem_fc2_func = match dep.func_order.unwrap() {
        0 => fk * fd * 8 + fn_ * 8,
        _ => (fk * fd + fk * fr * fd + fn_ * fr) * 8,
    };
    let mem_others = (dep.fc1.len() * dep.fc1[0].len()
        + dep.fc3.len() * dep.fc3[0].len()
        + dep.head.len() * dep.head[0].len()) * 8;
    let mem_dense_total = n2 * (n1 + 1) * 8 + mem_others;
    let mem_func_total = mem_fc2_func + mem_others;

    // 输出比较日志
    let mut log = String::new();
    log.push_str("==============================================================\n");
    log.push_str(" nrfunc 多对象检测部署对比(Rust 独立二进制)\n");
    log.push_str(" 任务:MNIST 多对象检测(YOLO式 2×2 网格,每格1对象)\n");
    log.push_str(" 模型:784 → 4096 → 2048 → 1024 → 网格头(60)\n");
    log.push_str("==============================================================\n\n");

    log.push_str("[路线 A] 成熟模型稠密部署(model.bin)\n");
    log.push_str("[路线 B] 函数化生成物函数部署(functional.bin,fc2 层函数化)\n\n");

    log.push_str("----------------------------------------------------------------\n");
    log.push_str("一、体积(磁盘占用)\n");
    log.push_str(&format!("  路线 A model.bin      {} B\n", vol_dense));
    log.push_str(&format!("  路线 B functional.bin {} B(fc1/fc3/head 稠密 + fc2 区域函数)\n", vol_func));
    log.push_str(&format!("  fc2 层权重内存:dense {} B vs 函数化 {} B(省 {:.1}%)\n",
        mem_fc2_dense, mem_fc2_func,
        (1.0 - mem_fc2_func as f64 / mem_fc2_dense as f64) * 100.0));
    log.push_str(&format!("  全模型权重内存:dense {} B vs 函数化 {} B(省 {:.1}%)\n",
        mem_dense_total, mem_func_total,
        (1.0 - mem_func_total as f64 / mem_dense_total as f64) * 100.0));
    log.push_str("\n");

    log.push_str("----------------------------------------------------------------\n");
    log.push_str(&format!("二、精度(同一批 {} 个输入,两路线输出对比)\n", b_batch));
    log.push_str(&format!("  objectness 命中一致率 = {:.4}%\n", obj_ratio * 100.0));
    log.push_str(&format!("  分类 argmax 一致率    = {:.4}%\n", cls_ratio * 100.0));
    log.push_str(&format!("  bbox 平均绝对差      = {:.6}\n", box_mad_avg));
    log.push_str(&format!("  (Python 侧真实掉点   = {} pt,见 stats.json)\n",
        stats["drop_pt"].as_f64().unwrap_or(0.0)));
    log.push_str("\n");

    log.push_str("----------------------------------------------------------------\n");
    log.push_str("三、CPU 耗时(Rust 同实现质量,同机同批;µs)\n");
    log.push_str(&format!("  单样本:稠密 {} µs vs 函数化 {} µs({:.2}x)\n",
        t_a_single, t_b_single, t_a_single / t_b_single));
    log.push_str(&format!("  批量({}):稠密 {} µs vs 函数化 {} µs({:.2}x)\n",
        b_batch, t_a_batch, t_b_batch, t_a_batch / t_b_batch));
    log.push_str(&format!("  int8 稠密前向(int8 域累加):单样本 {} µs vs 批量 {} µs,\n",
        t_int8_single, t_int8_batch));
    log.push_str(&format!("    int8 vs f64 稠密:单样本 {:.2}x / 批量 {:.2}x;int8 前向数值最大绝对差 = {:.3e}\n",
        t_a_single / t_int8_single, t_a_batch / t_int8_batch, int8_err));
    log.push_str(&format!("  函数化 int8 前向(fc2 函数化 int8):单样本 {} µs vs 批量 {} µs,\n",
        t_fi_single, t_fi_batch));
    log.push_str(&format!("    函数化 int8 vs 函数化 f64:单样本 {:.2}x / 批量 {:.2}x;\n",
        t_b_single / t_fi_single, t_b_batch / t_fi_batch));
    log.push_str(&format!("    函数化 int8 数值自检(vs 函数化 f64):最大绝对差 = {:.3e}\n",
        func_i8_err));
    log.push_str("\n");

    log.push_str("----------------------------------------------------------------\n");
    log.push_str("四、GPU 耗时(cudarc + cuBLAS,权重驻留 GPU,每次只传 x;µs)\n");
    log.push_str(&format!("  单样本:稠密 {} µs vs 函数化 {} µs({:.2}x)\n",
        t_ga_single, t_gb_single, t_ga_single / t_gb_single));
    log.push_str(&format!("  批量({}):稠密 {} µs vs 函数化 {} µs({:.2}x)\n",
        b_batch, t_ga_batch, t_gb_batch, t_ga_batch / t_gb_batch));
    log.push_str(&format!("  GPU vs CPU 数值自检:输出最大绝对差 = {:.3e}\n", gpu_err));
    log.push_str(&format!("  GPU int8 Tensor Core(gemm_ex CUDA_R_8I):单样本 {} µs vs 批量 {} µs,\n",
        t_gi_single, t_gi_batch));
    log.push_str(&format!("    GPU int8 vs GPU f64 稠密:单样本 {:.2}x / 批量 {:.2}x;\n",
        t_ga_single / t_gi_single, t_ga_batch / t_gi_batch));
    log.push_str(&format!("    GPU int8 数值自检:vs CPU f64 最大绝对差 = {:.3e},vs CPU int8 最大绝对差 = {:.3e}\n",
        gpu_i8_err, gpu_i8_err_i));
    log.push_str(&format!("  GPU int8 函数化前向:单样本 {} µs vs 批量 {} µs;\n",
        t_gfi_single, t_gfi_batch));
    log.push_str(&format!("    GPU int8 函数化 vs GPU f64 函数化:单样本 {:.2}x / 批量 {:.2}x;\n",
        t_gb_single / t_gfi_single, t_gb_batch / t_gfi_batch));
    log.push_str(&format!("    GPU int8 函数化数值自检(vs 函数化 f64):最大绝对差 = {:.3e}\n",
        gpu_fi_err));
    log.push_str("  (GPU 前向含 H2D(x)+D2H(y) 传输;int8 半量化含每层 host 侧量化开销)\n");
    log.push_str("\n");

    log.push_str("----------------------------------------------------------------\n");
    log.push_str("五、结论\n");
    log.push_str("  函数化收益集中在体积/内存;CPU 提速取决于 K≪n;GPU 提速取决于\n");
    log.push_str("  层规模(大层才超传输开销)。本表由 Rust 二进制独立跑出。\n");
    log.push_str("==============================================================\n");

    println!("{}", log);
    fs::write(format!("{}/comparison.log", art_dir), log).unwrap();
    eprintln!("[rust-deploy] 比较日志已写入 {}/comparison.log", art_dir);
}

// ---------- 确定性随机 ----------
struct SimpleRng(u64);
impl SimpleRng {
    fn next(&mut self) -> f64 {
        self.0 = self.0.wrapping_mul(6364136223846793005).wrapping_add(1442695040888963407);
        let v = (self.0 >> 11) as f64 / (1u64 << 53) as f64;
        v * 0.5
    }
}