use crate::error::{RealizarError, Result};
use crate::gguf::qwen3_moe_load::{load_qwen3_moe_layer, Qwen3MoeQuantizedLayer};
use crate::gguf::{MappedGGUFModel, OwnedQuantizedKVCache, OwnedQuantizedModel};
pub type Qwen3MoeSession<'a> = crate::session::Session<Qwen3MoeForward<'a>>;
pub struct Qwen3MoeForward<'a> {
mapped: &'a MappedGGUFModel,
model: &'a OwnedQuantizedModel,
moe_layers: Vec<Qwen3MoeQuantizedLayer>,
num_experts: usize,
num_experts_per_tok: usize,
moe_intermediate: usize,
context_length: usize,
cache: Option<OwnedQuantizedKVCache>,
capacity: usize,
held: usize,
notices: Vec<String>,
}
impl<'a> Qwen3MoeForward<'a> {
pub fn cpu(mapped: &'a MappedGGUFModel, model: &'a OwnedQuantizedModel) -> Result<Self> {
let config = model.config();
let arch = crate::tensor_names::normalize_architecture(&config.architecture);
if arch != "qwen3_moe" {
return Err(RealizarError::InvalidShape {
reason: format!(
"qwen3_moe session: arch '{}' (canonical '{arch}') is not qwen3_moe — \
the caller should open a DenseSession instead",
config.architecture
),
});
}
let missing = |key: &str| RealizarError::InvalidShape {
reason: format!(
"qwen3_moe session: missing '{}.{key}' in GGUF metadata",
config.architecture
),
};
let num_experts = mapped
.model
.expert_count()
.ok_or_else(|| missing("expert_count"))?;
let num_experts_per_tok = mapped
.model
.expert_used_count()
.ok_or_else(|| missing("expert_used_count"))?;
let moe_intermediate = mapped
.model
.expert_feed_forward_length()
.ok_or_else(|| missing("expert_feed_forward_length"))?;
let data = mapped.data();
let moe_layers = (0..config.num_layers)
.map(|layer_idx| load_qwen3_moe_layer(&mapped.model, data, layer_idx))
.collect::<Result<Vec<_>>>()?;
Ok(Self {
mapped,
model,
moe_layers,
num_experts,
num_experts_per_tok,
moe_intermediate,
context_length: config.context_length.max(1),
cache: None,
capacity: 0,
held: 0,
notices: vec!["Backend: CPU".to_string()],
})
}
}
impl crate::session::ArchForward for Qwen3MoeForward<'_> {
fn arch(&self) -> &'static str {
"qwen3_moe"
}
fn on_gpu(&self) -> bool {
false
}
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> {
let positions = positions.min(self.context_length).max(1);
if positions <= self.capacity && self.cache.is_some() {
return Ok(false);
}
let capacity = positions
.max(self.capacity.saturating_mul(2))
.min(self.context_length)
.max(positions);
match &mut self.cache {
Some(cache) => cache.grow_to(capacity),
None => {
self.cache = Some(OwnedQuantizedKVCache::from_config(
self.model.config(),
capacity,
));
},
}
self.capacity = capacity;
Ok(false)
}
fn forward(&mut self, tokens: &[u32], start: usize) -> Result<Vec<f32>> {
if self.cache.is_none() {
self.reserve(tokens.len())?;
}
let cache = self
.cache
.as_mut()
.ok_or_else(|| RealizarError::InvalidShape {
reason: "qwen3_moe session: no KV cache after reserve".to_string(),
})?;
let start = if start == self.held { start } else { 0 };
if start == 0 {
cache.reset();
self.held = 0;
}
let data = self.mapped.data();
let mut logits = Vec::new();
for (pos, &token) in tokens.iter().enumerate().skip(start) {
logits = self.model.forward_single_qwen3_moe_with_cache(
token,
cache,
pos,
&self.moe_layers,
self.num_experts,
self.num_experts_per_tok,
self.moe_intermediate,
data,
)?;
self.held = pos + 1;
}
if logits.is_empty() {
return Err(RealizarError::InvalidShape {
reason: format!(
"qwen3_moe session: forward over {} tokens from {start} produced no logits",
tokens.len()
),
});
}
Ok(logits)
}
}