use std::sync::Arc;
use crate::error::{RealizarError, Result};
use crate::gguf::{OwnedQuantizedKVCache, OwnedQuantizedModel};
pub const DENSE_GPU_FALLBACK_PREFIX: &str = "[dense] GPU forward failed, falling back to CPU";
pub type DenseSession = crate::session::Session<DenseForward>;
enum Backend {
Cpu {
model: Arc<OwnedQuantizedModel>,
cache: Option<OwnedQuantizedKVCache>,
},
#[cfg(feature = "cuda")]
Cuda {
model: Box<crate::gguf::OwnedQuantizedModelCuda>,
cache: Option<OwnedQuantizedKVCache>,
},
Moving,
}
enum Step {
#[cfg_attr(not(feature = "cuda"), allow(dead_code))]
Gpu(String),
Fatal(RealizarError),
}
impl From<RealizarError> for Step {
fn from(e: RealizarError) -> Self {
Self::Fatal(e)
}
}
pub struct DenseForward {
backend: Backend,
arch: &'static str,
context_length: usize,
capacity: usize,
turn_positions: usize,
held: usize,
notices: Vec<String>,
batched_prefills: usize,
}
impl DenseForward {
fn with_backend(backend: Backend) -> Self {
let mut forward = Self {
backend,
arch: "",
context_length: 1,
capacity: 0,
turn_positions: 0,
held: 0,
notices: Vec::new(),
batched_prefills: 0,
};
let config = &forward.model().config;
let (arch, context_length) = (
crate::tensor_names::normalize_architecture(&config.architecture),
config.context_length.max(1),
);
forward.arch = arch;
forward.context_length = context_length;
forward
}
#[must_use]
pub fn cpu(model: Arc<OwnedQuantizedModel>) -> Self {
let mut forward = Self::with_backend(Backend::Cpu { model, cache: None });
forward.notices.push("Backend: CPU".to_string());
forward
}
#[cfg(feature = "cuda")]
#[must_use]
pub fn cuda(model: crate::gguf::OwnedQuantizedModelCuda) -> Self {
let line = format!(
"Backend: GPU ({}, {} MB VRAM)",
model.device_name(),
model.vram_mb()
);
let mut forward = Self::with_backend(Backend::Cuda {
model: Box::new(model),
cache: None,
});
forward.notices.push(line);
forward
}
#[must_use]
pub fn model(&self) -> &OwnedQuantizedModel {
match &self.backend {
Backend::Cpu { model, .. } => model,
#[cfg(feature = "cuda")]
Backend::Cuda { model, .. } => model.model(),
Backend::Moving => {
unreachable!("a dense forward is only Moving inside fall_back_to_cpu")
},
}
}
fn ensure_capacity(&mut self) -> std::result::Result<(), Step> {
let positions = self.turn_positions.min(self.context_length).max(1);
let context_length = self.context_length;
let capacity = self.capacity;
let (target, built) = match &mut self.backend {
Backend::Cpu { model, cache } => {
if positions <= capacity && cache.is_some() {
return Ok(());
}
let target = positions
.max(capacity.saturating_mul(2))
.min(context_length)
.max(positions);
let config = &model.config;
(
target,
grow_or_build(cache, target, |n| {
OwnedQuantizedKVCache::from_config(config, n)
}),
)
},
#[cfg(feature = "cuda")]
Backend::Cuda { model, cache } => {
let device_max = model.executor().max_kv_len();
if device_max > 0 && positions > device_max {
return Err(Step::Gpu(format!(
"the turn needs {positions} positions and the device KV cache holds {device_max}"
)));
}
if positions <= capacity && cache.is_some() {
return Ok(());
}
let config = &model.model().config;
(
positions,
grow_or_build(cache, positions, |n| {
OwnedQuantizedKVCache::from_config(config, n)
}),
)
},
Backend::Moving => {
unreachable!("a dense forward is only Moving inside fall_back_to_cpu")
},
};
self.capacity = target;
if built {
self.held = 0;
}
Ok(())
}
#[cfg(feature = "cuda")]
fn fall_back_to_cpu(&mut self, reason: &str) -> Result<()> {
let line = format!("{DENSE_GPU_FALLBACK_PREFIX}: {reason}");
eprintln!("{line}");
self.notices.push(line);
if let Backend::Cuda { model, .. } = std::mem::replace(&mut self.backend, Backend::Moving) {
self.backend = Backend::Cpu {
model: Arc::new(model.into_model()),
cache: None,
};
}
self.capacity = 0;
self.held = 0;
match self.ensure_capacity() {
Ok(()) => Ok(()),
Err(Step::Fatal(e)) => Err(e),
Err(Step::Gpu(reason)) => Err(RealizarError::UnsupportedOperation {
operation: "dense_session".to_string(),
reason: format!("the CPU backend reported a GPU failure: {reason}"),
}),
}
}
fn recover(&mut self, step: Step) -> Result<bool> {
match step {
#[cfg(feature = "cuda")]
Step::Gpu(reason) => {
self.fall_back_to_cpu(&reason)?;
Ok(false)
},
#[cfg(not(feature = "cuda"))]
Step::Gpu(reason) => Err(RealizarError::UnsupportedOperation {
operation: "dense_session".to_string(),
reason,
}),
Step::Fatal(e) => Err(e),
}
}
fn resume_at(&mut self, start: usize) -> std::result::Result<usize, Step> {
let start = start.min(self.held);
if start == 0 {
self.reset_cache()?;
}
Ok(start)
}
fn reset_cache(&mut self) -> std::result::Result<(), Step> {
match &mut self.backend {
Backend::Cpu { cache, .. } => {
if let Some(cache) = cache {
cache.reset();
}
},
#[cfg(feature = "cuda")]
Backend::Cuda { model, cache } => {
if let Some(cache) = cache {
cuda_reset(model, cache);
}
},
Backend::Moving => {
unreachable!("a dense forward is only Moving inside fall_back_to_cpu")
},
}
self.held = 0;
Ok(())
}
fn try_forward(&mut self, tokens: &[u32], start: usize) -> std::result::Result<Vec<f32>, Step> {
let start = self.resume_at(start)?;
let logits = match &mut self.backend {
Backend::Cpu { model, cache } => {
let cache = cache.as_mut().ok_or_else(|| never_reserved("CPU"))?;
let mut logits = Vec::new();
for (pos, &token) in tokens.iter().enumerate().skip(start) {
logits = model.forward_single_with_cache(token, cache, pos)?;
}
logits
},
#[cfg(feature = "cuda")]
Backend::Cuda { model, cache } => {
let cache = cache.as_mut().ok_or_else(|| never_reserved("CUDA"))?;
cuda_forward(model, cache, tokens, start, &mut self.batched_prefills)
.map_err(Step::Gpu)?
},
Backend::Moving => {
unreachable!("a dense forward is only Moving inside fall_back_to_cpu")
},
};
self.held = tokens.len();
Ok(logits)
}
#[cfg_attr(
not(feature = "cuda"),
allow(clippy::unnecessary_wraps, clippy::unused_self)
)]
fn try_forward_greedy(
&mut self,
tokens: &[u32],
start: usize,
) -> std::result::Result<Option<u32>, Step> {
#[cfg(feature = "cuda")]
if matches!(self.backend, Backend::Cuda { .. }) {
let start = self.resume_at(start)?;
let Backend::Cuda { model, cache } = &mut self.backend else {
unreachable!("matched above");
};
let cache = cache.as_mut().ok_or_else(|| never_reserved("CUDA"))?;
let Some(next) =
cuda_forward_greedy(model, cache, tokens, start, &mut self.batched_prefills)
.map_err(Step::Gpu)?
else {
return Ok(None);
};
self.held = tokens.len();
return Ok(Some(next));
}
let _ = (tokens, start);
Ok(None)
}
}
fn never_reserved(backend: &str) -> Step {
Step::Fatal(RealizarError::InvalidShape {
reason: format!("dense session: the {backend} KV cache was never reserved"),
})
}
#[cfg(feature = "cuda")]
fn gpu_failed(what: &str, at: usize, e: RealizarError) -> String {
format!("the GPU {what} failed at {at}: {e}")
}
#[cfg(feature = "cuda")]
pub(crate) fn cuda_reset(
model: &mut crate::gguf::OwnedQuantizedModelCuda,
cache: &mut OwnedQuantizedKVCache,
) {
cache.reset();
model.executor_mut().reset_kv_cache_gpu();
}
#[cfg(feature = "cuda")]
pub(crate) fn cuda_forward(
model: &mut crate::gguf::OwnedQuantizedModelCuda,
cache: &mut OwnedQuantizedKVCache,
tokens: &[u32],
start: usize,
batched_prefills: &mut usize,
) -> std::result::Result<Vec<f32>, String> {
let last = tokens.len() - 1;
let mut from = start;
if start == 0 && last > 1 {
model
.run_prefill(tokens, cache, last, false, false)
.map_err(|e| gpu_failed("batched prefill", tokens.len(), e))?;
*batched_prefills += 1;
from = last;
}
let mut logits = Vec::new();
for (pos, &token) in tokens.iter().enumerate().skip(from) {
logits = model
.forward_gpu_resident(token, cache, pos)
.map_err(|e| gpu_failed("forward", pos, e))?;
}
Ok(logits)
}
#[cfg(feature = "cuda")]
pub(crate) fn cuda_forward_greedy(
model: &mut crate::gguf::OwnedQuantizedModelCuda,
cache: &mut OwnedQuantizedKVCache,
tokens: &[u32],
start: usize,
batched_prefills: &mut usize,
) -> std::result::Result<Option<u32>, String> {
let last = tokens.len() - 1;
if start == 0 && last > 0 {
let first = model
.run_prefill(tokens, cache, tokens.len(), false, true)
.map_err(|e| gpu_failed("batched prefill", tokens.len(), e))?;
*batched_prefills += 1;
return Ok(first);
}
for (pos, &token) in tokens.iter().enumerate().take(last).skip(start) {
model
.forward_gpu_resident(token, cache, pos)
.map_err(|e| gpu_failed("forward", pos, e))?;
}
model
.forward_gpu_resident_to_token_id(tokens[last], cache, last)
.map(Some)
.map_err(|e| gpu_failed("forward", last, e))
}
impl crate::session::ArchForward for DenseForward {
fn arch(&self) -> &'static str {
self.arch
}
fn on_gpu(&self) -> bool {
match self.backend {
#[cfg(feature = "cuda")]
Backend::Cuda { .. } => true,
_ => false,
}
}
fn context_length(&self) -> usize {
self.context_length
}
fn batched_prefills(&self) -> usize {
self.batched_prefills
}
fn notices(&self) -> &[String] {
&self.notices
}
fn reserve(&mut self, positions: usize) -> Result<bool> {
self.turn_positions = positions;
let held_before = self.held;
#[cfg(feature = "cuda")]
if let Backend::Cuda { model, .. } = &self.backend {
if let Err(e) = model.executor().make_current() {
self.fall_back_to_cpu(&format!(
"the CUDA context would not bind to this thread: {e}"
))?;
}
}
if let Err(step) = self.ensure_capacity() {
self.recover(step)?;
}
Ok(held_before > 0 && self.held == 0)
}
fn forward(&mut self, tokens: &[u32], start: usize) -> Result<Vec<f32>> {
loop {
match self.try_forward(tokens, start) {
Ok(logits) => return Ok(logits),
Err(step) => {
self.held = 0;
self.recover(step)?;
},
}
}
}
fn forward_greedy(&mut self, tokens: &[u32], start: usize) -> Result<Option<u32>> {
match self.try_forward_greedy(tokens, start) {
Ok(next) => Ok(next),
Err(step) => {
self.held = 0;
self.recover(step)?;
Ok(None)
},
}
}
}
pub fn dense_turn<F: crate::session::ArchForward>(
session: &mut crate::session::Session<F>,
prompt: &[u32],
config: &crate::gguf::QuantizedGenerateConfig,
) -> Result<(Vec<u32>, bool)> {
dense_stream(session, prompt, config, &mut |_| true)
}
pub fn dense_stream<F: crate::session::ArchForward>(
session: &mut crate::session::Session<F>,
prompt: &[u32],
config: &crate::gguf::QuantizedGenerateConfig,
on_token: &mut dyn FnMut(u32) -> bool,
) -> Result<(Vec<u32>, bool)> {
let maximum = session.context_length();
if prompt.len() > maximum {
return Err(RealizarError::ContextLimitExceeded {
provided: prompt.len(),
maximum,
});
}
let turn = session.generate(prompt, config, &mut |t| {
config.stop_tokens.contains(&t) || on_token(t)
})?;
let mut tokens = turn.tokens;
if tokens.len() > prompt.len()
&& tokens
.last()
.is_some_and(|t| config.stop_tokens.contains(t))
{
tokens.pop();
}
Ok((tokens, turn.used_gpu))
}
fn grow_or_build(
cache: &mut Option<OwnedQuantizedKVCache>,
target: usize,
build: impl FnOnce(usize) -> OwnedQuantizedKVCache,
) -> bool {
if let Some(cache) = cache {
cache.grow_to(target);
return false;
}
*cache = Some(build(target));
true
}
#[must_use]
pub fn cap_context(context_length: usize, device_kv: usize) -> usize {
let context = context_length.max(1);
match device_kv {
0 => context,
device => context.min(device),
}
}
#[cfg(test)]
#[path = "dense_session_tests.rs"]
mod tests;
#[cfg(test)]
mod cap_context_tests {
use super::cap_context;
#[test]
fn d5_the_device_kv_caps_a_longer_model_context() {
assert_eq!(cap_context(32768, 4096), 4096);
}
#[test]
fn d5_a_shorter_model_context_is_kept() {
assert_eq!(cap_context(2048, 4096), 2048);
}
#[test]
fn d5_no_device_cache_means_the_model_context() {
assert_eq!(cap_context(32768, 0), 32768);
}
#[test]
fn d5_a_zero_model_context_still_holds_one_position() {
assert_eq!(cap_context(0, 0), 1);
assert_eq!(cap_context(0, 4096), 1);
}
}