use crate::error::{RealizarError, Result};
pub type AprQ4kSession<S> = crate::session::Session<AprQ4kForward<S>>;
pub trait Q4kStep {
fn reset(&mut self);
fn step(&mut self, token: u32, position: usize) -> Result<Vec<f32>>;
}
pub struct AprQ4kForward<S: Q4kStep> {
step: S,
on_gpu: bool,
context_length: usize,
held: usize,
notices: Vec<String>,
}
impl<S: Q4kStep> AprQ4kForward<S> {
pub fn new(step: S, on_gpu: bool, context_length: Option<usize>) -> Self {
let backend = if on_gpu { "GPU" } else { "CPU" };
Self {
step,
on_gpu,
context_length: context_length.unwrap_or(usize::MAX).max(1),
held: 0,
notices: vec![format!("Backend: {backend}")],
}
}
}
impl<S: Q4kStep> crate::session::ArchForward for AprQ4kForward<S> {
fn arch(&self) -> &'static str {
"apr-q4k"
}
fn on_gpu(&self) -> bool {
self.on_gpu
}
fn context_length(&self) -> usize {
self.context_length
}
fn batched_prefills(&self) -> usize {
0
}
fn notices(&self) -> &[String] {
&self.notices
}
fn reserve(&mut self, _positions: usize) -> Result<bool> {
Ok(false)
}
fn forward(&mut self, tokens: &[u32], start: usize) -> Result<Vec<f32>> {
let start = if start == self.held && start > 0 {
start
} else {
self.step.reset();
self.held = 0;
0
};
let mut logits = Vec::new();
for (position, &token) in tokens.iter().enumerate().skip(start) {
self.held = 0;
logits = self.step.step(token, position)?;
self.held = position + 1;
}
if logits.is_empty() {
return Err(RealizarError::InvalidShape {
reason: format!(
"apr-q4k session: forward over {} tokens from {start} produced no logits",
tokens.len()
),
});
}
Ok(logits)
}
}
#[cfg(feature = "cuda")]
pub(crate) struct CudaQ4kStep<'a> {
executor: &'a mut crate::cuda::CudaExecutor,
config: &'a crate::gpu::adapters::apr_q4k::AprQ4KConfig,
weights: &'a crate::api::apr_q4k_scheduler::Q4kHostWeights,
kv_k: Vec<Vec<f32>>,
kv_v: Vec<Vec<f32>>,
}
#[cfg(feature = "cuda")]
impl<'a> CudaQ4kStep<'a> {
pub(crate) fn new(
executor: &'a mut crate::cuda::CudaExecutor,
config: &'a crate::gpu::adapters::apr_q4k::AprQ4KConfig,
weights: &'a crate::api::apr_q4k_scheduler::Q4kHostWeights,
) -> Self {
Self {
executor,
config,
weights,
kv_k: vec![Vec::new(); config.num_layers],
kv_v: vec![Vec::new(); config.num_layers],
}
}
}
#[cfg(feature = "cuda")]
impl Q4kStep for CudaQ4kStep<'_> {
fn reset(&mut self) {
self.kv_k.iter_mut().for_each(Vec::clear);
self.kv_v.iter_mut().for_each(Vec::clear);
}
fn step(&mut self, token: u32, position: usize) -> Result<Vec<f32>> {
crate::gpu::adapters::apr_q4k::forward_token_apr_q4k(
self.executor,
self.config,
&self.weights.embedding,
&self.weights.output_norm,
&self.weights.layer_norms,
&self.weights.qkv_biases,
&mut self.kv_k,
&mut self.kv_v,
token,
position,
)
}
}