| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| use crate::{coverage, coverage_checksums, energy, gradient, model::Vectors}; |
| use std::cell::RefCell; |
| use std::collections::HashMap; |
|
|
| thread_local! { |
| |
| |
| static STATE: RefCell<Option<(Vectors, Vec<f64>)>> = const { RefCell::new(None) }; |
| } |
|
|
| |
| |
| pub fn compute_json(v: &Vectors, target: &[f64]) -> String { |
| let cov = coverage::region_coverages(v); |
| let img = energy::compose(v, &cov); |
| let e_data = energy::e_data(&img, target, v.l0); |
| let grad = gradient::vertex_gradients(v, &img, target); |
|
|
| let cov_by_label: HashMap<i64, &Vec<f64>> = cov.iter().map(|(l, c)| (*l, c)).collect(); |
| let coverage_samples: Vec<f64> = v |
| .golden |
| .coverage_samples |
| .iter() |
| .map(|s| { |
| cov_by_label |
| .get(&s.label) |
| .map(|c| c[s.row * v.width + s.col]) |
| .unwrap_or(f64::NAN) |
| }) |
| .collect(); |
| let vertex_gradients: Vec<[f64; 2]> = v |
| .golden |
| .vertex_gradients |
| .iter() |
| .map(|s| { |
| grad.get(&s.edge) |
| .map(|pc| pc[s.cubic][s.vertex]) |
| .unwrap_or([f64::NAN, f64::NAN]) |
| }) |
| .collect(); |
|
|
| let checksums = coverage_checksums(&cov); |
| serde_json::json!({ |
| "name": v.name, |
| "engine": "wasm", |
| "e_data": e_data, |
| "coverage_checksums": checksums.iter() |
| .map(|(l, c)| (l.to_string(), serde_json::json!({"sum": c.sum, "sumsq": c.sumsq}))) |
| .collect::<serde_json::Map<_, _>>(), |
| "coverage_samples": coverage_samples, |
| "vertex_gradients": vertex_gradients, |
| }) |
| .to_string() |
| } |
|
|
| |
| |
| #[cfg_attr(not(any(target_arch = "wasm32", test)), allow(dead_code))] |
| fn store(text: &str) -> i32 { |
| match serde_json::from_str::<Vectors>(text) { |
| Ok(v) => { |
| let target = v.target(); |
| STATE.with(|s| *s.borrow_mut() = Some((v, target))); |
| 0 |
| } |
| Err(_) => -1, |
| } |
| } |
|
|
| #[cfg_attr(not(any(target_arch = "wasm32", test)), allow(dead_code))] |
| fn compute_stored() -> String { |
| STATE.with(|s| { |
| let b = s.borrow(); |
| let (v, target) = b.as_ref().expect("compute() called before load()"); |
| compute_json(v, target) |
| }) |
| } |
|
|
| |
| #[cfg_attr(not(any(target_arch = "wasm32", test)), allow(dead_code))] |
| fn leak_length_prefixed(payload: String) -> *mut u8 { |
| let bytes = payload.into_bytes(); |
| let mut buf = Vec::<u8>::with_capacity(4 + bytes.len()); |
| buf.extend_from_slice(&(bytes.len() as u32).to_le_bytes()); |
| buf.extend_from_slice(&bytes); |
| let p = buf.as_mut_ptr(); |
| std::mem::forget(buf); |
| p |
| } |
|
|
| |
| |
|
|
| #[cfg(target_arch = "wasm32")] |
| mod abi { |
| use super::*; |
|
|
| |
| |
| |
| |
| #[no_mangle] |
| pub extern "C" fn alloc(len: usize) -> *mut u8 { |
| let mut buf = Vec::<u8>::with_capacity(len); |
| let ptr = buf.as_mut_ptr(); |
| std::mem::forget(buf); |
| ptr |
| } |
|
|
| |
| |
| |
| |
| #[no_mangle] |
| pub unsafe extern "C" fn dealloc(ptr: *mut u8, len: usize) { |
| if !ptr.is_null() { |
| drop(Vec::from_raw_parts(ptr, 0, len)); |
| } |
| } |
|
|
| |
| |
| |
| |
| #[no_mangle] |
| pub unsafe extern "C" fn load(ptr: *const u8, len: usize) -> i32 { |
| let bytes = std::slice::from_raw_parts(ptr, len); |
| match std::str::from_utf8(bytes) { |
| Ok(text) => store(text), |
| Err(_) => -1, |
| } |
| } |
|
|
| |
| #[no_mangle] |
| pub extern "C" fn compute() -> *mut u8 { |
| leak_length_prefixed(compute_stored()) |
| } |
| } |
|
|
| #[cfg(test)] |
| mod tests { |
| use super::*; |
|
|
| |
| |
| #[test] |
| fn compute_json_round_trips_a_minimal_document() { |
| let doc = serde_json::json!({ |
| "name": "unit", "kind": "unit", "seed": 1, |
| "width": 4, "height": 4, "background": 1.0, "l0": 1.0, |
| "colors255": {"0": [255.0, 0.0, 0.0]}, |
| "edges": [], |
| "faces": [], |
| "target_u8": vec![0u8; 4 * 4 * 3], |
| "golden": { |
| "e_data": 0.0, |
| "coverage_checksums": {}, |
| "coverage_samples": [], |
| "vertex_gradients": [] |
| } |
| }) |
| .to_string(); |
| assert_eq!(store(&doc), 0, "minimal document should parse"); |
| let out = compute_stored(); |
| let parsed: serde_json::Value = serde_json::from_str(&out).unwrap(); |
| |
| assert!((parsed["e_data"].as_f64().unwrap() - 48.0).abs() < 1e-12); |
| } |
| } |
|
|