use std::time::Instant;
use crate::paged_kv::SeqId;
pub type Token = u32;
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct Priority(pub u8);
impl Default for Priority {
fn default() -> Self {
Priority(128) }
}
#[derive(Debug, Clone)]
pub struct InferenceRequest {
pub id: SeqId,
pub input_tokens: Vec<Token>,
pub max_new_tokens: usize,
pub priority: Priority,
pub arrival_time: Instant,
pub estimated_tokens: usize,
}
impl InferenceRequest {
pub fn new(id: SeqId, input_tokens: Vec<Token>, max_new_tokens: usize) -> Self {
let estimated_tokens = input_tokens.len() + max_new_tokens;
Self {
id,
input_tokens,
max_new_tokens,
priority: Priority::default(),
arrival_time: Instant::now(),
estimated_tokens,
}
}
pub fn with_priority(mut self, priority: Priority) -> Self {
self.priority = priority;
self
}
pub fn input_len(&self) -> usize {
self.input_tokens.len()
}
}
#[derive(Debug, Clone)]
pub struct SequenceGroup {
pub request: InferenceRequest,
pub output_tokens: Vec<Token>,
pub is_finished: bool,
pub last_access: Instant,
pub num_steps: usize,
}
impl SequenceGroup {
pub fn new(request: InferenceRequest) -> Self {
Self {
request,
output_tokens: Vec::new(),
is_finished: false,
last_access: Instant::now(),
num_steps: 0,
}
}
pub fn total_tokens(&self) -> usize {
self.request.input_tokens.len() + self.output_tokens.len()
}
pub fn remaining_tokens(&self) -> usize {
self.request
.max_new_tokens
.saturating_sub(self.output_tokens.len())
}
pub fn add_token(&mut self, token: Token) {
self.output_tokens.push(token);
self.last_access = Instant::now();
self.num_steps += 1;
if self.output_tokens.len() >= self.request.max_new_tokens {
self.is_finished = true;
}
}
pub fn finish(&mut self) {
self.is_finished = true;
}
}