// forward.rs —— 读取 nrfunc 二进制生成物,执行 0 阶函数前向。
//
// 零依赖(仅标准库 std)。做的事和 README 里说的一样:
//   1) 读 artifact.bin 全部字节
//   2) 按「小端」解析头部 + centroids + assign
//   3) y[i] = Σ x[j] * centroids[assign[i]][j]
//
// 运行:rustc forward.rs -O -o forward && ./forward
use std::fs;

// 从字节流里按小端读一个 u32(把 [u8] 拷到 [u8;4] 再交给 from_le_bytes)
fn read_u32(data: &[u8], pos: usize) -> u32 {
    let mut buf = [0u8; 4];
    buf.copy_from_slice(&data[pos..pos + 4]);
    u32::from_le_bytes(buf)
}

fn main() {
    // 1) 读二进制
    let data = fs::read("artifact.bin").expect("读 artifact.bin 失败");

    // 2) 解析头部(小端)
    //    magic "NRFN"(4B) | version(1) | order(1) | bits(1) | K(4) | n(4) | D(4) | r(4)
    //    偏移:  0..4        4           5          6        7..11   11..15  15..19  19..23
    assert_eq!(&data[0..4], b"NRFN", "不是 nrfunc 生成物");
    let order = data[5];
    // let bits = data[6];  // 本示例固定 bits=32,无需读
    let k = read_u32(&data, 7) as usize;
    let n = read_u32(&data, 11) as usize;
    let d = read_u32(&data, 15) as usize;
    // let r = read_u32(&data, 19) as usize; // 仅 order=1

    assert_eq!(order, 0, "本示例只演示 order=0");

    // 3) 读 centroids(K×D 个 float32)和 assign(n 个 int32)
    //    小端 float32:把连续 4 字节按小端拼成 u32 位模式,再转 f32
    let mut pos = 23usize;
    let mut centroids = vec![vec![0f32; d]; k];
    for i in 0..k {
        for j in 0..d {
            centroids[i][j] = f32::from_bits(read_u32(&data, pos));
            pos += 4;
        }
    }
    let mut assign = vec![0i32; n];
    for i in 0..n {
        assign[i] = read_u32(&data, pos) as i32;
        pos += 4;
    }

    // 4) 对输入 x 算前向
    let x = [0.5f32, -0.25f32, 1.0f32]; // 与 input.txt 一致
    let mut y = vec![0f32; n];
    for i in 0..n {
        let c = &centroids[assign[i] as usize];
        let mut s = 0f32;
        for j in 0..d {
            s += x[j] * c[j];
        }
        y[i] = s;
    }

    // 打印结果,对照 input.txt 里的 y_expected
    print!("y =");
    for v in y {
        print!(" {:.6}", v);
    }
    println!();
}