use crate::policy::radix::{align_ceil, align_down};
#[derive(Debug)]
pub struct SlotTable {
free: Vec<u32>,
capacity: usize,
}
impl SlotTable {
pub fn new(capacity: usize) -> Self {
SlotTable {
free: (0..capacity as u32).collect(),
capacity,
}
}
pub fn capacity(&self) -> usize {
self.capacity
}
pub fn available(&self) -> usize {
self.free.len()
}
pub fn allocate(&mut self) -> Option<u32> {
self.free.pop()
}
pub fn free(&mut self, slot: u32) {
debug_assert!(!self.free.contains(&slot), "slot {slot} was freed twice");
self.free.push(slot);
}
pub fn rebuild(&mut self, capacity: usize) {
assert_eq!(
self.free.len(),
self.capacity,
"a slot table may only be rebuilt while idle"
);
self.capacity = capacity;
self.free = (0..capacity as u32).collect();
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct DecodingRequest {
pub uid: u64,
pub remaining: usize,
}
#[derive(Debug, Default)]
pub struct DecodeSet {
running: Vec<DecodingRequest>,
}
impl DecodeSet {
pub fn new() -> Self {
DecodeSet::default()
}
pub fn len(&self) -> usize {
self.running.len()
}
pub fn is_empty(&self) -> bool {
self.running.is_empty()
}
pub fn runnable(&self) -> bool {
!self.running.is_empty()
}
pub fn admit(&mut self, requests: impl IntoIterator<Item = DecodingRequest>) {
for request in requests {
if let Some(slot) = self.running.iter_mut().find(|r| r.uid == request.uid) {
*slot = request;
} else {
self.running.push(request);
}
}
self.running.retain(|r| r.remaining > 0);
}
pub fn remove(&mut self, uid: u64) -> Option<DecodingRequest> {
let index = self.running.iter().position(|r| r.uid == uid)?;
Some(self.running.remove(index))
}
pub fn inflight_tokens(&self, page_size: usize) -> usize {
let reserved = (page_size - 1) * self.running.len();
self.running.iter().map(|r| r.remaining).sum::<usize>() + reserved
}
pub fn next_batch(&self) -> Vec<DecodingRequest> {
let mut batch = self.running.clone();
batch.sort_by_key(|r| r.uid);
batch
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct PendingRequest {
pub uid: u64,
pub prompt_len: usize,
pub output_len: usize,
pub chunk: Option<ChunkState>,
}
impl PendingRequest {
pub fn new(uid: u64, prompt_len: usize, output_len: usize) -> Self {
PendingRequest {
uid,
prompt_len,
output_len,
chunk: None,
}
}
pub fn is_continuation(&self) -> bool {
self.chunk.is_some()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ChunkState {
pub computed_len: usize,
pub slot: u32,
pub swa_evicted_len: usize,
pub locked_prefix_len: usize,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub struct Capacity {
pub request_slots: usize,
pub kv_tokens: usize,
pub swa_tokens: usize,
pub recurrent_slots: usize,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Geometry {
pub page_size: usize,
pub sliding_window: Option<usize>,
pub recurrent: bool,
}
impl Default for Geometry {
fn default() -> Self {
Geometry {
page_size: 1,
sliding_window: None,
recurrent: false,
}
}
}
const RECURRENT_SLOTS_PER_REQUEST: usize = 3;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum NotAdmitted {
NoRequestSlot,
KvBudget,
WindowBudget,
RecurrentBudget,
NoWholePage,
BudgetSpent,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct AdmittedChunk {
pub uid: u64,
pub slot: u32,
pub computed_len: usize,
pub chunk_len: usize,
pub chunked: bool,
pub admission: Option<PromptAdmission>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct PromptAdmission {
pub uid: u64,
pub prompt_tokens: usize,
pub cached_tokens: usize,
}
#[derive(Debug)]
pub struct PrefillPass {
token_budget: usize,
reserved_tokens: usize,
reserved_swa: usize,
geometry: Geometry,
}
impl PrefillPass {
pub fn new(token_budget: usize, inflight_tokens: usize, geometry: Geometry) -> Self {
PrefillPass {
token_budget,
reserved_tokens: inflight_tokens,
reserved_swa: 0,
geometry,
}
}
pub fn token_budget(&self) -> usize {
self.token_budget
}
pub fn reserved_tokens(&self) -> usize {
self.reserved_tokens
}
pub fn check_admission(
&self,
request: &PendingRequest,
cached_len: usize,
capacity: &Capacity,
) -> Result<(), NotAdmitted> {
if capacity.request_slots == 0 {
return Err(NotAdmitted::NoRequestSlot);
}
let extend_len = request.prompt_len.saturating_sub(cached_len);
let estimated = extend_len + request.output_len;
if estimated + self.reserved_tokens > capacity.kv_tokens {
return Err(NotAdmitted::KvBudget);
}
if self.geometry.recurrent && capacity.recurrent_slots < RECURRENT_SLOTS_PER_REQUEST {
return Err(NotAdmitted::RecurrentBudget);
}
if let Some(window) = self.geometry.sliding_window {
let need = align_ceil(extend_len.max(1).min(window) + 1, self.geometry.page_size);
if capacity.swa_tokens.saturating_sub(self.reserved_swa) < need {
return Err(NotAdmitted::WindowBudget);
}
}
Ok(())
}
pub fn take_chunk(
&mut self,
request: &PendingRequest,
cached_len: usize,
slot: u32,
capacity: &Capacity,
) -> Result<AdmittedChunk, NotAdmitted> {
if self.token_budget == 0 {
return Err(NotAdmitted::BudgetSpent);
}
let computed_len = match request.chunk {
Some(chunk) => chunk.computed_len,
None => cached_len,
};
let remaining = request.prompt_len.saturating_sub(computed_len);
let mut chunk_len = self.token_budget.min(remaining);
if let Some(window) = self.geometry.sliding_window {
chunk_len =
self.window_bounded_chunk(request, computed_len, chunk_len, window, capacity)?;
}
if chunk_len == 0 {
return Err(NotAdmitted::NoWholePage);
}
let chunked = chunk_len < remaining;
self.token_budget -= chunk_len;
if request.chunk.is_none() {
self.reserved_tokens += remaining + request.output_len;
}
Ok(AdmittedChunk {
uid: request.uid,
slot,
computed_len,
chunk_len,
chunked,
admission: request.chunk.is_none().then_some(PromptAdmission {
uid: request.uid,
prompt_tokens: request.prompt_len,
cached_tokens: cached_len,
}),
})
}
fn window_bounded_chunk(
&mut self,
request: &PendingRequest,
computed_len: usize,
chunk_len: usize,
window: usize,
capacity: &Capacity,
) -> Result<usize, NotAdmitted> {
let page = self.geometry.page_size;
let chunk = request.chunk.unwrap_or(ChunkState {
computed_len,
slot: 0,
swa_evicted_len: 0,
locked_prefix_len: 0,
});
let newly_evicted = align_down(computed_len.saturating_sub(window + page), page);
let already_gone = chunk.swa_evicted_len.max(chunk.locked_prefix_len);
let self_reclaim = newly_evicted.saturating_sub(already_gone);
let budget = (capacity.swa_tokens + self_reclaim).saturating_sub(self.reserved_swa);
let max_end = (computed_len.div_ceil(page) + budget / page) * page;
let mut bounded = chunk_len.min(max_end.saturating_sub(computed_len));
let remaining = request.prompt_len.saturating_sub(computed_len);
if bounded > 0 && bounded < remaining {
let aligned = align_down(computed_len + bounded, page).saturating_sub(computed_len);
if aligned == 0 {
return Err(NotAdmitted::NoWholePage);
}
bounded = aligned;
}
self.reserved_swa +=
((computed_len + bounded).div_ceil(page) - computed_len.div_ceil(page)) * page;
Ok(bounded)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum FinishReason {
Stop,
Length,
}
impl FinishReason {
pub fn as_str(self) -> &'static str {
match self {
FinishReason::Stop => "stop",
FinishReason::Length => "length",
}
}
}
pub fn finish_reason(
hit_eos: bool,
matched_stop: bool,
budget_exhausted: bool,
) -> Option<FinishReason> {
if hit_eos || matched_stop {
Some(FinishReason::Stop)
} else if budget_exhausted {
Some(FinishReason::Length)
} else {
None
}
}
#[cfg(test)]
pub(crate) mod tests {
use super::*;
pub(crate) fn geometry() -> Geometry {
Geometry {
page_size: 16,
..Geometry::default()
}
}
pub(crate) fn roomy() -> Capacity {
Capacity {
request_slots: 4,
kv_tokens: 100_000,
swa_tokens: 100_000,
recurrent_slots: 64,
}
}
#[test]
fn slots_are_handed_out_last_in_first_out() {
let mut table = SlotTable::new(3);
assert_eq!(table.available(), 3);
assert_eq!(table.allocate(), Some(2));
assert_eq!(table.allocate(), Some(1));
table.free(2);
assert_eq!(table.allocate(), Some(2), "the most recent slot comes back");
assert_eq!(table.allocate(), Some(0));
assert_eq!(table.allocate(), None);
}
#[test]
#[should_panic(expected = "only be rebuilt while idle")]
fn a_slot_table_with_a_live_request_may_not_be_rebuilt() {
let mut table = SlotTable::new(3);
table.allocate();
table.rebuild(8);
}
#[test]
fn inflight_tokens_reserve_a_page_per_running_request() {
let mut decode = DecodeSet::new();
decode.admit([
DecodingRequest {
uid: 1,
remaining: 100,
},
DecodingRequest {
uid: 2,
remaining: 50,
},
]);
assert_eq!(decode.inflight_tokens(16), 150 + 2 * 15);
assert_eq!(decode.inflight_tokens(1), 150, "no page, no reservation");
}
#[test]
fn a_finished_request_leaves_the_decode_set() {
let mut decode = DecodeSet::new();
decode.admit([DecodingRequest {
uid: 1,
remaining: 1,
}]);
assert!(decode.runnable());
decode.admit([DecodingRequest {
uid: 1,
remaining: 0,
}]);
assert!(!decode.runnable());
assert!(decode.next_batch().is_empty());
}
#[test]
fn decode_batches_come_out_in_uid_order() {
let mut decode = DecodeSet::new();
decode.admit([
DecodingRequest {
uid: 7,
remaining: 4,
},
DecodingRequest {
uid: 2,
remaining: 4,
},
]);
let uids: Vec<u64> = decode.next_batch().iter().map(|r| r.uid).collect();
assert_eq!(uids, vec![2, 7]);
}
#[test]
fn admission_reserves_the_whole_prompt_and_output() {
let mut pass = PrefillPass::new(8192, 0, geometry());
let request = PendingRequest::new(1, 1000, 500);
let capacity = Capacity {
request_slots: 1,
kv_tokens: 1400, ..roomy()
};
assert_eq!(
pass.check_admission(&request, 0, &capacity),
Err(NotAdmitted::KvBudget)
);
let capacity = Capacity {
kv_tokens: 1500,
..capacity
};
assert!(pass.check_admission(&request, 0, &capacity).is_ok());
let chunk = pass.take_chunk(&request, 0, 0, &capacity).unwrap();
assert_eq!(pass.reserved_tokens(), 1500);
assert_eq!(chunk.chunk_len, 1000, "the budget covers the whole prompt");
assert!(!chunk.chunked);
}
#[test]
fn a_prefix_hit_shrinks_the_reservation() {
let mut pass = PrefillPass::new(8192, 0, geometry());
let request = PendingRequest::new(1, 1000, 500);
let capacity = Capacity {
kv_tokens: 800,
..roomy()
};
assert_eq!(
pass.check_admission(&request, 0, &capacity),
Err(NotAdmitted::KvBudget)
);
assert!(pass.check_admission(&request, 800, &capacity).is_ok());
let chunk = pass.take_chunk(&request, 800, 0, &capacity).unwrap();
assert_eq!(chunk.chunk_len, 200);
assert_eq!(chunk.computed_len, 800);
assert_eq!(pass.reserved_tokens(), 700);
}
#[test]
fn no_free_slot_refuses_before_anything_else_is_computed() {
let pass = PrefillPass::new(8192, 0, geometry());
let capacity = Capacity {
request_slots: 0,
..roomy()
};
assert_eq!(
pass.check_admission(&PendingRequest::new(1, 8, 8), 0, &capacity),
Err(NotAdmitted::NoRequestSlot)
);
}
#[test]
fn a_recurrent_model_seats_three_slots_per_request() {
let geometry = Geometry {
recurrent: true,
..geometry()
};
let pass = PrefillPass::new(8192, 0, geometry);
let request = PendingRequest::new(1, 64, 8);
assert!(pass
.check_admission(
&request,
0,
&Capacity {
recurrent_slots: 3,
..roomy()
}
)
.is_ok());
assert_eq!(
pass.check_admission(
&request,
0,
&Capacity {
recurrent_slots: 2,
..roomy()
}
),
Err(NotAdmitted::RecurrentBudget)
);
}
#[test]
fn a_window_model_needs_a_seat_in_the_window_pool() {
let geometry = Geometry {
sliding_window: Some(512),
..geometry()
};
let pass = PrefillPass::new(8192, 0, geometry);
let request = PendingRequest::new(1, 2000, 100);
assert_eq!(
pass.check_admission(
&request,
0,
&Capacity {
swa_tokens: 16,
..roomy()
}
),
Err(NotAdmitted::WindowBudget)
);
assert!(pass
.check_admission(
&request,
0,
&Capacity {
swa_tokens: 528,
..roomy()
}
)
.is_ok());
}
#[test]
fn a_long_prompt_is_chunked_and_charged_once() {
let mut pass = PrefillPass::new(256, 0, geometry());
let request = PendingRequest::new(1, 1000, 100);
let capacity = roomy();
let first = pass.take_chunk(&request, 0, 3, &capacity).unwrap();
assert_eq!(first.chunk_len, 256);
assert!(first.chunked);
assert_eq!(
first.admission,
Some(PromptAdmission {
uid: 1,
prompt_tokens: 1000,
cached_tokens: 0
}),
"the whole prompt is reported, not the chunk"
);
assert_eq!(pass.reserved_tokens(), 1100);
assert_eq!(pass.token_budget(), 0);
let mut pass = PrefillPass::new(256, 1100, geometry());
let continued = PendingRequest {
chunk: Some(ChunkState {
computed_len: 256,
slot: 3,
swa_evicted_len: 0,
locked_prefix_len: 0,
}),
..request
};
let second = pass.take_chunk(&continued, 0, 3, &capacity).unwrap();
assert_eq!(second.computed_len, 256);
assert_eq!(second.chunk_len, 256);
assert_eq!(second.admission, None, "reported once, on the first chunk");
assert_eq!(
pass.reserved_tokens(),
1100,
"a continuation reserves nothing new"
);
}
#[test]
fn a_window_continuation_stops_on_a_page_boundary() {
let geometry = Geometry {
page_size: 128,
sliding_window: Some(512),
..Geometry::default()
};
let mut pass = PrefillPass::new(724, 0, geometry);
let request = PendingRequest::new(1, 2000, 100);
let chunk = pass.take_chunk(&request, 0, 0, &roomy()).unwrap();
assert_eq!(chunk.chunk_len, 640, "5 whole pages, not 724");
assert!(chunk.chunked);
}
#[test]
fn the_last_chunk_may_end_mid_page() {
let geometry = Geometry {
page_size: 128,
sliding_window: Some(512),
..Geometry::default()
};
let mut pass = PrefillPass::new(8192, 0, geometry);
let request = PendingRequest::new(1, 700, 100);
let chunk = pass.take_chunk(&request, 0, 0, &roomy()).unwrap();
assert_eq!(chunk.chunk_len, 700);
assert!(!chunk.chunked);
}
#[test]
fn a_chunk_that_cannot_reach_a_page_boundary_is_refused() {
let geometry = Geometry {
page_size: 128,
sliding_window: Some(512),
..Geometry::default()
};
let mut pass = PrefillPass::new(8192, 0, geometry);
let request = PendingRequest::new(1, 2000, 100);
let capacity = Capacity {
swa_tokens: 64, ..roomy()
};
assert_eq!(
pass.take_chunk(&request, 0, 0, &capacity),
Err(NotAdmitted::NoWholePage)
);
}
#[test]
fn a_sliding_window_request_reclaims_its_own_state() {
let geometry = Geometry {
page_size: 128,
sliding_window: Some(512),
..Geometry::default()
};
let capacity = Capacity {
swa_tokens: 256,
..roomy()
};
let deep = PendingRequest {
chunk: Some(ChunkState {
computed_len: 4096,
slot: 0,
swa_evicted_len: 0,
locked_prefix_len: 0,
}),
..PendingRequest::new(1, 8192, 100)
};
let mut pass = PrefillPass::new(8192, 0, geometry);
let chunk = pass.take_chunk(&deep, 0, 0, &capacity).unwrap();
assert!(
chunk.chunk_len > 256,
"the request's own slid-out window paid for more than the pool had ({})",
chunk.chunk_len
);
let already = PendingRequest {
chunk: Some(ChunkState {
computed_len: 4096,
slot: 0,
swa_evicted_len: 4096,
locked_prefix_len: 0,
}),
..deep
};
let mut pass = PrefillPass::new(8192, 0, geometry);
let chunk = pass.take_chunk(&already, 0, 0, &capacity).unwrap();
assert_eq!(chunk.chunk_len, 256);
}
#[test]
fn one_pass_charges_each_admission_against_the_last() {
let capacity = Capacity {
kv_tokens: 1000,
..roomy()
};
let mut pass = PrefillPass::new(8192, 0, geometry());
let first = PendingRequest::new(1, 400, 200);
assert!(pass.check_admission(&first, 0, &capacity).is_ok());
pass.take_chunk(&first, 0, 0, &capacity).unwrap();
let second = PendingRequest::new(2, 400, 200);
assert_eq!(
pass.check_admission(&second, 0, &capacity),
Err(NotAdmitted::KvBudget),
"600 already reserved leaves 400 for a request that needs 600"
);
}
#[test]
fn stopping_beats_running_out_of_budget() {
assert_eq!(finish_reason(true, false, true), Some(FinishReason::Stop));
assert_eq!(finish_reason(false, true, true), Some(FinishReason::Stop));
assert_eq!(
finish_reason(false, false, true),
Some(FinishReason::Length)
);
assert_eq!(finish_reason(false, false, false), None);
assert_eq!(FinishReason::Stop.as_str(), "stop");
}
}