use ferrox_core::par;
pub trait Engine: Sync {
type State: Send;
fn new_state(&self) -> Self::State;
fn vocab_size(&self) -> usize;
fn forward_token_on_worker(
&self,
token_id: usize,
pos: usize,
state: &mut Self::State,
) -> Vec<f32>;
fn forward_token(&self, token_id: usize, pos: usize, state: &mut Self::State) -> Vec<f32> {
par::on_workers(move || self.forward_token_on_worker(token_id, pos, state))
}
}
#[cfg(test)]
pub(crate) fn on_workers_promotes_here() -> bool {
let before = par::cold_regions();
par::on_workers(|| {});
par::cold_regions() - before == 1
}
#[cfg(test)]
pub(crate) fn assert_one_pool_entry_per_step<E: Engine>(engine: &E, token_id: usize) {
if !on_workers_promotes_here() {
return;
}
let mut state = engine.new_state();
let _ = engine.forward_token_on_worker(token_id, 0, &mut state);
let before = par::cold_regions();
let _ = engine.forward_token_on_worker(token_id, 1, &mut state);
let raw = par::cold_regions() - before;
let before = par::cold_regions();
let _ = Engine::forward_token(engine, token_id, 2, &mut state);
let promoted = par::cold_regions() - before;
assert_eq!(
promoted, 1,
"a decode step must enter the pool once, not once per matvec"
);
assert!(
raw > promoted,
"the unpromoted body must open a region per parallel section \
({raw} vs {promoted}); equal counts mean this engine's step is \
too small to be measuring the promotion at all"
);
}
#[cfg(test)]
mod tests {
use super::{assert_one_pool_entry_per_step, on_workers_promotes_here, Engine};
use ferrox_core::par;
use std::path::{Path, PathBuf};
fn crate_sources() -> Vec<PathBuf> {
fn walk(dir: &Path, out: &mut Vec<PathBuf>) {
for entry in std::fs::read_dir(dir).expect("src/ is readable") {
let path = entry.expect("dir entry").path();
if path.is_dir() {
walk(&path, out);
} else if path.extension().is_some_and(|e| e == "rs") {
out.push(path);
}
}
}
let src = Path::new(env!("CARGO_MANIFEST_DIR")).join("src");
let mut out = Vec::new();
walk(&src, &mut out);
assert!(!out.is_empty(), "walked src/ and found no Rust at all");
out
}
#[test]
fn no_engine_may_override_the_promoted_forward_token() {
let decoder_entry = Path::new("decoder").join("entry.rs");
let engine_entry = Path::new("engine").join("entry.rs");
let mut stray: Vec<String> = Vec::new();
for path in crate_sources() {
if path.ends_with(&decoder_entry) || path.ends_with(&engine_entry) {
continue;
}
let body = std::fs::read_to_string(&path).expect("source reads");
for (i, line) in body.lines().enumerate() {
let t = line.trim_start();
if t.starts_with("fn forward_token(") || t.starts_with("pub fn forward_token(") {
stray.push(format!("{}:{}", path.display(), i + 1));
}
}
}
assert!(
stray.is_empty(),
"these declarations bypass the pool wrapper on Engine::forward_token: {stray:?}"
);
}
struct TwelveRegions;
impl Engine for TwelveRegions {
type State = Vec<f32>;
fn new_state(&self) -> Vec<f32> {
vec![0.0; 4096]
}
fn vocab_size(&self) -> usize {
1
}
fn forward_token_on_worker(
&self,
_token_id: usize,
_pos: usize,
state: &mut Vec<f32>,
) -> Vec<f32> {
for _ in 0..4 {
for _ in 0..3 {
par::items_mut(state, 1, |_, x| *x += 1.0);
}
}
vec![0.0]
}
}
#[test]
fn a_step_through_the_trait_enters_the_pool_once_and_the_raw_body_does_not() {
assert_one_pool_entry_per_step(&TwelveRegions, 0);
}
#[test]
fn the_shared_assertion_rejects_a_step_too_small_to_measure() {
struct OneRegion;
impl Engine for OneRegion {
type State = Vec<f32>;
fn new_state(&self) -> Vec<f32> {
vec![0.0; 64]
}
fn vocab_size(&self) -> usize {
1
}
fn forward_token_on_worker(
&self,
_token_id: usize,
_pos: usize,
state: &mut Vec<f32>,
) -> Vec<f32> {
par::items_mut(state, 1, |_, x| *x += 1.0);
vec![0.0]
}
}
if !on_workers_promotes_here() {
return;
}
let caught = std::panic::catch_unwind(|| assert_one_pool_entry_per_step(&OneRegion, 0));
assert!(
caught.is_err(),
"a one-region step must be refused as unmeasurable, not accepted as promoted"
);
}
}