use rane::mil::{self, OutputDtype};
use rane::{f32_to_fp16, Buffer, Program};
fn main() -> Result<(), Box<dyn std::error::Error>> {
let oc = 64usize;
let seq = 64usize;
if std::env::var("DUMP_MIL").is_ok() {
let p = mil::matmul_split(64, oc, seq, 2);
println!("---- MIL for split(64,64,64,2) ----\n{}", p.text);
return Ok(());
}
println!("ANE split-matmul + fp32 sum probe");
println!(" hypothesis: matmul fp16 accum saturates at ~32768, but fp32 add is unconstrained");
println!(" oc={oc}, seq={seq}\n");
println!("--- structural test: ic=64, v=2, expected=256 (fp16 sum) ---");
test_split("fp16 1chunk", 64, oc, seq, 2.0, 1, 256, true)?;
test_split("fp16 2chunk", 64, oc, seq, 2.0, 2, 256, true)?;
test_split("fp32 2chunk", 64, oc, seq, 2.0, 2, 256, false)?;
println!("\n--- ic=64, v=31, expected=61504 (single fail, 2-chunk should win) ---");
test_split("single fp32", 64, oc, seq, 31.0, 1, 61504, false)?;
test_split("fp16 2chunk", 64, oc, seq, 31.0, 2, 61504, true)?;
test_split("fp32 2chunk", 64, oc, seq, 31.0, 2, 61504, false)?;
test_split("fp32 4chunk", 64, oc, seq, 31.0, 4, 61504, false)?;
println!("\n--- ic=128, v=31, expected=123008 ---");
test_split("fp32 4chunk", 128, oc, seq, 31.0, 4, 123008, false)?;
println!("\n--- ic=2176, v=31, expected=2091136 (Pearl full rank) ---");
test_split("fp32 68chunk", 2176, oc, seq, 31.0, 68, 2091136, false)?;
Ok(())
}
fn test_split(
label: &str,
ic: usize,
oc: usize,
seq: usize,
fill_val: f32,
n_chunks: usize,
expected: i32,
use_fp16_sum: bool,
) -> Result<(), Box<dyn std::error::Error>> {
let program = if n_chunks == 1 && !use_fp16_sum {
mil::matmul_cast(ic, oc, seq, OutputDtype::Fp32)
} else if use_fp16_sum {
mil::matmul_split_fp16(ic, oc, seq, n_chunks)
} else {
mil::matmul_split(ic, oc, seq, n_chunks)
};
let mut model = match Program::compile(&program, &[]) {
Ok(m) => m,
Err(e) => {
println!(
" [{label:14}] compile FAILED: {}",
short_err(&format!("{e}"))
);
return Ok(());
}
};
if let Err(e) = model.load() {
println!(
" [{label:14}] load FAILED: {}",
short_err(&format!("{e}"))
);
return Ok(());
}
let input = Buffer::new(program.input_size())?;
let output = Buffer::new(program.output_size())?;
let sp = program.input_spatial;
let vfill = f32_to_fp16(fill_val);
input.write(|data| {
for d in data.iter_mut() {
*d = 0;
}
for ch in 0..ic {
for s in 0..seq {
data[ch * sp + s] = vfill;
}
for o in 0..oc {
data[ch * sp + seq + o] = vfill;
}
}
});
if let Err(e) = model.run(&input, &output) {
println!(
" [{label:14}] eval FAILED: {}",
short_err(&format!("{e}"))
);
return Ok(());
}
let (val_str, correct) = if program.output_dtype == OutputDtype::Fp16 {
output.read(|data| {
let n = oc * seq;
let v0 = rane::fp16_to_f32(data[0]);
let n_correct = data[..n]
.iter()
.filter(|&&v| rane::fp16_to_f32(v) as i32 == expected)
.count();
(format!("{v0:.1}"), n_correct == n)
})
} else {
output.read_f32(|data| {
let n = oc * seq;
let v0 = data[0];
let n_correct = data[..n].iter().filter(|&&v| v as i32 == expected).count();
(format!("{v0:.1}"), n_correct == n)
})
};
let status = if correct {
"ALL CORRECT β"
} else {
"WRONG β"
};
println!(" [{label:14}] got={val_str:>12} expected={expected:>8} β {status}");
Ok(())
}
fn short_err(s: &str) -> String {
if let Some(idx) = s.find("status=") {
let tail: String = s[idx..].chars().take(60).collect();
return tail;
}
if let Some(idx) = s.find("err=(") {
let tail: String = s[idx..].chars().take(80).collect();
return tail;
}
s.chars().take(120).collect()
}