use crate::error::InferenceError;
use crate::grammar::GrammarEngine;
use crate::stop_reason::StopReason;
use std::path::Path;
use std::sync::Arc;
pub const QWEN_CHAT_IM_END_TOKEN_ID: u32 = 248_046;
pub const QWEN3_THINK_OPEN_TOKEN_ID: u32 = 248_068;
pub const QWEN3_THINK_CLOSE_TOKEN_ID: u32 = 248_069;
pub const QWEN3_NEWLINE_TOKEN_ID: u32 = 198;
pub const QWEN3_NO_THINK_PREFIX: [u32; 6] = [
QWEN3_THINK_OPEN_TOKEN_ID,
QWEN3_NEWLINE_TOKEN_ID,
QWEN3_NEWLINE_TOKEN_ID,
QWEN3_THINK_CLOSE_TOKEN_ID,
QWEN3_NEWLINE_TOKEN_ID,
QWEN3_NEWLINE_TOKEN_ID,
];
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum LayerType {
LinearAttention,
FullAttention,
}
#[derive(Debug, Clone, serde::Deserialize)]
#[serde(default)]
pub struct Qwen35Config {
pub hidden_size: usize,
pub num_hidden_layers: usize,
pub vocab_size: usize,
pub intermediate_size: usize,
pub rms_norm_eps: f32,
pub num_attention_heads: usize,
pub num_key_value_heads: usize,
pub head_dim: usize,
pub rope_theta: f64,
pub partial_rotary_factor: f32,
#[serde(default)]
pub rope_parameters: Option<RopeParams>,
pub linear_num_key_heads: usize,
#[serde(default)]
pub linear_num_value_heads: Option<usize>,
pub linear_key_head_dim: usize,
pub linear_value_head_dim: usize,
pub linear_conv_kernel_dim: usize,
#[serde(default)]
pub num_experts: Option<usize>,
#[serde(default)]
pub num_experts_per_tok: Option<usize>,
#[serde(default)]
pub moe_intermediate_size: Option<usize>,
#[serde(default)]
pub shared_expert_intermediate_size: Option<usize>,
#[serde(default)]
pub output_router_logits: bool,
#[serde(default)]
pub router_aux_loss_coef: Option<f32>,
#[serde(default = "default_tie_word_embeddings")]
pub tie_word_embeddings: bool,
#[serde(default)]
pub mtp_num_hidden_layers: usize,
#[serde(default)]
pub mtp_use_dedicated_embeddings: bool,
#[serde(default)]
pub quarot_rotation_seed: Option<u64>,
pub full_attention_interval: usize,
pub layer_types: Vec<LayerType>,
#[serde(default)]
pub layer_mask: Vec<bool>,
pub eos_token_id: u32,
#[serde(default = "default_max_position_embeddings")]
pub max_position_embeddings: usize,
}
fn default_tie_word_embeddings() -> bool {
true
}
fn default_max_position_embeddings() -> usize {
4096
}
#[derive(Debug, Clone, Default, serde::Deserialize)]
pub struct RopeParams {
#[serde(default)]
pub rope_theta: f64,
#[serde(default)]
pub partial_rotary_factor: Option<f32>,
}
#[derive(Debug, serde::Deserialize)]
struct HfQwenConfigFile {
#[serde(default)]
text_config: Option<Qwen35Config>,
#[serde(default)]
tie_word_embeddings: Option<bool>,
}
impl Default for Qwen35Config {
fn default() -> Self {
Self::qwen36_35b_a3b()
}
}
impl Qwen35Config {
pub fn qwen35_2b() -> Self {
let num_hidden_layers = 24;
let full_attention_interval = 4;
let layer_types = compute_layer_types(num_hidden_layers, full_attention_interval);
Self {
hidden_size: 2048,
num_hidden_layers,
vocab_size: 248_320,
intermediate_size: 6144,
rms_norm_eps: 1e-6,
num_attention_heads: 8,
num_key_value_heads: 2,
head_dim: 256,
rope_theta: 10_000_000.0,
partial_rotary_factor: 0.25,
rope_parameters: None,
linear_num_key_heads: 16,
linear_num_value_heads: Some(16),
linear_key_head_dim: 128,
linear_value_head_dim: 128,
linear_conv_kernel_dim: 4,
num_experts: None,
num_experts_per_tok: None,
moe_intermediate_size: None,
shared_expert_intermediate_size: None,
output_router_logits: false,
router_aux_loss_coef: None,
tie_word_embeddings: true,
full_attention_interval,
layer_types,
layer_mask: vec![true; num_hidden_layers],
eos_token_id: 248_044,
max_position_embeddings: 262_144,
mtp_num_hidden_layers: 0,
mtp_use_dedicated_embeddings: false,
quarot_rotation_seed: None,
}
}
pub fn qwen35_0_8b() -> Self {
let num_hidden_layers = 24;
let full_attention_interval = 4;
let layer_types = compute_layer_types(num_hidden_layers, full_attention_interval);
Self {
hidden_size: 1024,
num_hidden_layers,
vocab_size: 248_320,
intermediate_size: 3584,
rms_norm_eps: 1e-6,
num_attention_heads: 8,
num_key_value_heads: 2,
head_dim: 256,
rope_theta: 10_000_000.0,
partial_rotary_factor: 0.25,
rope_parameters: None,
linear_num_key_heads: 16,
linear_num_value_heads: Some(16),
linear_key_head_dim: 128,
linear_value_head_dim: 128,
linear_conv_kernel_dim: 4,
num_experts: None,
num_experts_per_tok: None,
moe_intermediate_size: None,
shared_expert_intermediate_size: None,
output_router_logits: false,
router_aux_loss_coef: None,
tie_word_embeddings: true,
full_attention_interval,
layer_types,
layer_mask: vec![true; num_hidden_layers],
eos_token_id: 248_044,
max_position_embeddings: 262_144,
mtp_num_hidden_layers: 1,
mtp_use_dedicated_embeddings: false,
quarot_rotation_seed: None,
}
}
pub fn qwen36_35b_a3b() -> Self {
let num_hidden_layers = 40;
let full_attention_interval = 4;
let layer_types = compute_layer_types(num_hidden_layers, full_attention_interval);
Self {
hidden_size: 2048,
num_hidden_layers,
vocab_size: 248_320,
intermediate_size: 6144,
rms_norm_eps: 1e-6,
num_attention_heads: 16,
num_key_value_heads: 2,
head_dim: 256,
rope_theta: 10_000_000.0,
partial_rotary_factor: 0.25,
rope_parameters: None,
linear_num_key_heads: 16,
linear_num_value_heads: Some(32),
linear_key_head_dim: 128,
linear_value_head_dim: 128,
linear_conv_kernel_dim: 4,
num_experts: Some(256),
num_experts_per_tok: Some(8),
moe_intermediate_size: Some(512),
shared_expert_intermediate_size: Some(512),
output_router_logits: false,
router_aux_loss_coef: Some(0.001),
tie_word_embeddings: false,
full_attention_interval,
layer_types,
layer_mask: vec![true; num_hidden_layers],
eos_token_id: 248_044,
max_position_embeddings: 262_144,
mtp_num_hidden_layers: 1,
mtp_use_dedicated_embeddings: false,
quarot_rotation_seed: None,
}
}
pub fn qwen36_27b() -> Self {
let num_hidden_layers = 64;
let full_attention_interval = 4;
let layer_types = compute_layer_types(num_hidden_layers, full_attention_interval);
Self {
hidden_size: 5120,
num_hidden_layers,
vocab_size: 248_320,
intermediate_size: 17408,
rms_norm_eps: 1e-6,
num_attention_heads: 24,
num_key_value_heads: 4,
head_dim: 256,
rope_theta: 10_000_000.0,
partial_rotary_factor: 0.25,
rope_parameters: None,
linear_num_key_heads: 16,
linear_num_value_heads: Some(48),
linear_key_head_dim: 128,
linear_value_head_dim: 128,
linear_conv_kernel_dim: 4,
num_experts: None,
num_experts_per_tok: None,
moe_intermediate_size: None,
shared_expert_intermediate_size: None,
output_router_logits: false,
router_aux_loss_coef: None,
tie_word_embeddings: false,
full_attention_interval,
layer_types,
layer_mask: vec![true; num_hidden_layers],
eos_token_id: 248_044,
max_position_embeddings: 262_144,
mtp_num_hidden_layers: 1,
mtp_use_dedicated_embeddings: false,
quarot_rotation_seed: None,
}
}
pub fn from_config_json(path: &Path) -> Result<Self, InferenceError> {
let json = std::fs::read_to_string(path).map_err(InferenceError::Io)?;
Self::from_config_json_str(&json)
}
pub fn from_config_json_str(json: &str) -> Result<Self, InferenceError> {
let parsed: HfQwenConfigFile = serde_json::from_str(json)
.map_err(|e| InferenceError::Inference(format!("invalid Qwen config.json: {e}")))?;
let mut cfg = parsed
.text_config
.unwrap_or_else(Qwen35Config::qwen36_35b_a3b);
if let Some(tie) = parsed.tie_word_embeddings {
cfg.tie_word_embeddings = tie;
}
if let Some(rp) = &cfg.rope_parameters {
if cfg.rope_theta == 0.0 && rp.rope_theta > 0.0 {
cfg.rope_theta = rp.rope_theta;
}
if let Some(prf) = rp.partial_rotary_factor {
cfg.partial_rotary_factor = prf;
}
}
if cfg.layer_types.len() != cfg.num_hidden_layers {
if cfg.full_attention_interval == 0 {
return Err(InferenceError::Inference(
"invalid Qwen config.json: full_attention_interval must be > 0 when \
layer_types is absent or its length differs from num_hidden_layers"
.to_string(),
));
}
cfg.layer_types =
compute_layer_types(cfg.num_hidden_layers, cfg.full_attention_interval);
}
cfg.normalize_layer_mask();
if cfg.num_attention_heads == 0 {
return Err(InferenceError::Inference(
"invalid Qwen config.json: num_attention_heads must be > 0".to_string(),
));
}
if cfg.num_key_value_heads == 0 {
return Err(InferenceError::Inference(
"invalid Qwen config.json: num_key_value_heads must be > 0".to_string(),
));
}
if !cfg
.num_attention_heads
.is_multiple_of(cfg.num_key_value_heads)
{
return Err(InferenceError::Inference(format!(
"invalid Qwen config.json: num_attention_heads ({}) must be divisible by \
num_key_value_heads ({})",
cfg.num_attention_heads, cfg.num_key_value_heads
)));
}
if cfg.head_dim == 0 {
return Err(InferenceError::Inference(
"invalid Qwen config.json: head_dim must be > 0".to_string(),
));
}
if cfg.num_hidden_layers == 0 {
return Err(InferenceError::Inference(
"invalid Qwen config.json: num_hidden_layers must be > 0".to_string(),
));
}
if cfg.linear_conv_kernel_dim == 0 {
return Err(InferenceError::Inference(
"invalid Qwen config.json: linear_conv_kernel_dim must be > 0".to_string(),
));
}
if !(cfg.partial_rotary_factor.is_finite()
&& cfg.partial_rotary_factor >= 0.0
&& cfg.partial_rotary_factor <= 1.0)
{
return Err(InferenceError::Inference(format!(
"invalid Qwen config.json: partial_rotary_factor ({}) must be in [0.0, 1.0]",
cfg.partial_rotary_factor
)));
}
let rope_dim = cfg.rope_dim();
if rope_dim < 2 || !rope_dim.is_multiple_of(2) || rope_dim > cfg.head_dim {
return Err(InferenceError::Inference(format!(
"invalid Qwen config.json: derived rope_dim ({rope_dim}) must be even, >= 2, \
and <= head_dim ({hd}) (partial_rotary_factor={prf})",
hd = cfg.head_dim,
prf = cfg.partial_rotary_factor,
)));
}
let key_heads = cfg.linear_num_key_heads;
let value_heads = cfg.linear_num_value_heads();
if key_heads == 0 {
return Err(InferenceError::Inference(
"invalid Qwen config.json: linear_num_key_heads must be > 0".to_string(),
));
}
if value_heads == 0 || !value_heads.is_multiple_of(key_heads) {
return Err(InferenceError::Inference(format!(
"invalid Qwen config.json: linear_num_value_heads ({value_heads}) must be a \
positive multiple of linear_num_key_heads ({key_heads})"
)));
}
Ok(cfg)
}
pub fn linear_num_value_heads(&self) -> usize {
self.linear_num_value_heads.unwrap_or(32)
}
pub fn is_moe(&self) -> bool {
self.num_experts.is_some()
|| self.num_experts_per_tok.is_some()
|| self.moe_intermediate_size.is_some()
|| self.shared_expert_intermediate_size.is_some()
}
pub fn moe_intermediate_size(&self) -> usize {
self.moe_intermediate_size.unwrap_or(self.intermediate_size)
}
pub fn shared_expert_intermediate_size(&self) -> usize {
self.shared_expert_intermediate_size
.unwrap_or_else(|| self.moe_intermediate_size())
}
pub fn num_full_attention_layers(&self) -> usize {
self.layer_types
.iter()
.filter(|t| **t == LayerType::FullAttention)
.count()
}
pub fn num_linear_attention_layers(&self) -> usize {
self.layer_types
.iter()
.filter(|t| **t == LayerType::LinearAttention)
.count()
}
pub fn full_q_dim(&self) -> usize {
self.num_attention_heads * self.head_dim
}
pub fn full_kv_dim(&self) -> usize {
self.num_key_value_heads * self.head_dim
}
pub fn kv_cache_layer_count(&self) -> usize {
self.num_full_attention_layers()
}
pub fn kv_bytes_per_token(&self, dtype_bytes: usize) -> usize {
self.num_full_attention_layers() * 2 * self.full_kv_dim() * dtype_bytes
}
pub fn rope_dim(&self) -> usize {
(self.head_dim as f32 * self.partial_rotary_factor) as usize
}
pub fn linear_qkv_dim(&self) -> usize {
self.linear_num_key_heads * self.linear_key_head_dim + self.linear_num_key_heads * self.linear_key_head_dim + self.linear_num_value_heads() * self.linear_value_head_dim }
pub fn linear_output_dim(&self) -> usize {
self.linear_num_value_heads() * self.linear_value_head_dim
}
pub fn is_full_attention(&self, layer_idx: usize) -> bool {
self.layer_types
.get(layer_idx)
.copied()
.unwrap_or(LayerType::LinearAttention)
== LayerType::FullAttention
}
fn normalize_layer_mask(&mut self) {
if self.layer_mask.len() != self.num_hidden_layers {
self.layer_mask = vec![true; self.num_hidden_layers];
}
}
pub fn is_layer_active(&self, layer_idx: usize) -> bool {
self.layer_mask.get(layer_idx).copied().unwrap_or(true)
}
pub fn num_active_layers(&self) -> usize {
(0..self.num_hidden_layers)
.filter(|&i| self.is_layer_active(i))
.count()
}
pub fn num_active_linear_attention_layers(&self) -> usize {
(0..self.num_hidden_layers)
.filter(|&i| self.is_layer_active(i) && !self.is_full_attention(i))
.count()
}
pub fn num_active_full_attention_layers(&self) -> usize {
(0..self.num_hidden_layers)
.filter(|&i| self.is_layer_active(i) && self.is_full_attention(i))
.count()
}
pub fn apply_layer_mask(&mut self, mask: Vec<bool>) {
assert_eq!(
mask.len(),
self.num_hidden_layers,
"layer_mask length {} does not match num_hidden_layers {}",
mask.len(),
self.num_hidden_layers
);
assert!(
mask.iter().any(|&active| active),
"layer_mask must keep at least one active layer"
);
self.layer_mask = mask;
}
pub fn pruned_config(&self, mask: Vec<bool>) -> Self {
let mut cfg = self.clone();
cfg.apply_layer_mask(mask);
cfg
}
}
#[derive(Clone)]
pub struct GenerateConfig {
pub max_new_tokens: usize,
pub temperature: f32,
pub top_k: usize,
pub top_p: f32,
pub repetition_penalty: f32,
pub seed: Option<u64>,
pub stop_token_ids: Vec<u32>,
pub enable_thinking: bool,
pub enable_mtp: Option<bool>,
pub grammar: Option<Arc<GrammarEngine>>,
pub stop_strings: Vec<String>,
pub reasoning_budget: Option<usize>,
pub logprobs: Option<usize>,
}
impl std::fmt::Debug for GenerateConfig {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("GenerateConfig")
.field("max_new_tokens", &self.max_new_tokens)
.field("temperature", &self.temperature)
.field("top_k", &self.top_k)
.field("top_p", &self.top_p)
.field("repetition_penalty", &self.repetition_penalty)
.field("seed", &self.seed)
.field("stop_token_ids", &self.stop_token_ids)
.field("enable_thinking", &self.enable_thinking)
.field("enable_mtp", &self.enable_mtp)
.field("grammar", &self.grammar.as_ref().map(|_| "<GrammarEngine>"))
.field("stop_strings", &self.stop_strings)
.field("reasoning_budget", &self.reasoning_budget)
.field("logprobs", &self.logprobs)
.finish()
}
}
impl Default for GenerateConfig {
fn default() -> Self {
Self {
max_new_tokens: 256,
temperature: 0.7,
top_k: 50,
top_p: 0.9,
repetition_penalty: 1.1,
seed: None,
stop_token_ids: vec![QWEN_CHAT_IM_END_TOKEN_ID],
enable_thinking: true,
enable_mtp: None,
grammar: None,
stop_strings: vec![],
reasoning_budget: None,
logprobs: None,
}
}
}
#[inline]
pub(crate) fn force_close_think(
reasoning_budget: Option<usize>,
enable_thinking: bool,
thinking_closed: bool,
generated_so_far: usize,
close_id: Option<u32>,
) -> Option<u32> {
let budget = reasoning_budget?;
let close = close_id?;
if enable_thinking && !thinking_closed && budget > 0 && generated_so_far >= budget {
Some(close)
} else {
None
}
}
#[inline]
pub(crate) fn decode_cap(reasoning_budget: Option<usize>, max_new_tokens: usize) -> usize {
match reasoning_budget {
Some(rb) if rb > 0 => rb.saturating_add(max_new_tokens).saturating_add(1),
_ => max_new_tokens,
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct TopLogprob {
pub token_id: u32,
pub logprob: f32,
}
#[derive(Debug, Clone)]
pub struct TokenLogprob {
pub token_id: u32,
pub logprob: f32,
pub top: Vec<TopLogprob>,
}
#[derive(Debug, Clone)]
pub struct GenerateOutput {
pub text: String,
pub token_ids: Vec<u32>,
pub prompt_tokens: usize,
pub generated_tokens: usize,
pub stopped: bool,
pub stop_reason: Option<StopReason>,
pub token_logprobs: Vec<TokenLogprob>,
}
pub(crate) fn compute_layer_types(num_layers: usize, interval: usize) -> Vec<LayerType> {
(0..num_layers)
.map(|i| {
if (i + 1) % interval == 0 {
LayerType::FullAttention
} else {
LayerType::LinearAttention
}
})
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_config_construction_and_layer_types() {
let cfg = Qwen35Config::qwen35_2b();
assert_eq!(cfg.num_hidden_layers, 24);
assert_eq!(cfg.hidden_size, 2048);
assert_eq!(cfg.vocab_size, 248_320);
assert_eq!(cfg.layer_types.len(), 24);
assert_eq!(cfg.num_full_attention_layers(), 6);
assert_eq!(cfg.num_linear_attention_layers(), 18);
let full_indices: Vec<usize> = cfg
.layer_types
.iter()
.enumerate()
.filter(|(_, t)| **t == LayerType::FullAttention)
.map(|(i, _)| i)
.collect();
assert_eq!(full_indices, vec![3, 7, 11, 15, 19, 23]);
for i in [
0, 1, 2, 4, 5, 6, 8, 9, 10, 12, 13, 14, 16, 17, 18, 20, 21, 22,
] {
assert_eq!(cfg.layer_types[i], LayerType::LinearAttention);
}
}
#[test]
fn test_generate_config_defaults() {
let gen_cfg = GenerateConfig::default();
assert_eq!(gen_cfg.max_new_tokens, 256);
assert!((gen_cfg.temperature - 0.7).abs() < 1e-6);
assert_eq!(gen_cfg.top_k, 50);
assert!((gen_cfg.top_p - 0.9).abs() < 1e-6);
assert!((gen_cfg.repetition_penalty - 1.1).abs() < 1e-6);
assert!(
gen_cfg.stop_token_ids.contains(&QWEN_CHAT_IM_END_TOKEN_ID),
"default stop tokens must include im_end"
);
}
#[test]
fn test_dimension_helpers() {
let cfg = Qwen35Config::qwen35_2b();
assert_eq!(cfg.full_q_dim(), 8 * 256); assert_eq!(cfg.full_kv_dim(), 2 * 256); assert_eq!(cfg.rope_dim(), 64);
assert_eq!(cfg.linear_qkv_dim(), 6144);
assert_eq!(cfg.linear_output_dim(), 2048); }
#[test]
fn test_is_full_attention() {
let cfg = Qwen35Config::qwen35_2b();
assert!(!cfg.is_full_attention(0));
assert!(!cfg.is_full_attention(1));
assert!(!cfg.is_full_attention(2));
assert!(cfg.is_full_attention(3));
assert!(!cfg.is_full_attention(4));
assert!(cfg.is_full_attention(7));
assert!(cfg.is_full_attention(23));
assert!(!cfg.is_full_attention(100));
}
#[test]
fn test_qwen36_hf_config_fixture_parse_fields() {
let json = include_str!(concat!(
env!("CARGO_MANIFEST_DIR"),
"/tests/fixtures/qwen36_config.json"
));
let cfg = Qwen35Config::from_config_json_str(json).expect("Qwen3.6 HF config parses");
assert_eq!(cfg.hidden_size, 2048);
assert_eq!(cfg.num_attention_heads, 16);
assert_eq!(cfg.num_key_value_heads, 2);
assert_eq!(cfg.linear_num_value_heads(), 32);
assert_eq!(cfg.num_hidden_layers, 40);
assert_eq!(cfg.num_experts, Some(256));
assert_eq!(cfg.num_experts_per_tok, Some(8));
assert_eq!(cfg.moe_intermediate_size, Some(512));
assert_eq!(cfg.shared_expert_intermediate_size, Some(512));
assert!(!cfg.tie_word_embeddings);
assert!(cfg.is_moe());
assert_eq!(cfg.vocab_size, 248320);
assert_eq!(cfg.head_dim, 256);
assert_eq!(cfg.intermediate_size, 6144);
assert_eq!(cfg.linear_num_key_heads, 16);
assert_eq!(cfg.linear_key_head_dim, 128);
assert_eq!(cfg.linear_value_head_dim, 128);
assert_eq!(cfg.linear_conv_kernel_dim, 4);
assert_eq!(cfg.full_attention_interval, 4);
assert_eq!(cfg.eos_token_id, 248044_u32);
assert_eq!(cfg.max_position_embeddings, 262144);
assert_eq!(cfg.rope_theta, 10_000_000.0_f64);
assert_eq!(cfg.partial_rotary_factor, 0.25_f32);
let full_indices: Vec<usize> = cfg
.layer_types
.iter()
.enumerate()
.filter(|(_, t)| **t == LayerType::FullAttention)
.map(|(i, _)| i)
.collect();
assert_eq!(full_indices, vec![3, 7, 11, 15, 19, 23, 27, 31, 35, 39]);
}
#[test]
fn test_qwen36_27b_preset_dimensions() {
let cfg = Qwen35Config::qwen36_27b();
assert_eq!(cfg.hidden_size, 5120);
assert_eq!(cfg.num_hidden_layers, 64);
assert_eq!(cfg.vocab_size, 248_320);
assert_eq!(cfg.intermediate_size, 17408);
assert_eq!(cfg.num_attention_heads, 24);
assert_eq!(cfg.num_key_value_heads, 4);
assert_eq!(cfg.head_dim, 256);
assert!((cfg.rope_theta - 10_000_000.0_f64).abs() < 1.0);
assert!((cfg.partial_rotary_factor - 0.25_f32).abs() < 1e-6);
assert_eq!(cfg.eos_token_id, 248_044_u32);
assert_eq!(cfg.max_position_embeddings, 262_144);
}
#[test]
fn test_qwen36_27b_preset_layer_types() {
let cfg = Qwen35Config::qwen36_27b();
assert_eq!(cfg.layer_types.len(), 64);
assert_eq!(cfg.full_attention_interval, 4);
assert_eq!(
cfg.layer_types
.iter()
.filter(|t| **t == LayerType::LinearAttention)
.count(),
48
);
assert_eq!(
cfg.layer_types
.iter()
.filter(|t| **t == LayerType::FullAttention)
.count(),
16
);
for i in 0..64_usize {
let expected = (i + 1) % 4 == 0;
assert_eq!(
cfg.layer_types[i] == LayerType::FullAttention,
expected,
"layer {i} type mismatch"
);
}
}
#[test]
fn test_qwen36_27b_preset_not_moe() {
let cfg = Qwen35Config::qwen36_27b();
assert!(!cfg.is_moe());
assert!(cfg.num_experts.is_none());
assert!(cfg.num_experts_per_tok.is_none());
assert!(cfg.moe_intermediate_size.is_none());
assert!(cfg.shared_expert_intermediate_size.is_none());
assert!(cfg.router_aux_loss_coef.is_none());
assert!(!cfg.output_router_logits);
}
#[test]
fn test_qwen36_27b_preset_gdn_fields() {
let cfg = Qwen35Config::qwen36_27b();
assert_eq!(cfg.linear_num_key_heads, 16);
assert_eq!(cfg.linear_num_value_heads, Some(48));
assert_eq!(cfg.linear_num_value_heads(), 48);
assert_eq!(cfg.linear_key_head_dim, 128);
assert_eq!(cfg.linear_value_head_dim, 128);
assert_eq!(cfg.linear_conv_kernel_dim, 4);
assert!(!cfg.tie_word_embeddings);
assert_eq!(cfg.mtp_num_hidden_layers, 1);
assert!(!cfg.mtp_use_dedicated_embeddings);
}
#[test]
fn test_qwen36_27b_from_config_json() {
let home = std::env::var("HOME").unwrap_or_else(|_| "/root".to_string());
let path =
std::path::PathBuf::from(format!("{home}/.lattice/models/qwen3.6-27b/config.json"));
if !path.exists() {
return; }
let cfg = Qwen35Config::from_config_json(&path).expect("27B config.json parses");
assert_eq!(cfg.hidden_size, 5120);
assert_eq!(cfg.num_hidden_layers, 64);
assert_eq!(cfg.vocab_size, 248_320);
assert_eq!(cfg.intermediate_size, 17408);
assert_eq!(cfg.num_attention_heads, 24);
assert_eq!(cfg.num_key_value_heads, 4);
assert_eq!(cfg.head_dim, 256);
assert!(!cfg.tie_word_embeddings);
assert_eq!(cfg.layer_types.len(), 64);
assert!(!cfg.is_moe());
assert_eq!(cfg.mtp_num_hidden_layers, 1);
assert!(
(cfg.rope_theta - 10_000_000.0_f64).abs() < 1.0,
"rope_theta should be extracted from nested rope_parameters"
);
}
#[test]
fn test_layer_mask_default_all_true_27b() {
let cfg = Qwen35Config::qwen36_27b();
assert_eq!(cfg.layer_mask.len(), 64);
assert!(cfg.layer_mask.iter().all(|&active| active));
assert_eq!(cfg.num_active_layers(), 64);
assert_eq!(cfg.num_active_linear_attention_layers(), 48);
assert_eq!(cfg.num_active_full_attention_layers(), 16);
}
#[test]
fn test_num_active_layers_partial_mask() {
let mut cfg = Qwen35Config::qwen36_27b();
let mut mask = vec![true; 64];
mask[0] = false;
mask[3] = false;
mask[4] = false;
cfg.apply_layer_mask(mask);
assert_eq!(cfg.num_active_layers(), 61);
assert_eq!(cfg.num_active_linear_attention_layers(), 46);
assert_eq!(cfg.num_active_full_attention_layers(), 15);
}
#[test]
#[should_panic(expected = "layer_mask length")]
fn test_apply_layer_mask_wrong_length_panics() {
let mut cfg = Qwen35Config::qwen36_27b();
cfg.apply_layer_mask(vec![true; 32]);
}
#[test]
fn test_pruned_config_preserves_fields() {
let cfg = Qwen35Config::qwen36_27b();
let mut mask = vec![true; 64];
mask[5] = false;
mask[10] = false;
let pruned = cfg.pruned_config(mask.clone());
assert_eq!(pruned.hidden_size, cfg.hidden_size);
assert_eq!(pruned.num_hidden_layers, cfg.num_hidden_layers);
assert_eq!(pruned.vocab_size, cfg.vocab_size);
assert_eq!(pruned.layer_types, cfg.layer_types);
assert_eq!(pruned.layer_mask, mask);
assert_eq!(pruned.num_active_layers(), 62);
}
#[test]
#[should_panic(expected = "at least one active layer")]
fn test_apply_layer_mask_all_false_panics() {
let mut cfg = Qwen35Config::qwen36_27b();
cfg.apply_layer_mask(vec![false; 64]);
}
#[test]
fn test_layer_mask_normalizes_on_parse() {
let json = r#"{
"text_config": {
"hidden_size": 2048,
"num_hidden_layers": 4,
"full_attention_interval": 4,
"eos_token_id": 1
}
}"#;
let cfg = Qwen35Config::from_config_json_str(json).unwrap();
assert_eq!(
cfg.layer_mask.len(),
4,
"normalize_layer_mask must fill mask to num_hidden_layers"
);
assert!(
cfg.layer_mask.iter().all(|&v| v),
"normalized mask must be all-true"
);
}
#[test]
fn test_zero_full_attention_interval_errors_not_panics() {
let json = r#"{
"text_config": {
"hidden_size": 2048,
"num_hidden_layers": 4,
"full_attention_interval": 0,
"eos_token_id": 1
}
}"#;
let result = Qwen35Config::from_config_json_str(json);
assert!(
result.is_err(),
"full_attention_interval: 0 must yield an InferenceError, not panic"
);
}
#[test]
fn test_zero_num_key_value_heads_errors_not_panics() {
let json = r#"{"text_config": {"num_key_value_heads": 0}}"#;
let err = Qwen35Config::from_config_json_str(json)
.expect_err("num_key_value_heads: 0 must yield an InferenceError, not a panic")
.to_string();
assert!(
err.contains("num_key_value_heads"),
"wrong guard fired: {err}"
);
}
#[test]
fn test_indivisible_head_counts_error_not_panics() {
let json = r#"{"text_config": {"num_attention_heads": 3, "num_key_value_heads": 2}}"#;
let err = Qwen35Config::from_config_json_str(json)
.expect_err("indivisible head counts must yield an InferenceError, not OOB/panic")
.to_string();
assert!(err.contains("divisible"), "wrong guard fired: {err}");
}
#[test]
fn test_zero_head_dim_errors() {
let json = r#"{"text_config": {"head_dim": 0}}"#;
let err = Qwen35Config::from_config_json_str(json)
.expect_err("head_dim: 0 must yield an InferenceError")
.to_string();
assert!(err.contains("head_dim"), "wrong guard fired: {err}");
}
#[test]
fn test_zero_num_hidden_layers_errors() {
let json = r#"{"text_config": {"num_hidden_layers": 0}}"#;
let err = Qwen35Config::from_config_json_str(json)
.expect_err("num_hidden_layers: 0 must yield an InferenceError")
.to_string();
assert!(
err.contains("num_hidden_layers"),
"wrong guard fired: {err}"
);
}
#[test]
fn test_zero_linear_conv_kernel_dim_errors() {
let json = r#"{"text_config": {"linear_conv_kernel_dim": 0}}"#;
let err = Qwen35Config::from_config_json_str(json)
.expect_err("linear_conv_kernel_dim: 0 must yield an InferenceError, not underflow")
.to_string();
assert!(
err.contains("linear_conv_kernel_dim"),
"wrong guard fired: {err}"
);
}
#[test]
fn test_zero_linear_num_key_heads_errors() {
let json = r#"{"text_config": {"linear_num_key_heads": 0}}"#;
let err = Qwen35Config::from_config_json_str(json)
.expect_err("linear_num_key_heads: 0 must yield an InferenceError, not divide-by-zero")
.to_string();
assert!(
err.contains("linear_num_key_heads"),
"wrong guard fired: {err}"
);
}
#[test]
fn test_value_heads_below_key_heads_errors() {
let json = r#"{"text_config": {"linear_num_value_heads": 1}}"#;
let err = Qwen35Config::from_config_json_str(json)
.expect_err("value_heads < key_heads must yield an InferenceError, not divide-by-zero")
.to_string();
assert!(
err.contains("linear_num_value_heads"),
"wrong guard fired: {err}"
);
}
#[test]
fn test_value_heads_multiple_of_key_heads_accepted() {
let json = r#"{"text_config": {"linear_num_key_heads": 16, "linear_num_value_heads": 32}}"#;
assert!(
Qwen35Config::from_config_json_str(json).is_ok(),
"key 16 / value 32 (ratio 2) is a real GDN config and must be accepted"
);
}
#[test]
fn test_partial_rotary_factor_above_one_errors() {
let json = r#"{"text_config": {"partial_rotary_factor": 3.0}}"#;
let err = Qwen35Config::from_config_json_str(json)
.expect_err("partial_rotary_factor > 1.0 must yield an InferenceError, not OOB")
.to_string();
assert!(
err.contains("partial_rotary_factor"),
"wrong guard fired: {err}"
);
}
#[test]
fn test_partial_rotary_factor_one_accepted() {
let json = r#"{"text_config": {"partial_rotary_factor": 1.0}}"#;
assert!(
Qwen35Config::from_config_json_str(json).is_ok(),
"partial_rotary_factor == 1.0 (full rotary) must be accepted"
);
}
#[test]
fn test_odd_rope_dim_errors_not_panics() {
let json = r#"{"text_config": {"head_dim": 10, "partial_rotary_factor": 0.3}}"#;
let err = Qwen35Config::from_config_json_str(json)
.expect_err("odd rope_dim must yield an InferenceError, not silent wrong output")
.to_string();
assert!(err.contains("rope_dim"), "wrong guard fired: {err}");
}
#[test]
fn test_zero_rope_dim_errors_not_panics() {
let json = r#"{"text_config": {"partial_rotary_factor": 0.0}}"#;
let err = Qwen35Config::from_config_json_str(json)
.expect_err("zero rope_dim must yield an InferenceError, not capacity-zero surprise")
.to_string();
assert!(err.contains("rope_dim"), "wrong guard fired: {err}");
}
#[test]
fn test_rope_dim_exceeds_head_dim_via_f32_rounding_errors() {
let json = r#"{"text_config": {"head_dim": 16777219, "partial_rotary_factor": 1.0}}"#;
let err = Qwen35Config::from_config_json_str(json)
.expect_err("rope_dim > head_dim (f32 rounding) must yield an InferenceError, not OOB")
.to_string();
assert!(err.contains("rope_dim"), "wrong guard fired: {err}");
}
#[test]
fn test_generate_config_enable_thinking_default_and_toggle() {
let default_cfg = GenerateConfig::default();
assert!(
default_cfg.enable_thinking,
"default must have thinking enabled"
);
let no_think = GenerateConfig {
enable_thinking: false,
..GenerateConfig::default()
};
assert!(!no_think.enable_thinking);
assert_eq!(no_think.max_new_tokens, 256);
assert!((no_think.temperature - 0.7).abs() < 1e-6);
}
#[test]
fn test_qwen36_27b_layer_distribution_matches_config_json() {
let types = compute_layer_types(64, 4);
for chunk_start in (0..64).step_by(4) {
assert_eq!(types[chunk_start], LayerType::LinearAttention);
assert_eq!(types[chunk_start + 1], LayerType::LinearAttention);
assert_eq!(types[chunk_start + 2], LayerType::LinearAttention);
assert_eq!(types[chunk_start + 3], LayerType::FullAttention);
}
}
#[test]
fn test_qwen35_config_backward_compat() {
let json = r#"{
"text_config": {
"hidden_size": 2048,
"num_hidden_layers": 24,
"vocab_size": 248320,
"num_attention_heads": 8,
"num_key_value_heads": 2,
"head_dim": 256,
"rms_norm_eps": 0.000001,
"intermediate_size": 6144,
"linear_num_key_heads": 16,
"linear_num_value_heads": 16,
"linear_key_head_dim": 128,
"linear_value_head_dim": 128,
"linear_conv_kernel_dim": 4,
"full_attention_interval": 4,
"eos_token_id": 248044,
"max_position_embeddings": 262144,
"rope_theta": 10000000.0,
"partial_rotary_factor": 0.25
}
}"#;
let cfg = Qwen35Config::from_config_json_str(json).expect("backward-compat config parses");
assert_eq!(cfg.num_hidden_layers, 24);
assert_eq!(cfg.linear_num_value_heads(), 16);
assert!(!cfg.is_moe(), "Qwen3.5 must not be detected as MoE");
assert!(
cfg.tie_word_embeddings,
"default tie_word_embeddings is true"
);
assert_eq!(cfg.num_experts, None);
assert_eq!(cfg.num_experts_per_tok, None);
assert_eq!(
cfg.mtp_num_hidden_layers, 0,
"Qwen3.5 mtp_num_hidden_layers must default to 0"
);
assert!(!cfg.mtp_use_dedicated_embeddings);
}
#[test]
fn qwen35_config_missing_max_position_embeddings_defaults_4096() {
let json = r#"{
"text_config": {
"hidden_size": 2048,
"num_hidden_layers": 24,
"vocab_size": 248320,
"num_attention_heads": 8,
"num_key_value_heads": 2,
"head_dim": 256,
"rms_norm_eps": 0.000001,
"intermediate_size": 6144,
"linear_num_key_heads": 16,
"linear_num_value_heads": 16,
"linear_key_head_dim": 128,
"linear_value_head_dim": 128,
"linear_conv_kernel_dim": 4,
"full_attention_interval": 4,
"eos_token_id": 248044,
"rope_theta": 10000000.0,
"partial_rotary_factor": 0.25
}
}"#;
let cfg = Qwen35Config::from_config_json_str(json)
.expect("config without max_position_embeddings still parses");
assert_eq!(cfg.max_position_embeddings, 4096);
}
#[test]
fn test_qwen35_0_8b_preset_dimensions() {
let cfg = Qwen35Config::qwen35_0_8b();
assert_eq!(cfg.hidden_size, 1024);
assert_eq!(cfg.num_hidden_layers, 24);
assert_eq!(cfg.vocab_size, 248_320);
assert_eq!(cfg.intermediate_size, 3584);
assert_eq!(cfg.num_attention_heads, 8);
assert_eq!(cfg.num_key_value_heads, 2);
assert_eq!(cfg.head_dim, 256);
assert_eq!(cfg.rope_dim(), 64); assert_eq!(cfg.linear_num_key_heads, 16);
assert_eq!(cfg.linear_num_value_heads(), 16);
assert_eq!(cfg.linear_key_head_dim, 128);
assert_eq!(cfg.linear_value_head_dim, 128);
assert_eq!(cfg.eos_token_id, 248_044);
assert_eq!(cfg.max_position_embeddings, 262_144);
assert_eq!(cfg.mtp_num_hidden_layers, 1);
assert!(cfg.tie_word_embeddings);
assert!(!cfg.is_moe(), "Qwen3.5-0.8B is dense, not MoE");
assert_eq!(cfg.layer_types.len(), 24);
assert_eq!(cfg.num_full_attention_layers(), 6);
assert_eq!(cfg.num_linear_attention_layers(), 18);
}
#[test]
fn test_qwen35_0_8b_config_json_fixture_parses() {
let json = include_str!(concat!(
env!("CARGO_MANIFEST_DIR"),
"/tests/fixtures/qwen35_0_8b_config.json"
));
let cfg =
Qwen35Config::from_config_json_str(json).expect("Qwen3.5-0.8B config.json parses");
assert_eq!(cfg.hidden_size, 1024);
assert_eq!(cfg.num_hidden_layers, 24);
assert_eq!(cfg.vocab_size, 248_320);
assert_eq!(cfg.intermediate_size, 3584);
assert_eq!(cfg.num_attention_heads, 8);
assert_eq!(cfg.num_key_value_heads, 2);
assert_eq!(cfg.head_dim, 256);
assert_eq!(cfg.linear_num_key_heads, 16);
assert_eq!(cfg.linear_num_value_heads(), 16);
assert_eq!(cfg.linear_key_head_dim, 128);
assert_eq!(cfg.linear_value_head_dim, 128);
assert_eq!(cfg.linear_conv_kernel_dim, 4);
assert_eq!(cfg.full_attention_interval, 4);
assert_eq!(cfg.eos_token_id, 248_044);
assert_eq!(cfg.max_position_embeddings, 262_144);
assert_eq!(cfg.mtp_num_hidden_layers, 1);
assert!(!cfg.is_moe(), "Qwen3.5-0.8B is dense, not MoE");
assert_eq!(cfg.rope_theta, 10_000_000.0);
assert!((cfg.partial_rotary_factor - 0.25).abs() < 1e-6);
assert_eq!(cfg.rope_dim(), 64);
assert_eq!(cfg.layer_types.len(), 24);
assert_eq!(cfg.num_full_attention_layers(), 6);
assert_eq!(cfg.num_linear_attention_layers(), 18);
let full_indices: Vec<usize> = cfg
.layer_types
.iter()
.enumerate()
.filter(|(_, t)| **t == LayerType::FullAttention)
.map(|(i, _)| i)
.collect();
assert_eq!(full_indices, vec![3, 7, 11, 15, 19, 23]);
assert!(cfg.tie_word_embeddings);
}
#[test]
fn kv_layer_count_excludes_linear_layers() {
let cfg = Qwen35Config::qwen35_0_8b();
assert_eq!(
cfg.kv_cache_layer_count(),
6,
"must be full-attention count"
);
assert_ne!(
cfg.kv_cache_layer_count(),
cfg.num_hidden_layers,
"kv_cache_layer_count must not equal num_hidden_layers (would 4× decode memory)"
);
}
#[test]
fn full_plus_linear_equals_total() {
for (name, cfg) in [
("qwen35_0_8b", Qwen35Config::qwen35_0_8b()),
("qwen35_2b", Qwen35Config::qwen35_2b()),
("qwen36_35b_a3b", Qwen35Config::qwen36_35b_a3b()),
("qwen36_27b", Qwen35Config::qwen36_27b()),
] {
assert_eq!(
cfg.num_full_attention_layers() + cfg.num_linear_attention_layers(),
cfg.num_hidden_layers,
"{name}: full + linear must equal num_hidden_layers"
);
}
}
#[test]
fn kv_bytes_per_token_f16() {
let cfg = Qwen35Config::qwen35_0_8b();
assert_eq!(cfg.kv_bytes_per_token(2), 12_288);
}
#[test]
fn kv_bytes_per_token_identity() {
let cfg = Qwen35Config::qwen35_0_8b();
assert_eq!(cfg.kv_bytes_per_token(1), 6_144);
assert_eq!(
cfg.kv_bytes_per_token(1),
cfg.num_hidden_layers * cfg.head_dim
);
}
#[test]
fn decode_cap_none_budget_returns_max() {
assert_eq!(decode_cap(None, 512), 512);
assert_eq!(decode_cap(None, 0), 0);
}
#[test]
fn decode_cap_zero_budget_returns_max() {
assert_eq!(decode_cap(Some(0), 512), 512);
assert_eq!(decode_cap(Some(0), 1), 1);
}
#[test]
fn decode_cap_nonzero_budget_adds_budgets() {
assert_eq!(decode_cap(Some(2048), 512), 2561);
assert_eq!(decode_cap(Some(1), 1), 3);
assert_eq!(decode_cap(Some(100), 200), 301);
}
#[test]
fn decode_cap_saturates_on_overflow() {
assert_eq!(decode_cap(Some(usize::MAX), 1), usize::MAX);
assert_eq!(decode_cap(Some(1), usize::MAX), usize::MAX);
}
#[test]
fn force_close_think_disabled_when_budget_none() {
assert_eq!(
force_close_think(None, true, false, 100, Some(99)),
None,
"None budget must disable forcing"
);
}
#[test]
fn force_close_think_disabled_when_budget_zero() {
assert_eq!(
force_close_think(Some(0), true, false, 100, Some(99)),
None,
"budget=0 must disable forcing"
);
}
#[test]
fn force_close_think_disabled_when_enable_thinking_false() {
assert_eq!(
force_close_think(Some(10), false, false, 20, Some(99)),
None,
"enable_thinking=false must disable forcing"
);
}
#[test]
fn force_close_think_disabled_when_already_closed() {
assert_eq!(
force_close_think(Some(10), true, true, 20, Some(99)),
None,
"already-closed thinking block must not force again"
);
}
#[test]
fn force_close_think_disabled_when_close_id_none() {
assert_eq!(
force_close_think(Some(10), true, false, 20, None),
None,
"close_id=None must disable forcing"
);
}
#[test]
fn force_close_think_fires_at_budget_boundary() {
let close_id = 248_069_u32;
assert_eq!(
force_close_think(Some(10), true, false, 10, Some(close_id)),
Some(close_id),
"must force when generated_so_far equals budget"
);
assert_eq!(
force_close_think(Some(10), true, false, 11, Some(close_id)),
Some(close_id),
"must force when generated_so_far exceeds budget"
);
}
#[test]
fn force_close_think_does_not_fire_before_budget() {
let close_id = 248_069_u32;
assert_eq!(
force_close_think(Some(10), true, false, 9, Some(close_id)),
None,
"must not force when generated_so_far is one below budget"
);
assert_eq!(
force_close_think(Some(10), true, false, 0, Some(close_id)),
None,
"must not force when zero tokens generated"
);
}
}