use super::qwen3_moe_generate::run_qwen3_moe_generate;
use crate::error::Result;
use crate::gguf::{MappedGGUFModel, OwnedQuantizedModel, QuantizedGenerateConfig};
pub const QWEN3MOE_GPU_FALLBACK_PREFIX: &str = "qwen3moe: the CUDA forward did not serve this run";
pub fn run_qwen3_moe_generate_dispatch(
mapped: &MappedGGUFModel,
model: &OwnedQuantizedModel,
input_tokens: &[u32],
gen_config: &QuantizedGenerateConfig,
no_gpu: bool,
) -> Result<(Vec<u32>, bool)> {
#[cfg(feature = "cuda")]
if !no_gpu {
match gpu::run_qwen3_moe_generate_gpu(mapped, model, input_tokens, gen_config) {
Ok(tokens) => return Ok((tokens, true)),
Err(reason) => {
eprintln!("{QWEN3MOE_GPU_FALLBACK_PREFIX}, falling back to CPU: {reason}");
},
}
}
#[cfg(not(feature = "cuda"))]
let _ = no_gpu;
let tokens = run_qwen3_moe_generate(mapped, model, input_tokens, gen_config)?;
Ok((tokens, false))
}
#[cfg(feature = "cuda")]
pub fn qwen3_moe_shape(
mapped: &MappedGGUFModel,
) -> std::result::Result<crate::gguf::Qwen3MoeShape, String> {
let m = &mapped.model;
let arch = m.architecture().unwrap_or("qwen3moe");
Ok(crate::gguf::Qwen3MoeShape {
num_experts: m
.expert_count()
.ok_or_else(|| format!("'{arch}.expert_count' is missing from the GGUF metadata"))?,
num_experts_per_tok: m.expert_used_count().ok_or_else(|| {
format!("'{arch}.expert_used_count' is missing from the GGUF metadata")
})?,
expert_dim: m.expert_feed_forward_length().ok_or_else(|| {
format!("'{arch}.expert_feed_forward_length' is missing from the GGUF metadata")
})?,
})
}
#[cfg(feature = "cuda")]
pub(crate) mod gpu {
use super::qwen3_moe_shape;
use crate::gguf::qwen3_moe_load::{load_qwen3_moe_layer, Qwen3MoeQuantizedLayer};
use crate::gguf::{
MappedGGUFModel, OwnedQuantizedKVCache, OwnedQuantizedModel, QuantizedGenerateConfig,
Qwen3MoeCudaModel, Qwen3MoeShape,
};
use crate::infer::qwen3_moe_generate::sample_from_logits;
pub(crate) const QWEN3MOE_F2_PROBE_MAX: usize = 64;
pub(crate) fn load_moe_layers(
mapped: &MappedGGUFModel,
num_layers: usize,
) -> std::result::Result<Vec<Qwen3MoeQuantizedLayer>, String> {
(0..num_layers)
.map(|il| {
load_qwen3_moe_layer(&mapped.model, mapped.data(), il)
.map_err(|e| format!("layer {il}'s MoE tensors would not load: {e}"))
})
.collect()
}
pub(crate) fn run_qwen3_moe_generate_gpu(
mapped: &MappedGGUFModel,
model: &OwnedQuantizedModel,
input_tokens: &[u32],
gen_config: &QuantizedGenerateConfig,
) -> std::result::Result<Vec<u32>, String> {
if input_tokens.is_empty() {
return Err("the prompt is empty".to_string());
}
let shape = qwen3_moe_shape(mapped)?;
let moe_layers = load_moe_layers(mapped, model.config().num_layers)?;
let build_start = std::time::Instant::now();
let executor = crate::cuda::CudaExecutor::new(0)
.map_err(|e| format!("CUDA initialization failed: {e}"))?;
let max_seq_len = input_tokens.len() + gen_config.max_tokens + 1;
let mut gpu = Qwen3MoeCudaModel::with_max_seq_len(
model,
&moe_layers,
shape,
mapped.data(),
executor,
max_seq_len,
)
.map_err(|e| format!("the CUDA model would not build: {e}"))?;
let build_ms = build_start.elapsed().as_secs_f64() * 1000.0;
let (device, vram_mb) = gpu.device_summary();
eprintln!(
"Backend: GPU (CUDA, {device}, {vram_mb} MB VRAM) [qwen3moe routed-expert forward, \
#3714; weights resident in {build_ms:.0} ms]"
);
f2_validate_qwen3_moe(
&mut gpu,
model,
&moe_layers,
shape,
mapped.data(),
input_tokens,
)?;
let decode_start = std::time::Instant::now();
crate::infer::mark_generation_start(); let tokens = decode(&mut gpu, input_tokens, gen_config)?;
let generated = tokens.len().saturating_sub(input_tokens.len());
let decode_s = decode_start.elapsed().as_secs_f64();
eprintln!(
"qwen3moe CUDA: {} prompt + {generated} generated tokens in {:.0} ms ({:.1} tok/s \
including the token-by-token prefill)",
input_tokens.len(),
decode_s * 1000.0,
(input_tokens.len() + generated) as f64 / decode_s.max(1e-9)
);
Ok(tokens)
}
fn decode(
gpu: &mut Qwen3MoeCudaModel<'_>,
input_tokens: &[u32],
gen_config: &QuantizedGenerateConfig,
) -> std::result::Result<Vec<u32>, String> {
use rand::SeedableRng;
let max_seq_len = input_tokens.len() + gen_config.max_tokens + 1;
let mut state = gpu
.new_state()
.map_err(|e| format!("the decode state would not allocate: {e}"))?;
let mut rng = rand::rngs::StdRng::seed_from_u64(gen_config.seed);
let mut logits = Vec::new();
for (pos, &token) in input_tokens.iter().enumerate() {
logits = gpu
.forward_single(token, &mut state, pos)
.map_err(|e| format!("the GPU forward failed at prompt position {pos}: {e}"))?;
}
let mut tokens = input_tokens.to_vec();
for _ in 0..gen_config.max_tokens {
let next = sample_from_logits(&logits, gen_config, &mut rng, &tokens)
.map_err(|e| format!("sampling failed: {e}"))?;
tokens.push(next);
if gen_config.stop_tokens.contains(&next) || tokens.len() >= max_seq_len {
break;
}
let pos = tokens.len() - 1;
logits = gpu
.forward_single(next, &mut state, pos)
.map_err(|e| format!("the GPU forward failed at decode position {pos}: {e}"))?;
}
Ok(tokens)
}
pub(crate) fn cpu_reference(
model: &OwnedQuantizedModel,
moe_layers: &[Qwen3MoeQuantizedLayer],
shape: Qwen3MoeShape,
data: &[u8],
probe: &[u32],
) -> std::result::Result<Vec<Vec<f32>>, String> {
crate::quantize::with_fp32_activations(|| {
cpu_reference_q8k_or_fp32(model, moe_layers, shape, data, probe)
})
}
fn cpu_reference_q8k_or_fp32(
model: &OwnedQuantizedModel,
moe_layers: &[Qwen3MoeQuantizedLayer],
shape: Qwen3MoeShape,
data: &[u8],
probe: &[u32],
) -> std::result::Result<Vec<Vec<f32>>, String> {
let mut cache = OwnedQuantizedKVCache::from_config(model.config(), probe.len() + 2);
let mut per_pos = Vec::with_capacity(probe.len() + 1);
let forward = |cache: &mut OwnedQuantizedKVCache, token: u32, pos: usize| {
model
.forward_single_qwen3_moe_with_cache(
token,
cache,
pos,
moe_layers,
shape.num_experts,
shape.num_experts_per_tok,
shape.expert_dim,
data,
)
.map_err(|e| format!("the CPU reference failed at position {pos}: {e}"))
};
for (pos, &token) in probe.iter().enumerate() {
per_pos.push(forward(&mut cache, token, pos)?);
}
let next = per_pos.last().map_or(0, |l| crate::infer::argmax_u32(l));
per_pos.push(forward(&mut cache, next, probe.len())?);
Ok(per_pos)
}
pub(crate) fn gpu_logits(
gpu: &mut Qwen3MoeCudaModel<'_>,
probe: &[u32],
decode_token: u32,
) -> std::result::Result<Vec<Vec<f32>>, String> {
let mut state = gpu
.new_state()
.map_err(|e| format!("the probe state would not allocate: {e}"))?;
let mut per_pos = Vec::with_capacity(probe.len() + 1);
for (pos, &token) in probe
.iter()
.chain(std::iter::once(&decode_token))
.enumerate()
{
per_pos.push(
gpu.forward_single(token, &mut state, pos)
.map_err(|e| format!("the GPU probe failed at position {pos}: {e}"))?,
);
}
Ok(per_pos)
}
fn f2_validate_qwen3_moe(
gpu: &mut Qwen3MoeCudaModel<'_>,
model: &OwnedQuantizedModel,
moe_layers: &[Qwen3MoeQuantizedLayer],
shape: Qwen3MoeShape,
data: &[u8],
prompt: &[u32],
) -> std::result::Result<(), String> {
if std::env::var("SKIP_PARITY_GATE").is_ok_and(|v| v == "1") {
eprintln!("F2 guard: SKIP_PARITY_GATE=1 — nothing was compared on this run");
return Ok(());
}
let probe = &prompt[prompt.len().saturating_sub(QWEN3MOE_F2_PROBE_MAX)..];
if probe.len() < 2 {
eprintln!(
"F2 guard: a {}-token prompt has no real position to compare; nothing was judged",
probe.len()
);
return Ok(());
}
let start = std::time::Instant::now();
let cpu = cpu_reference(model, moe_layers, shape, data, probe)?;
let decode_token = cpu
.get(probe.len() - 1)
.map_or(0, |l| crate::infer::argmax_u32(l));
let gpu_per_pos = gpu_logits(gpu, probe, decode_token)?;
let report = crate::infer::f2_multi_position_report(&cpu, &gpu_per_pos);
let ms = start.elapsed().as_secs_f64() * 1000.0;
if report.accepted {
eprintln!(
"F2 guard: GPU matches the CPU forward on {} positions (min cosine {:.4}) in {ms:.0} ms",
cpu.len(),
report.min_cosine_real
);
Ok(())
} else {
Err(crate::infer::f2_divergence_msg(
&report,
crate::infer::F2ProbePath::Serial,
))
}
}
}