#![cfg(not(target_arch = "wasm32"))]
use fxtranslate::engine::Engine;
const MODEL: &str = "../../../data/models/enfr/model.enfr.intgemm.alphas.bin";
const VOCAB: &str = "../../../data/models/enfr/vocab.enfr.spm";
fn engine() -> Option<Engine> {
if !std::path::Path::new(MODEL).exists() {
eprintln!("skipping batched-decode parity: {MODEL} absent");
return None;
}
Some(Engine::load(MODEL, VOCAB, VOCAB).expect("engine loads"))
}
#[test]
fn greedy_batch_matches_single_per_sentence() {
let Some(eng) = engine() else { return };
let texts = [
"The cat sat on the mat.",
"Dogs run.",
"Scientists carefully explained the experiment to the students.",
"Birds fly south.",
];
let ids: Vec<Vec<u32>> = texts.iter().map(|t| eng.src_ids(t)).collect();
let batched = eng.greedy_batch(&ids);
for (b, sid) in ids.iter().enumerate() {
let single = eng.greedy(sid);
eprintln!(
"sentence {b}: single {} toks, batched {} toks",
single.len(),
batched[b].len()
);
assert_eq!(
batched[b], single,
"sentence {b}: batched greedy diverges from single-sentence greedy"
);
}
}
#[test]
fn translate_batch_matches_single() {
let Some(eng) = engine() else { return };
let texts = [
"Hello world.",
"The quick brown fox jumps.",
"Good morning.",
];
let batched = eng.translate_batch(&texts);
for (b, t) in texts.iter().enumerate() {
assert_eq!(
batched[b],
eng.translate(t),
"sentence {b}: translate_batch diverges from translate"
);
}
}
#[cfg(feature = "threads")]
#[test]
fn greedy_batch_parallel_matches_serial() {
let Some(eng) = engine() else { return };
let texts = [
"The cat sat on the mat.",
"Dogs run.",
"Scientists carefully explained the experiment to the students.",
"Birds fly south.",
"Hello world.",
"The quick brown fox jumps over the lazy dog.",
"Good morning.",
"She sells seashells by the seashore.",
"Winter is coming soon.",
"They travelled across the country by train.",
];
let ids: Vec<Vec<u32>> = texts.iter().map(|t| eng.src_ids(t)).collect();
let serial = eng.greedy_batch(&ids); let eng = eng.with_threads(8);
let parallel = eng.greedy_batch(&ids);
assert_eq!(
parallel, serial,
"data-parallel greedy_batch diverges from the serial batch"
);
}