use std::collections::VecDeque;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use super::decoder::{self, DecoderWeightCacheI8};
use super::rswa::BatchedRingCache;
use super::sampler::{self, DecodeParams};
use super::tensor::Mat;
use crate::error::FocrResult;
pub const DEFAULT_BATCH_SIZE: usize = 128;
pub const MAX_BATCH_SIZE: usize = 256;
#[must_use]
pub fn parse_batch_size(val: Option<&str>) -> usize {
match val.and_then(|s| s.trim().parse::<usize>().ok()) {
Some(0) | None => DEFAULT_BATCH_SIZE,
Some(n) => n.min(MAX_BATCH_SIZE),
}
}
#[must_use]
pub fn scheduler_batch_size() -> usize {
parse_batch_size(std::env::var("FOCR_BATCH_SIZE").ok().as_deref())
}
#[must_use]
pub fn spine_enabled() -> bool {
decoder::batch_spine_enabled()
}
const BATCH_PACK_ENV: &str = "FOCR_BATCH_PACK";
#[must_use]
pub fn batch_pack_enabled() -> bool {
static FLAG: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*FLAG.get_or_init(|| std::env::var_os(BATCH_PACK_ENV).is_some())
}
#[must_use]
fn admission_order(streams: &[PageStream], pack: bool) -> Vec<usize> {
let mut order: Vec<usize> = (0..streams.len()).collect();
if pack {
order.sort_by_key(|&i| streams[i].prefill_len);
}
order
}
#[derive(Debug, Clone)]
pub struct PageStream {
pub input_index: usize,
pub prefill_len: usize,
pub position: usize,
pub generated: Vec<u32>,
pub prompt_len: usize,
pub last_hidden: Mat,
pub done: bool,
pub eos: bool,
pub max_emit: Option<usize>,
}
impl PageStream {
#[must_use]
pub fn new(
input_index: usize,
prefill_len: usize,
prompt_ids: &[u32],
last_hidden: Mat,
) -> Self {
Self {
input_index,
prefill_len,
position: prefill_len,
generated: prompt_ids.to_vec(),
prompt_len: prompt_ids.len(),
last_hidden,
done: false,
eos: false,
max_emit: None,
}
}
#[must_use]
pub fn with_max_emit(mut self, cap: usize) -> Self {
self.max_emit = Some(cap);
self
}
#[must_use]
pub fn emitted(&self) -> &[u32] {
&self.generated[self.prompt_len..]
}
}
pub struct StreamSlot<'a> {
pub slot_index: usize,
pub history: &'a [u32],
pub hidden: &'a Mat,
pub position: usize,
}
pub struct StreamOut {
pub token: u32,
pub is_eos: bool,
pub new_hidden: Mat,
}
pub trait BatchStep {
fn step(&mut self, slots: &[StreamSlot<'_>]) -> FocrResult<Vec<StreamOut>>;
}
#[derive(Debug, Clone, Copy)]
pub struct SchedulerStats {
pub max_concurrent_forwards: usize,
pub guard_held_during_fanout: bool,
pub max_active: usize,
pub total_steps: usize,
}
pub struct BatchScheduler {
batch_size: usize,
max_length: usize,
live_forwards: AtomicUsize,
max_forwards: AtomicUsize,
max_active: AtomicUsize,
guard_violation: AtomicBool,
steps: usize,
}
impl BatchScheduler {
#[must_use]
pub fn new(batch_size: usize, max_length: usize) -> Self {
Self {
batch_size: batch_size.clamp(1, MAX_BATCH_SIZE),
max_length,
live_forwards: AtomicUsize::new(0),
max_forwards: AtomicUsize::new(0),
max_active: AtomicUsize::new(0),
guard_violation: AtomicBool::new(false),
steps: 0,
}
}
#[must_use]
pub fn from_env(max_length: usize) -> Self {
Self::new(scheduler_batch_size(), max_length)
}
#[must_use]
pub fn batch_size(&self) -> usize {
self.batch_size
}
pub fn note_guard_held_during_fanout(&self) {
self.guard_violation.store(true, Ordering::SeqCst);
}
#[must_use]
pub fn stats(&self) -> SchedulerStats {
SchedulerStats {
max_concurrent_forwards: self.max_forwards.load(Ordering::SeqCst),
guard_held_during_fanout: self.guard_violation.load(Ordering::SeqCst),
max_active: self.max_active.load(Ordering::SeqCst),
total_steps: self.steps,
}
}
pub fn run<S: BatchStep>(
&mut self,
streams: Vec<PageStream>,
step: &mut S,
) -> FocrResult<Vec<Vec<u32>>> {
self.run_with_pack(streams, step, batch_pack_enabled())
}
fn run_with_pack<S: BatchStep>(
&mut self,
mut streams: Vec<PageStream>,
step: &mut S,
pack: bool,
) -> FocrResult<Vec<Vec<u32>>> {
let mut pending: VecDeque<usize> = admission_order(&streams, pack).into();
let mut active: Vec<usize> = Vec::with_capacity(self.batch_size);
Self::admit(&mut active, &mut pending, self.batch_size);
while !active.is_empty() {
self.max_active.fetch_max(active.len(), Ordering::SeqCst);
let slots: Vec<StreamSlot<'_>> = active
.iter()
.map(|&i| StreamSlot {
slot_index: i,
history: &streams[i].generated,
hidden: &streams[i].last_hidden,
position: streams[i].position,
})
.collect();
let fwd = self.enter_forward();
let result = step.step(&slots);
drop(fwd);
self.exit_forward();
let outs = result?;
self.steps += 1;
debug_assert_eq!(outs.len(), active.len(), "one StreamOut per active stream");
let mut retire: Vec<usize> = Vec::new();
for (k, (&i, out)) in active.iter().zip(outs).enumerate() {
let s = &mut streams[i];
s.generated.push(out.token);
s.position += 1;
let cap = s
.max_emit
.map_or(self.max_length, |m| m.min(self.max_length));
if out.is_eos || s.emitted().len() >= cap {
s.done = true;
s.eos = out.is_eos;
retire.push(k);
} else {
s.last_hidden = out.new_hidden;
}
}
for &k in retire.iter().rev() {
active.remove(k);
}
Self::admit(&mut active, &mut pending, self.batch_size);
}
let mut out: Vec<(usize, Vec<u32>)> = streams
.into_iter()
.map(|s| (s.input_index, s.emitted().to_vec()))
.collect();
out.sort_by_key(|(idx, _)| *idx);
Ok(out.into_iter().map(|(_, toks)| toks).collect())
}
fn admit(active: &mut Vec<usize>, pending: &mut VecDeque<usize>, cap: usize) {
while active.len() < cap {
match pending.pop_front() {
Some(i) => active.push(i),
None => break,
}
}
}
fn enter_forward(&self) -> super::ForwardPass {
let n = self.live_forwards.fetch_add(1, Ordering::SeqCst) + 1;
self.max_forwards.fetch_max(n, Ordering::SeqCst);
super::enter_forward()
}
fn exit_forward(&self) {
self.live_forwards.fetch_sub(1, Ordering::SeqCst);
}
}
pub struct DecoderBatchStep<'a> {
pub wc: &'a DecoderWeightCacheI8,
pub caches: &'a mut BatchedRingCache,
pub embed_table: &'a super::EmbedTable<'a>,
pub params: &'a DecodeParams,
}
impl BatchStep for DecoderBatchStep<'_> {
fn step(&mut self, slots: &[StreamSlot<'_>]) -> FocrResult<Vec<StreamOut>> {
let hidden_dim = self.embed_table.cols();
let vocab = self.embed_table.rows();
let hiddens: Vec<Mat> = slots.iter().map(|s| s.hidden.clone()).collect();
let logits_rows = decoder::batched_lm_head_i8(self.wc, &hiddens)?;
let b = slots.len();
let mut stacked = Vec::with_capacity(b * vocab);
for row in &logits_rows {
stacked.extend_from_slice(&row.data);
}
let logits = Mat::from_vec(b, vocab, stacked);
let histories: Vec<&[u32]> = slots.iter().map(|s| s.history).collect();
let decoded = sampler::batched_decode_step(&logits, &histories, self.params)?;
let mut token_embeds: Vec<Mat> = Vec::with_capacity(b);
let mut positions: Vec<usize> = Vec::with_capacity(b);
let mut stream_ids: Vec<usize> = Vec::with_capacity(b);
for (out, slot) in decoded.iter().zip(slots.iter()) {
let t = out.token_id as usize;
if t >= vocab {
return Err(crate::error::FocrError::Other(anyhow::anyhow!(
"batch_scheduler::DecoderBatchStep: token id {t} outside embed vocab {vocab}"
)));
}
let row = self.embed_table.row_f32(t);
token_embeds.push(Mat::from_vec(1, hidden_dim, row));
positions.push(slot.position);
stream_ids.push(slot.slot_index);
}
let new_hiddens = decoder::batched_decode_step_i8_streams(
self.wc,
self.caches,
&stream_ids,
&token_embeds,
&positions,
)?;
Ok(decoded
.into_iter()
.zip(new_hiddens)
.map(|(out, new_hidden)| StreamOut {
token: out.token_id,
is_eos: out.is_eos,
new_hidden,
})
.collect())
}
}
#[cfg(test)]
mod tests {
use super::*;
struct MockStep {
batch_sizes: Vec<usize>,
}
impl MockStep {
fn new() -> Self {
Self {
batch_sizes: Vec::new(),
}
}
}
impl BatchStep for MockStep {
fn step(&mut self, slots: &[StreamSlot<'_>]) -> FocrResult<Vec<StreamOut>> {
self.batch_sizes.push(slots.len());
Ok(slots
.iter()
.map(|s| {
let tag = s.hidden.data[0];
let eos_after = s.hidden.data[1] as usize;
let emitted_before = s.history.len(); let token_id = emitted_before as u32;
let is_eos = eos_after != 0 && emitted_before + 1 >= eos_after;
StreamOut {
token: token_id,
is_eos,
new_hidden: Mat::from_vec(1, 2, vec![tag, eos_after as f32]),
}
})
.collect())
}
}
fn stream(input_index: usize, eos_after: usize) -> PageStream {
let hidden = Mat::from_vec(1, 2, vec![input_index as f32, eos_after as f32]);
PageStream::new(
input_index,
4,
&[],
hidden,
)
}
#[test]
fn parse_batch_size_clamps_and_defaults() {
assert_eq!(parse_batch_size(None), DEFAULT_BATCH_SIZE);
assert_eq!(parse_batch_size(Some("")), DEFAULT_BATCH_SIZE);
assert_eq!(parse_batch_size(Some("garbage")), DEFAULT_BATCH_SIZE);
assert_eq!(parse_batch_size(Some("0")), DEFAULT_BATCH_SIZE);
assert_eq!(parse_batch_size(Some("1")), 1);
assert_eq!(parse_batch_size(Some("64")), 64);
assert_eq!(parse_batch_size(Some("256")), MAX_BATCH_SIZE);
assert_eq!(parse_batch_size(Some("100000")), MAX_BATCH_SIZE);
}
#[test]
fn new_clamps_batch_size() {
assert_eq!(BatchScheduler::new(0, 16).batch_size(), 1);
assert_eq!(BatchScheduler::new(9999, 16).batch_size(), MAX_BATCH_SIZE);
assert_eq!(BatchScheduler::new(8, 16).batch_size(), 8);
}
#[test]
fn each_stream_emits_exactly_eos_after_tokens() {
let mut sched = BatchScheduler::new(4, 100);
let streams = vec![stream(0, 3), stream(1, 1), stream(2, 5)];
let mut step = MockStep::new();
let out = sched.run(streams, &mut step).expect("run");
assert_eq!(out[0], vec![0, 1, 2]);
assert_eq!(out[1], vec![0]);
assert_eq!(out[2], vec![0, 1, 2, 3, 4]);
}
#[test]
fn outputs_resorted_to_input_order_despite_shuffled_submission() {
let mut sched = BatchScheduler::new(8, 100);
let streams = vec![stream(2, 2), stream(0, 4), stream(1, 1)];
let mut step = MockStep::new();
let out = sched.run(streams, &mut step).expect("run");
assert_eq!(out[0], vec![0, 1, 2, 3]); assert_eq!(out[1], vec![0]); assert_eq!(out[2], vec![0, 1]); }
#[test]
fn one_live_forward_and_guard_not_held() {
let mut sched = BatchScheduler::new(4, 100);
let streams = vec![stream(0, 3), stream(1, 4), stream(2, 2)];
let mut step = MockStep::new();
let _ = sched.run(streams, &mut step).expect("run");
let st = sched.stats();
assert_eq!(
st.max_concurrent_forwards, 1,
"exactly one live forward per step"
);
assert!(
!st.guard_held_during_fanout,
"model-cache guard never held during fan-out"
);
}
#[test]
fn active_set_never_exceeds_batch_size_and_backfills() {
let mut sched = BatchScheduler::new(2, 100);
let streams = vec![
stream(0, 1),
stream(1, 3),
stream(2, 2),
stream(3, 1),
stream(4, 4),
];
let mut step = MockStep::new();
let out = sched.run(streams, &mut step).expect("run");
assert_eq!(out.len(), 5);
assert_eq!(out[0].len(), 1);
assert_eq!(out[1].len(), 3);
assert_eq!(out[2].len(), 2);
assert_eq!(out[3].len(), 1);
assert_eq!(out[4].len(), 4);
let st = sched.stats();
assert!(
st.max_active <= 2,
"active set bounded by batch_size (got {})",
st.max_active
);
assert_eq!(st.max_concurrent_forwards, 1);
assert!(step.batch_sizes.iter().all(|&n| n <= 2));
}
#[test]
fn batch_size_one_serializes_all_streams_correctly() {
let mut sched = BatchScheduler::new(1, 100);
let streams = vec![stream(0, 2), stream(1, 3)];
let mut step = MockStep::new();
let out = sched.run(streams, &mut step).expect("run");
assert_eq!(out[0], vec![0, 1]);
assert_eq!(out[1], vec![0, 1, 2]);
let st = sched.stats();
assert!(st.max_active <= 1);
assert!(step.batch_sizes.iter().all(|&n| n == 1));
}
#[test]
fn max_length_retires_unbounded_streams() {
let mut sched = BatchScheduler::new(4, 5);
let streams = vec![stream(0, 0), stream(1, 2)];
let mut step = MockStep::new();
let out = sched.run(streams, &mut step).expect("run");
assert_eq!(out[0].len(), 5, "unbounded stream capped at max_length");
assert_eq!(out[0], vec![0, 1, 2, 3, 4]);
assert_eq!(out[1], vec![0, 1]);
}
#[test]
fn empty_input_yields_empty_output() {
let mut sched = BatchScheduler::new(4, 100);
let mut step = MockStep::new();
let out = sched.run(Vec::new(), &mut step).expect("run");
assert!(out.is_empty());
assert_eq!(sched.stats().total_steps, 0);
}
#[test]
fn note_guard_held_is_observable() {
let sched = BatchScheduler::new(2, 10);
assert!(!sched.stats().guard_held_during_fanout);
sched.note_guard_held_during_fanout();
assert!(sched.stats().guard_held_during_fanout);
}
struct SlotTraceStep {
traces: Vec<Vec<usize>>,
}
impl BatchStep for SlotTraceStep {
fn step(&mut self, slots: &[StreamSlot<'_>]) -> FocrResult<Vec<StreamOut>> {
self.traces
.push(slots.iter().map(|s| s.slot_index).collect());
Ok(slots
.iter()
.map(|s| {
let eos_after = s.hidden.data[1] as usize;
let emitted_before = s.history.len();
let is_eos = eos_after != 0 && emitted_before + 1 >= eos_after;
StreamOut {
token: emitted_before as u32,
is_eos,
new_hidden: s.hidden.clone(),
}
})
.collect())
}
}
#[test]
fn slot_index_identifies_cache_stream_through_retire_and_backfill() {
let mut sched = BatchScheduler::new(2, 100);
let streams = vec![stream(0, 1), stream(1, 2), stream(2, 1), stream(3, 2)];
let mut step = SlotTraceStep { traces: Vec::new() };
let out = sched.run(streams, &mut step).expect("run");
assert_eq!(out[0].len(), 1);
assert_eq!(out[1].len(), 2);
assert_eq!(out[2].len(), 1);
assert_eq!(out[3].len(), 2);
assert_eq!(
step.traces,
vec![vec![0usize, 1], vec![1, 2], vec![3], vec![3]],
"slot_index must track stable stream identity, not active position"
);
}
fn stream_len(input_index: usize, prefill_len: usize, eos_after: usize) -> PageStream {
let hidden = Mat::from_vec(1, 2, vec![input_index as f32, eos_after as f32]);
PageStream::new(input_index, prefill_len, &[], hidden)
}
fn mean_window_len_variance(streams: &[PageStream], order: &[usize], width: usize) -> f64 {
let mut total = 0.0_f64;
let mut windows = 0_usize;
for chunk in order.chunks(width) {
let n = chunk.len() as f64;
let mean = chunk
.iter()
.map(|&i| streams[i].prefill_len as f64)
.sum::<f64>()
/ n;
let var = chunk
.iter()
.map(|&i| {
let d = streams[i].prefill_len as f64 - mean;
d * d
})
.sum::<f64>()
/ n;
total += var;
windows += 1;
}
if windows == 0 {
0.0
} else {
total / windows as f64
}
}
#[test]
fn packed_admission_order_is_a_bijection_over_the_streams() {
let streams = vec![
stream_len(0, 30, 2),
stream_len(1, 5, 2),
stream_len(2, 30, 2), stream_len(3, 12, 2),
stream_len(4, 5, 2),
stream_len(5, 99, 2),
];
let order = admission_order(&streams, true);
assert_eq!(order.len(), streams.len(), "no stream lost or duplicated");
let mut seen = order;
seen.sort_unstable();
assert_eq!(
seen,
(0..streams.len()).collect::<Vec<_>>(),
"packed order is a permutation of 0..n (every stream admitted exactly once)"
);
}
#[test]
fn packing_reduces_mean_window_length_variance() {
let streams = vec![
stream_len(0, 10, 2),
stream_len(1, 100, 2),
stream_len(2, 12, 2),
stream_len(3, 98, 2),
stream_len(4, 11, 2),
stream_len(5, 105, 2),
stream_len(6, 9, 2),
stream_len(7, 101, 2),
];
let width = 2;
let naive_var =
mean_window_len_variance(&streams, &admission_order(&streams, false), width);
let packed_var =
mean_window_len_variance(&streams, &admission_order(&streams, true), width);
assert!(
packed_var <= naive_var,
"packing must never increase window length variance (packed {packed_var} > naive {naive_var})"
);
assert!(
packed_var < naive_var,
"on a varied pool packing should strictly group similar lengths (packed {packed_var} vs naive {naive_var})"
);
}
#[test]
fn unpacked_admission_order_is_byte_identical_to_submission_order() {
let streams = vec![
stream_len(0, 30, 2),
stream_len(1, 5, 2),
stream_len(2, 99, 2),
stream_len(3, 12, 2),
];
assert_eq!(
admission_order(&streams, false),
(0..streams.len()).collect::<Vec<_>>(),
"pack-off admission order must equal today's submission order, byte-for-byte"
);
}
#[test]
fn packing_does_not_change_any_stream_output() {
let make = || {
vec![
stream_len(0, 10, 3),
stream_len(1, 100, 1),
stream_len(2, 12, 5),
stream_len(3, 98, 2),
stream_len(4, 11, 4),
stream_len(5, 105, 1),
]
};
let mut sched_off = BatchScheduler::new(2, 100);
let mut step_off = MockStep::new();
let out_off = sched_off
.run_with_pack(make(), &mut step_off, false)
.expect("run pack off");
let mut sched_on = BatchScheduler::new(2, 100);
let mut step_on = MockStep::new();
let out_on = sched_on
.run_with_pack(make(), &mut step_on, true)
.expect("run pack on");
assert_eq!(
out_on, out_off,
"packed admission must not change any stream's emitted tokens (lossless)"
);
assert!(sched_on.stats().max_active <= 2);
assert_eq!(sched_on.stats().max_concurrent_forwards, 1);
}
}