use crate::error::FocrResult;
use super::connector;
use super::decoder;
use super::tensor::Mat;
use super::vision_sam::{self, Linear};
use super::weights::Weights;
#[cfg(not(target_arch = "wasm32"))]
use std::time::Instant;
#[cfg(target_arch = "wasm32")]
use web_time::Instant;
pub const VISION_TOKENS: usize = 256;
pub const HIDDEN: usize = 768;
pub fn vision_features(
weights: &Weights,
statics: &OnechartStatics,
image: &Mat,
) -> FocrResult<Mat> {
let side = (image.cols as f64).sqrt() as usize;
if side * side != image.cols || image.rows != 3 {
return Err(crate::FocrError::Other(anyhow::anyhow!(
"onechart vision: expected [3, side*side] input, got [{}, {}]",
image.rows,
image.cols
)));
}
let sam = match statics.sam.as_ref() {
Some(tower) => vision_sam::forward_with(tower, image, side, side)?, None => vision_sam::forward_streamed(weights, image, &statics.prefix)?,
};
let sam_t = transpose(&sam); statics.proj.apply(&sam_t) }
pub struct OnechartStatics {
pub sam: Option<vision_sam::SamWeights>,
pub prefix: String,
pub proj: Linear,
pub embed: Mat,
}
pub fn hydrate_statics(
weights: &Weights,
prefix: &str,
stream_vision: bool,
) -> FocrResult<OnechartStatics> {
let th = Instant::now();
let statics = OnechartStatics {
sam: if stream_vision {
None
} else {
Some(vision_sam::sam_weights_from(weights, prefix)?)
},
prefix: prefix.to_string(),
proj: Linear::from_row_major(
&weights.vec("model.mm_projector.weight")?,
weights.vec("model.mm_projector.bias")?,
HIDDEN,
1024,
)?,
embed: weights.mat("model.decoder.embed_tokens.weight")?,
};
super::timing_log(&format!(
" onechart.hydrate({}) {:.2}s",
if stream_vision { "streamed" } else { "cached" },
th.elapsed().as_secs_f64()
));
Ok(statics)
}
pub fn build_inputs_embeds(
statics: &OnechartStatics,
vision: &Mat,
prompt_ids: &[u32],
) -> FocrResult<Mat> {
let embed = &statics.embed;
let (vocab, hidden) = (embed.rows, embed.cols);
let mut inputs_embeds = decoder::embed_tokens(&embed.data, vocab, hidden, prompt_ids)?;
let mask: Vec<bool> = prompt_ids
.iter()
.map(|&id| id == crate::tokenizer::special_opt::IMG_PAD)
.collect();
connector::masked_scatter(&mut inputs_embeds, vision, &mask)?;
Ok(inputs_embeds)
}
pub fn number_head(weights: &Weights, hidden_row: &[f32]) -> FocrResult<Vec<f32>> {
let lin = |i: usize, out, in_| -> FocrResult<Linear> {
Linear::from_row_major(
&weights.vec(&format!("num_decoder.{i}.weight"))?,
weights.vec(&format!("num_decoder.{i}.bias"))?,
out,
in_,
)
};
let x = Mat::from_vec(1, HIDDEN, hidden_row.to_vec());
let mut x = lin(0, 384, HIDDEN)?.apply(&x)?;
super::nn::relu(&mut x);
let mut x = lin(2, 384, 384)?.apply(&x)?;
super::nn::relu(&mut x);
Ok(lin(4, 256, 384)?.apply(&x)?.data)
}
pub const RELIABLE_THRESHOLD: f64 = 0.1;
#[must_use]
pub fn extract_gt_values(values: &serde_json::Value) -> Option<Vec<f64>> {
fn walk(v: &serde_json::Value, out: &mut Vec<f64>) -> bool {
match v {
serde_json::Value::Object(map) => map.values().all(|x| walk(x, out)),
serde_json::Value::Array(_) => false,
serde_json::Value::Number(n) => {
if let Some(f) = n.as_f64() {
out.push(f);
}
true
}
serde_json::Value::String(s) => {
let cleaned = strip_index_spans(s);
let filtered: String = cleaned
.chars()
.filter(|c| c.is_ascii_digit() || *c == '.' || *c == '-')
.collect();
if matches!(filtered.as_str(), "-" | "*" | "none" | "None" | "") {
return true;
}
if let Ok(f) = filtered.parse::<f64>() {
out.push(f);
}
true
}
_ => true,
}
}
let mut out = Vec::new();
walk(values, &mut out).then_some(out)
}
fn strip_index_spans(s: &str) -> String {
let chars: Vec<char> = s.chars().collect();
let mut out = String::with_capacity(s.len());
let mut i = 0;
while i < chars.len() {
let (open, close) = match chars[i] {
'(' => ('(', ')'),
'[' => ('[', ']'),
_ => {
out.push(chars[i]);
i += 1;
continue;
}
};
let _ = open;
let mut j = i + 1;
while j < chars.len() && chars[j].is_ascii_digit() {
j += 1;
}
if j > i + 1 && j < chars.len() && chars[j] == close {
i = j + 1; } else {
out.push(chars[i]);
i += 1;
}
}
out
}
#[must_use]
pub fn normalize_gt(xs: &[f64]) -> Vec<f64> {
let round4 = |x: f64| (x * 10_000.0).round_ties_even() / 10_000.0;
if xs.len() < 2 {
return xs.iter().map(|&x| round4(x)).collect();
}
let (lo, hi) = xs
.iter()
.fold((f64::INFINITY, f64::NEG_INFINITY), |(l, h), &x| {
(l.min(x), h.max(x))
});
xs.iter()
.map(|&x| round4((x - lo) / (hi - lo + 1e-9)))
.collect()
}
#[must_use]
pub fn reliable_distance(pred_locs: &[f32], gt: &[f64]) -> f64 {
if gt.is_empty() {
return f64::INFINITY;
}
let n = gt.len().min(pred_locs.len());
gt[..n]
.iter()
.zip(&pred_locs[..n])
.map(|(&g, &p)| (g - f64::from(p)).abs())
.sum::<f64>()
/ n as f64
}
pub fn chart_prompt_ids(tk: &crate::tokenizer::Tokenizer) -> FocrResult<Vec<u32>> {
let imgpad = "<imgpad>".repeat(VISION_TOKENS);
let prompt = format!(
"A chat between a curious user and an artificial intelligence assistant. \
The assistant gives helpful, detailed, and polite answers to the user's \
questions. USER: <img>{imgpad}</img>Convert the key information of the \
chart to a python dict:\n ASSISTANT:"
);
tk.encode(&prompt)
}
#[must_use]
pub fn complete_json_string(s: &str) -> String {
let mut depth = 0i32;
let mut in_str = false;
let mut esc = false;
for c in s.chars() {
if esc {
esc = false;
continue;
}
match c {
'\\' if in_str => esc = true,
'"' => in_str = !in_str,
'{' if !in_str => depth += 1,
'}' if !in_str => depth -= 1,
_ => {}
}
}
let mut out = s.to_string();
if in_str {
out.push('"');
}
for _ in 0..depth.max(0) {
out.push('}');
}
out
}
#[derive(Debug, Clone)]
pub struct ChartResult {
pub json_text: String,
pub pred_locs: Option<Vec<f32>>,
pub reliable_distance: Option<f64>,
pub reliable: Option<bool>,
}
pub fn recognize(
weights: &Weights,
statics: &OnechartStatics,
tk: &crate::tokenizer::Tokenizer,
img: &image::DynamicImage,
max_new: usize,
) -> FocrResult<ChartResult> {
let tv = Instant::now();
let image = crate::preprocess::onechart_view_tensor(img);
let vision = vision_features(weights, statics, &image)?;
let prompt_ids = chart_prompt_ids(tk)?;
let embeds = build_inputs_embeds(statics, &vision, &prompt_ids)?;
super::timing_log(&format!(
" onechart.vision+splice {:.2}s",
tv.elapsed().as_secs_f64()
));
let tg = Instant::now();
let cfg = super::decoder_qwen2::DecoderConfig::onechart();
let max_new = max_new.min(4096usize.saturating_sub(embeds.rows));
let ids = super::decoder_qwen2::generate_greedy_kvcache(
weights,
&cfg,
&embeds,
max_new,
crate::tokenizer::special_opt::BOS_EOS,
)?;
super::timing_log(&format!(
" onechart.generate {} tokens {:.2}s",
ids.len(),
tg.elapsed().as_secs_f64()
));
finish_recognition(weights, statics, tk, &cfg, &prompt_ids, &vision, ids)
}
fn finish_recognition(
weights: &Weights,
statics: &OnechartStatics,
tk: &crate::tokenizer::Tokenizer,
cfg: &super::decoder_qwen2::DecoderConfig,
prompt_ids: &[u32],
vision: &Mat,
ids: Vec<u32>,
) -> FocrResult<ChartResult> {
let pred_locs = match ids
.iter()
.position(|&id| id == crate::tokenizer::special_opt::NUMBER)
{
Some(pos) => {
let mut full = prompt_ids.to_vec();
full.extend_from_slice(&ids[..=pos]);
let embeds_tap = build_inputs_embeds(statics, vision, &full)?;
let hidden = super::decoder_qwen2::prefill_final_hidden(weights, cfg, &embeds_tap)?;
let last = &hidden.data[(hidden.rows - 1) * hidden.cols..];
let mut locs = number_head(weights, last)?;
locs.truncate(100);
Some(locs)
}
None => None,
};
let json_text = complete_json_string(tk.decode_skip_special(&ids)?.trim());
let (reliable_distance_v, reliable) = match (&pred_locs, parse_values(&json_text)) {
(Some(locs), Some(gt)) if !gt.is_empty() => {
let d = reliable_distance(locs, &normalize_gt(>));
(Some(d), Some(d < RELIABLE_THRESHOLD))
}
_ => (None, None),
};
Ok(ChartResult {
json_text,
pred_locs,
reliable_distance: reliable_distance_v,
reliable,
})
}
fn parse_values(json_text: &str) -> Option<Vec<f64>> {
let v: serde_json::Value = serde_json::from_str(json_text).ok()?;
let values = v.get("values").or_else(|| v.get("data"))?;
extract_gt_values(values)
}
fn transpose(m: &Mat) -> Mat {
let (r, c) = (m.rows, m.cols);
let mut out = vec![0.0f32; r * c];
for i in 0..r {
for j in 0..c {
out[j * r + i] = m.data[i * c + j];
}
}
Mat::from_vec(c, r, out)
}
pub fn recognize_batch(
weights: &Weights,
statics: &OnechartStatics,
tk: &crate::tokenizer::Tokenizer,
imgs: &[&image::DynamicImage],
max_new: usize,
) -> FocrResult<Vec<ChartResult>> {
let tv = Instant::now();
let prompt_ids = chart_prompt_ids(tk)?;
let mut visions: Vec<Mat> = Vec::with_capacity(imgs.len());
let mut embeds_list: Vec<Mat> = Vec::with_capacity(imgs.len());
let mut caps: Vec<usize> = Vec::with_capacity(imgs.len());
for img in imgs {
let image = crate::preprocess::onechart_view_tensor(img);
let vision = vision_features(weights, statics, &image)?;
let embeds = build_inputs_embeds(statics, &vision, &prompt_ids)?;
caps.push(max_new.min(4096usize.saturating_sub(embeds.rows)));
visions.push(vision);
embeds_list.push(embeds);
}
super::timing_log(&format!(
" onechart.vision+splice(batch of {}) {:.2}s",
imgs.len(),
tv.elapsed().as_secs_f64()
));
let tg = Instant::now();
let cfg = super::decoder_qwen2::DecoderConfig::onechart();
let id_streams = super::decoder_qwen2::generate_greedy_batched(
weights,
&cfg,
&embeds_list,
&caps,
crate::tokenizer::special_opt::BOS_EOS,
)?;
super::timing_log(&format!(
" onechart.generate(batch of {}) {} tokens {:.2}s",
imgs.len(),
id_streams.iter().map(Vec::len).sum::<usize>(),
tg.elapsed().as_secs_f64()
));
id_streams
.into_iter()
.zip(&visions)
.map(|(ids, vision)| {
finish_recognition(weights, statics, tk, &cfg, &prompt_ids, vision, ids)
})
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn transpose_round_trips() {
let m = Mat::from_vec(2, 3, vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]);
let t = transpose(&m);
assert_eq!((t.rows, t.cols), (3, 2));
assert_eq!(t.data, vec![1.0, 4.0, 2.0, 5.0, 3.0, 6.0]);
assert_eq!(transpose(&t).data, m.data);
}
#[test]
fn vision_features_error_handling() {
let w = Weights::default();
assert!(hydrate_statics(&w, "model.vision_tower", false).is_err());
}
#[test]
fn opt_prefill_matches_torch_oracle() {
let Ok(dir) = std::env::var("FOCR_ONECHART_DIR") else {
return;
};
let proj_path = format!("{dir}/onechart_proj_out.bin");
let logits_path = format!("{dir}/onechart_final_logits.bin");
let model_path = format!("{dir}/model.safetensors");
if !std::path::Path::new(&proj_path).is_file() {
eprintln!("skip-with-SUCCESS: {proj_path} absent (run the oracle script)");
return;
}
let read_f32 = |p: &str| -> Vec<f32> {
std::fs::read(p)
.expect("oracle blob reads")
.as_chunks::<4>()
.0
.iter()
.map(|c| f32::from_le_bytes(*c))
.collect()
};
let fx: serde_json::Value = serde_json::from_str(
&std::fs::read_to_string(concat!(
env!("CARGO_MANIFEST_DIR"),
"/tests/fixtures/onechart/oracle_fixtures.json"
))
.expect("oracle fixtures read"),
)
.expect("oracle fixtures parse");
let prompt_ids: Vec<u32> = fx["l0c_prompt"]["ids"]
.as_array()
.unwrap()
.iter()
.map(|v| v.as_u64().unwrap() as u32)
.collect();
assert_eq!(
prompt_ids.len(),
fx["l0c_prompt"]["n"].as_u64().unwrap() as usize,
"prompt drifted from its own fixture"
);
assert_eq!(prompt_ids.len(), 308, "measured census prompt length");
let weights = Weights::load(std::path::Path::new(&model_path)).expect("weights");
let statics = hydrate_statics(&weights, "model.vision_tower", false).expect("statics");
let vision = Mat::from_vec(VISION_TOKENS, HIDDEN, read_f32(&proj_path));
let embeds = build_inputs_embeds(&statics, &vision, &prompt_ids).expect("splice");
let cfg = super::super::decoder_qwen2::DecoderConfig::onechart();
let logits =
super::super::decoder_qwen2::forward_prefill(&weights, &cfg, &embeds).expect("prefill");
let ours = &logits.data[(logits.rows - 1) * logits.cols..];
let want = read_f32(&logits_path);
assert_eq!(ours.len(), want.len(), "vocab width");
let argmax = |v: &[f32]| {
v.iter()
.enumerate()
.max_by(|a, b| a.1.total_cmp(b.1))
.map(|(i, _)| i)
.unwrap()
};
let mut dot = 0.0f64;
let (mut na, mut nb) = (0.0f64, 0.0f64);
let mut max_abs = 0.0f64;
for (a, b) in ours.iter().zip(&want) {
let (a, b) = (f64::from(*a), f64::from(*b));
dot += a * b;
na += a * a;
nb += b * b;
max_abs = max_abs.max((a - b).abs());
}
let cos = dot / (na.sqrt() * nb.sqrt());
eprintln!(
"[D4 prefill] argmax={} (oracle {}) cos={cos:.8} maxabs={max_abs:.3e}",
argmax(ours),
argmax(&want)
);
assert_eq!(argmax(ours), argmax(&want), "next-token argmax diverged");
assert!(cos >= 0.9999, "prefill logit cosine {cos:.8} < 0.9999");
}
#[test]
fn complete_json_string_balances_braces() {
assert_eq!(
complete_json_string(r#"{"a": {"b": 1}"#),
r#"{"a": {"b": 1}}"#
);
assert_eq!(complete_json_string(r#"{"a": 1}"#), r#"{"a": 1}"#);
assert_eq!(complete_json_string(r#"{"a": "{{"#), r#"{"a": "{{"}"#);
assert_eq!(complete_json_string(""), "");
}
#[test]
fn chart_prompt_ids_match_oracle_l0c() {
let Ok(dir) = std::env::var("FOCR_ONECHART_DIR") else {
return;
};
let tk = crate::tokenizer::Tokenizer::from_opt_dir(std::path::Path::new(&dir))
.expect("onechart tokenizer");
let fx: serde_json::Value = serde_json::from_str(
&std::fs::read_to_string(concat!(
env!("CARGO_MANIFEST_DIR"),
"/tests/fixtures/onechart/oracle_fixtures.json"
))
.expect("oracle fixtures read"),
)
.expect("oracle fixtures parse");
let want: Vec<u32> = fx["l0c_prompt"]["ids"]
.as_array()
.unwrap()
.iter()
.map(|v| v.as_u64().unwrap() as u32)
.collect();
let got = chart_prompt_ids(&tk).expect("prompt encode");
assert_eq!(got, want, "chart prompt diverged from the 308-id oracle");
eprintln!("[D8 L0c] {} prompt ids exact", got.len());
}
#[test]
fn recognize_reads_the_committed_chart() {
let Ok(dir) = std::env::var("FOCR_ONECHART_DIR") else {
return;
};
let f32_path = format!("{dir}/model.safetensors");
if !std::path::Path::new(&f32_path).is_file() {
eprintln!("skip-with-SUCCESS: {f32_path} absent");
return;
}
let tk = crate::tokenizer::Tokenizer::from_opt_dir(std::path::Path::new(&dir))
.expect("onechart tokenizer");
let img = image::open(concat!(
env!("CARGO_MANIFEST_DIR"),
"/tests/fixtures/onechart/sample_chart.png"
))
.expect("sample chart decodes");
let weights = Weights::load(std::path::Path::new(&f32_path)).expect("f32 weights");
let statics = hydrate_statics(&weights, "model.vision_tower", false).expect("f32 statics");
let res = recognize(&weights, &statics, &tk, &img, 512).expect("recognize");
eprintln!("[D6 e2e f32] json: {}", res.json_text);
eprintln!(
"[D6 e2e f32] pred_locs[:4]: {:?} distance {:?} reliable {:?}",
res.pred_locs.as_ref().map(|l| &l[..4.min(l.len())]),
res.reliable_distance,
res.reliable
);
assert!(res.json_text.trim_start().starts_with('{'), "dict open");
let n_vals = ["30", "45", "25", "10"]
.iter()
.filter(|v| res.json_text.contains(**v))
.count();
eprintln!("[D6 e2e f32] text values: {n_vals}/4");
assert!(
n_vals >= 2,
"text values collapsed below the measured floor"
);
let locs = res.pred_locs.as_ref().expect("<Number> must fire");
for (i, want) in [0.5714, 1.0, 0.4286, 0.0].iter().enumerate() {
assert!(
(f64::from(locs[i]) - want).abs() < 0.1,
"pred_locs[{i}] = {} vs normalized truth {want}",
locs[i]
);
}
if res.pred_locs.is_some() && res.reliable_distance.is_some() {
assert!(res.reliable.is_some());
}
let int8_path = format!("{dir}/onechart.int8.focrq");
if std::path::Path::new(&int8_path).is_file() {
let w8 = Weights::load(std::path::Path::new(&int8_path)).expect("int8 artifact");
let statics8 = hydrate_statics(&w8, "model.vision_tower", false).expect("int8 statics");
if let Ok(r8) = recognize(&w8, &statics8, &tk, &img, 256) {
let n_vals = ["30", "45", "25", "10"]
.iter()
.filter(|v| r8.json_text.contains(**v))
.count();
eprintln!(
"[D6 e2e int8] {n_vals}/4 values, pred_locs[:4]: {:?}",
r8.pred_locs.as_ref().map(|l| &l[..4.min(l.len())])
);
}
}
}
#[test]
fn corpus_quality_scrm_proxy() {
let Ok(dir) = std::env::var("FOCR_ONECHART_DIR") else {
return;
};
let int8_path = format!("{dir}/onechart.int8.focrq");
if !std::path::Path::new(&int8_path).is_file() {
eprintln!("skip-with-SUCCESS: {int8_path} absent");
return;
}
let corpus_dir = concat!(
env!("CARGO_MANIFEST_DIR"),
"/tests/fixtures/onechart/corpus"
);
let manifest: serde_json::Value = serde_json::from_str(
&std::fs::read_to_string(format!("{corpus_dir}/manifest.json"))
.expect("corpus manifest"),
)
.expect("manifest parses");
let tk = crate::tokenizer::Tokenizer::from_opt_dir(std::path::Path::new(&dir))
.expect("onechart tokenizer");
let mut legs: Vec<(&str, Weights)> = vec![(
"int8",
Weights::load(std::path::Path::new(&int8_path)).expect("int8 artifact"),
)];
let f32_path = format!("{dir}/model.safetensors");
if std::path::Path::new(&f32_path).is_file() {
legs.push((
"f32",
Weights::load(std::path::Path::new(&f32_path)).expect("f32 weights"),
));
}
for (label, weights) in &legs {
run_corpus_leg(label, weights, &tk, corpus_dir, &manifest);
}
}
fn run_corpus_leg(
label: &str,
weights: &Weights,
tk: &crate::tokenizer::Tokenizer,
corpus_dir: &str,
manifest: &serde_json::Value,
) {
let mut n_valid_json = 0usize;
let mut head_dists = Vec::new();
let mut value_errs = Vec::new();
let statics = hydrate_statics(weights, "model.vision_tower", false).expect("statics");
let charts = manifest["charts"].as_array().unwrap();
for chart in charts {
let file = chart["file"].as_str().unwrap();
let img = image::open(format!("{corpus_dir}/{file}")).expect("chart decodes");
let res = recognize(weights, &statics, tk, &img, 512).expect("recognize");
let gt = extract_gt_values(&chart["values"]).expect("manifest GT is list-free");
let gt_norm = normalize_gt(>);
let parsed: Option<serde_json::Value> = serde_json::from_str(&res.json_text).ok();
let valid = parsed.is_some();
n_valid_json += usize::from(valid);
let ours_vals = parsed
.as_ref()
.and_then(|p| p.get("values").or_else(|| p.get("data")).cloned())
.and_then(|v| extract_gt_values(&v));
let rel_err = ours_vals.as_ref().map(|ov| {
let mut a = ov.clone();
let mut b = gt.clone();
a.sort_by(f64::total_cmp);
b.sort_by(f64::total_cmp);
let n = a.len().min(b.len()).max(1);
let pair_err: f64 = a
.iter()
.zip(&b)
.map(|(x, y)| (x - y).abs() / y.abs().max(1.0))
.sum::<f64>()
/ n as f64;
let miss = (a.len() as f64 - b.len() as f64).abs() / b.len().max(1) as f64;
pair_err + miss
});
if let Some(e) = rel_err {
value_errs.push(e);
}
let head_dist = res
.pred_locs
.as_ref()
.map(|locs| reliable_distance(locs, >_norm));
if let Some(d) = head_dist {
head_dists.push(d);
}
eprintln!(
"[bd-2lje {label}] {file}: valid_json={valid} rel_err={rel_err:?} head_dist={head_dist:?}"
);
let snip: String = res.json_text.chars().take(160).collect();
eprintln!("[bd-2lje {label}] {file}: text={snip:?}");
}
let mean = |v: &[f64]| v.iter().sum::<f64>() / v.len().max(1) as f64;
eprintln!(
"[bd-2lje {label}] SUMMARY: valid_json {n_valid_json}/{} | mean rel_err {:.3} (n={}) | \
mean head_dist {:.3} (n={})",
charts.len(),
mean(&value_errs),
value_errs.len(),
mean(&head_dists),
head_dists.len()
);
assert_eq!(charts.len(), 6, "corpus shrank");
assert_eq!(
head_dists.len(),
6,
"{label}: the number head must fire on every chart"
);
assert!(
mean(&head_dists) < 0.05,
"{label}: mean head distance {} regressed past the 0.05 gate (measured 0.015)",
mean(&head_dists)
);
assert!(n_valid_json >= 1, "{label}: valid-JSON collapsed");
}
#[test]
fn reliable_check_matches_upstream_goldens() {
let case = |json: &str| -> Option<Vec<f64>> {
extract_gt_values(&serde_json::from_str(json).unwrap()).map(|v| normalize_gt(&v))
};
assert_eq!(
case(r#"{"A":"30","B":"45","C":"25","D":"10"}"#).unwrap(),
vec![0.5714, 1.0, 0.4286, 0.0]
);
assert_eq!(
case(r#"{"x":"6.12%","y":"1,234"}"#).unwrap(),
vec![0.0, 1.0]
);
assert_eq!(
case(r#"{"s1":{"a":1,"b":3},"s2":{"c":5}}"#).unwrap(),
vec![0.0, 0.5, 1.0]
);
assert_eq!(case(r#"{"a":"none","b":"5","c":"-"}"#).unwrap(), vec![5.0]);
assert_eq!(case(r#"{"t":"ab(3)cd [7] 12.5x"}"#).unwrap(), vec![12.5]);
assert_eq!(case(r#"{"a":[1,2]}"#), None);
let gt = vec![0.5714, 1.0, 0.4286, 0.0];
let pred: Vec<f32> = vec![0.57, 1.0, 0.43, 0.0];
let d = reliable_distance(&pred, >);
assert!(d < RELIABLE_THRESHOLD, "near-exact pred must verify ({d})");
assert!(reliable_distance(&[0.9, 0.1], &[0.0, 1.0]) > RELIABLE_THRESHOLD);
assert!(reliable_distance(&pred, &[]).is_infinite());
}
#[test]
fn number_head_matches_golden() {
let Ok(dir) = std::env::var("FOCR_ONECHART_DIR") else {
return;
};
let model_path = format!("{dir}/model.safetensors");
if !std::path::Path::new(&model_path).is_file() {
eprintln!("skip-with-SUCCESS: {model_path} absent");
return;
}
let fx: serde_json::Value = serde_json::from_str(
&std::fs::read_to_string(concat!(
env!("CARGO_MANIFEST_DIR"),
"/tests/fixtures/onechart/num_decoder_golden.json"
))
.expect("golden read"),
)
.expect("golden parse");
let hidden: Vec<f32> = fx["input_hidden"]
.as_array()
.unwrap()
.iter()
.map(|v| v.as_f64().unwrap() as f32)
.collect();
let want: Vec<f32> = fx["pred_locs_256"]
.as_array()
.unwrap()
.iter()
.map(|v| v.as_f64().unwrap() as f32)
.collect();
let weights = Weights::load(std::path::Path::new(&model_path)).expect("weights");
let ours = number_head(&weights, &hidden).expect("number head");
assert_eq!(ours.len(), 256);
let mut max_abs = 0.0f64;
let mut dot = 0.0f64;
let (mut na, mut nb) = (0.0f64, 0.0f64);
for (a, b) in ours.iter().zip(&want) {
let (a, b) = (f64::from(*a), f64::from(*b));
max_abs = max_abs.max((a - b).abs());
dot += a * b;
na += a * a;
nb += b * b;
}
let cos = dot / (na.sqrt() * nb.sqrt());
eprintln!("[D5 parity] num_decoder cos={cos:.8} maxabs={max_abs:.3e}");
assert!(cos >= 0.9999, "num_decoder cosine {cos:.8}");
assert!(max_abs <= 1e-4, "num_decoder maxabs {max_abs:.3e}");
}
#[test]
fn opt_kvcache_matches_greedy_and_oracle() {
let Ok(dir) = std::env::var("FOCR_ONECHART_DIR") else {
return;
};
let proj_path = format!("{dir}/onechart_proj_out.bin");
let model_path = format!("{dir}/model.safetensors");
if !std::path::Path::new(&proj_path).is_file() {
eprintln!("skip-with-SUCCESS: {proj_path} absent (run the oracle script)");
return;
}
let read_f32 = |p: &str| -> Vec<f32> {
std::fs::read(p)
.expect("oracle blob reads")
.as_chunks::<4>()
.0
.iter()
.map(|c| f32::from_le_bytes(*c))
.collect()
};
let fx: serde_json::Value = serde_json::from_str(
&std::fs::read_to_string(concat!(
env!("CARGO_MANIFEST_DIR"),
"/tests/fixtures/onechart/oracle_fixtures.json"
))
.expect("oracle fixtures read"),
)
.expect("oracle fixtures parse");
let prompt_ids: Vec<u32> = fx["l0c_prompt"]["ids"]
.as_array()
.unwrap()
.iter()
.map(|v| v.as_u64().unwrap() as u32)
.collect();
let int8_path = format!("{dir}/onechart.int8.focrq");
let weights = if std::path::Path::new(&int8_path).is_file() {
Weights::load(std::path::Path::new(&int8_path)).expect("int8 artifact")
} else {
Weights::load(std::path::Path::new(&model_path)).expect("weights")
};
let vision = Mat::from_vec(VISION_TOKENS, HIDDEN, read_f32(&proj_path));
let statics = hydrate_statics(&weights, "model.vision_tower", false).expect("statics");
let embeds = build_inputs_embeds(&statics, &vision, &prompt_ids).expect("splice");
let cfg = super::super::decoder_qwen2::DecoderConfig::onechart();
let ids_kv = super::super::decoder_qwen2::generate_greedy_kvcache(
&weights,
&cfg,
&embeds,
24,
crate::tokenizer::special_opt::BOS_EOS,
)
.expect("kvcache greedy");
let ids_greedy = super::super::decoder_qwen2::generate_greedy(
&weights,
&cfg,
&embeds,
24,
crate::tokenizer::special_opt::BOS_EOS,
)
.expect("re-prefill greedy");
eprintln!("[D4 decode] kvcache: {ids_kv:?}");
let b9_prefix = ids_kv
.iter()
.zip(&ids_greedy)
.take_while(|(a, b)| a == b)
.count();
eprintln!("[D4 decode] kvcache-vs-greedy exact prefix: {b9_prefix}/24");
assert!(
b9_prefix >= 12,
"kvcache vs re-prefill diverged at step {b9_prefix} — earlier than the \
measured near-tie horizon (13); a structural decode-path defect"
);
assert_eq!(
ids_kv[0],
crate::tokenizer::special_opt::NUMBER,
"first generated id must be the <Number> trigger (census §8)"
);
let tk = crate::tokenizer::Tokenizer::from_opt_dir(std::path::Path::new(&dir))
.expect("onechart tokenizer");
let ours = tk.decode_skip_special(&ids_kv).expect("decode");
let oracle: String = fx["l4_chat"]["answer"]
.as_str()
.unwrap()
.chars()
.take(60)
.collect();
eprintln!("[D4 decode] ours: {:?}", ours.trim());
eprintln!("[D4 decode] oracle: {oracle:?}");
assert!(
ours.trim_start().starts_with('{'),
"decoded output does not open the chart dict: {ours:?}"
);
}
#[test]
fn vision_features_match_torch_oracle() {
let Ok(dir) = std::env::var("FOCR_ONECHART_DIR") else {
return;
};
let pre_path = format!("{dir}/onechart_preproc.bin");
let want_path = format!("{dir}/onechart_proj_out.bin");
let model_path = format!("{dir}/model.safetensors");
if !std::path::Path::new(&pre_path).is_file() {
eprintln!("skip-with-SUCCESS: {pre_path} absent (run the oracle script)");
return;
}
let read_f32 = |p: &str| -> Vec<f32> {
std::fs::read(p)
.expect("oracle blob reads")
.as_chunks::<4>()
.0
.iter()
.map(|c| f32::from_le_bytes(*c))
.collect()
};
let pre = read_f32(&pre_path);
assert_eq!(pre.len(), 3 * 1024 * 1024, "preproc not [3,1024,1024]");
let want = read_f32(&want_path);
assert_eq!(want.len(), VISION_TOKENS * HIDDEN, "proj_out not [256,768]");
let weights = Weights::load(std::path::Path::new(&model_path)).expect("weights");
let statics = hydrate_statics(&weights, "model.vision_tower", false).expect("statics");
let image = Mat::from_vec(3, 1024 * 1024, pre);
let ours = vision_features(&weights, &statics, &image).expect("vision");
assert_eq!((ours.rows, ours.cols), (VISION_TOKENS, HIDDEN));
let mut dot = 0.0f64;
let (mut na, mut nb) = (0.0f64, 0.0f64);
let mut max_abs = 0.0f64;
for (a, b) in ours.data.iter().zip(&want) {
let (a, b) = (f64::from(*a), f64::from(*b));
dot += a * b;
na += a * a;
nb += b * b;
max_abs = max_abs.max((a - b).abs());
}
let cos = dot / (na.sqrt() * nb.sqrt());
eprintln!("[D3 parity] proj_out cos={cos:.8} maxabs={max_abs:.3e}");
assert!(cos >= 0.9999, "OneChart vision cosine {cos:.8} < 0.9999");
assert!(
max_abs <= 1e-2,
"OneChart proj_out maxabs {max_abs:.3e} > 1e-2 — investigate before tightening"
);
}
}