1000 lines
38 KiB
Rust
1000 lines
38 KiB
Rust
//! Persistent module binding for the installed converted Qwen checkpoint.
|
|
//! Arrays/quantized packs have one owning bank; module views borrow that bank.
|
|
use super::array::{Array, Dtype};
|
|
use super::mlp::{Linear, QuantizedLinear};
|
|
use super::model::TextModel;
|
|
use super::stream::{Stream, Streams};
|
|
use super::{
|
|
decoder, gdn, hyper, model, moe, mtp, ngram, ngram_stage, ops, ple, qsa_attention, qsa_indexer,
|
|
rope, weights,
|
|
};
|
|
use std::collections::{BTreeMap, BTreeSet};
|
|
use std::path::Path;
|
|
|
|
#[derive(serde::Deserialize)]
|
|
#[serde(default)]
|
|
pub(super) struct Config {
|
|
hidden_size: i32,
|
|
hc_count: i32,
|
|
hc_lowrank: i32,
|
|
num_hidden_layers: usize,
|
|
layer_types: Option<Vec<String>>,
|
|
full_attention_interval: usize,
|
|
num_attention_heads: i32,
|
|
num_key_value_heads: i32,
|
|
head_dim: i32,
|
|
num_experts_per_tok: i32,
|
|
moe_intermediate_size: i32,
|
|
shared_expert_intermediate_size: i32,
|
|
norm_topk_prob: bool,
|
|
indexer_n_heads: i32,
|
|
indexer_kv_heads: i32,
|
|
indexer_head_dim: i32,
|
|
indexer_budget: i32,
|
|
indexer_compress_ratio: i32,
|
|
rms_norm_eps: f64,
|
|
ple_layer_ids: Vec<usize>,
|
|
ple_embed_dim: Option<i32>,
|
|
ngram_sidecar: bool,
|
|
eos_token_id: serde_json::Value,
|
|
tie_word_embeddings: bool,
|
|
partial_rotary_factor: f64,
|
|
rope_theta: f64,
|
|
rope_parameters: serde_json::Value,
|
|
mrope_section: Option<Vec<i32>>,
|
|
mrope_interleaved: bool,
|
|
}
|
|
|
|
impl Default for Config {
|
|
fn default() -> Self {
|
|
Self {
|
|
hidden_size: 2560,
|
|
hc_count: 4,
|
|
hc_lowrank: 320,
|
|
num_hidden_layers: 48,
|
|
layer_types: None,
|
|
full_attention_interval: 4,
|
|
num_attention_heads: 24,
|
|
num_key_value_heads: 2,
|
|
head_dim: 256,
|
|
num_experts_per_tok: 10,
|
|
moe_intermediate_size: 640,
|
|
shared_expert_intermediate_size: 640,
|
|
norm_topk_prob: true,
|
|
indexer_n_heads: 4,
|
|
indexer_kv_heads: 1,
|
|
indexer_head_dim: 128,
|
|
indexer_budget: 2048,
|
|
indexer_compress_ratio: 4,
|
|
rms_norm_eps: 1e-6,
|
|
ple_layer_ids: Vec::new(),
|
|
ple_embed_dim: None,
|
|
ngram_sidecar: false,
|
|
eos_token_id: serde_json::Value::Null,
|
|
tie_word_embeddings: false,
|
|
partial_rotary_factor: 0.25,
|
|
rope_theta: 10_000_000.,
|
|
rope_parameters: serde_json::Value::Null,
|
|
mrope_section: None,
|
|
mrope_interleaved: false,
|
|
}
|
|
}
|
|
}
|
|
|
|
impl Config {
|
|
fn parse(value: &serde_json::Value) -> Result<Self, String> {
|
|
let mut c: Self = serde_json::from_value(value.clone()).map_err(|e| e.to_string())?;
|
|
if c.hidden_size != 2560
|
|
|| c.hc_count != 4
|
|
|| c.hc_lowrank != 320
|
|
|| c.full_attention_interval == 0
|
|
|| c.num_hidden_layers == 0
|
|
|| c.indexer_compress_ratio <= 0
|
|
{
|
|
return Err("checkpoint does not have the installed Qwen geometry".into());
|
|
}
|
|
if c.layer_types.is_none() {
|
|
c.layer_types = Some(
|
|
(0..c.num_hidden_layers)
|
|
.map(|i| {
|
|
if (i + 1) % c.full_attention_interval == 0 {
|
|
"full_attention"
|
|
} else {
|
|
"linear_attention"
|
|
}
|
|
.to_owned()
|
|
})
|
|
.collect(),
|
|
);
|
|
}
|
|
if c.layer_types.as_ref().unwrap().len() != c.num_hidden_layers {
|
|
return Err("Qwen layer_types length differs from num_hidden_layers".into());
|
|
}
|
|
c.ple_layer_ids.sort_unstable();
|
|
c.ple_layer_ids.dedup();
|
|
c.ple_embed_dim.get_or_insert(c.hidden_size);
|
|
let p = &c.rope_parameters;
|
|
c.partial_rotary_factor = p["partial_rotary_factor"]
|
|
.as_f64()
|
|
.unwrap_or(c.partial_rotary_factor);
|
|
c.rope_theta = p["rope_theta"].as_f64().unwrap_or(c.rope_theta);
|
|
if let Some(section) = p["mrope_section"].as_array().filter(|s| !s.is_empty()) {
|
|
c.mrope_section = Some(
|
|
section
|
|
.iter()
|
|
.map(|v| {
|
|
v.as_i64()
|
|
.and_then(|v| i32::try_from(v).ok())
|
|
.ok_or("invalid mrope section".to_owned())
|
|
})
|
|
.collect::<Result<_, _>>()?,
|
|
);
|
|
}
|
|
c.mrope_interleaved = p["mrope_interleaved"]
|
|
.as_bool()
|
|
.unwrap_or(c.mrope_interleaved);
|
|
Ok(c)
|
|
}
|
|
|
|
fn input_width(&self, name: &str) -> i32 {
|
|
if name.ends_with("switch_mlp.down_proj") {
|
|
self.moe_intermediate_size
|
|
} else if name.ends_with("shared_expert.down_proj") {
|
|
self.shared_expert_intermediate_size
|
|
} else if name.ends_with("linear_attn.out_proj") {
|
|
48 * 128
|
|
} else if name.ends_with("self_attn.o_proj") {
|
|
self.num_attention_heads * self.head_dim
|
|
} else if name.ends_with("input_mix_weight_down") || name.ends_with("block_inject_weight") {
|
|
self.hidden_size * self.hc_count
|
|
} else if name.ends_with("input_mix_weight_up") {
|
|
self.hc_lowrank
|
|
} else if name.ends_with("ple.key_proj") || name.ends_with("ple.value_proj") {
|
|
self.ple_embed_dim.unwrap()
|
|
} else {
|
|
self.hidden_size
|
|
}
|
|
}
|
|
|
|
fn rope(&self, stream: Stream) -> Result<Rope, String> {
|
|
let p = &self.rope_parameters;
|
|
let kind = p["rope_type"]
|
|
.as_str()
|
|
.filter(|v| !v.is_empty())
|
|
.unwrap_or("default")
|
|
.trim()
|
|
.to_ascii_lowercase();
|
|
let yarn = match kind.as_str() {
|
|
"default" => None,
|
|
"yarn" => Some(rope::Yarn {
|
|
factor: p["factor"].as_f64().ok_or("YaRN factor missing")?,
|
|
original_max: p["original_max_position_embeddings"]
|
|
.as_i64()
|
|
.ok_or("YaRN original context missing")?,
|
|
attention_factor: p["attention_factor"].as_f64(),
|
|
mscale: p["mscale"].as_f64(),
|
|
mscale_all_dim: p["mscale_all_dim"].as_f64(),
|
|
beta_fast: p["beta_fast"].as_f64().unwrap_or(32.),
|
|
beta_slow: p["beta_slow"].as_f64().unwrap_or(1.),
|
|
truncate: p["truncate"].as_bool().unwrap_or(true),
|
|
}),
|
|
_ => return Err("unsupported Qwen RoPE type".into()),
|
|
};
|
|
let dim = (self.head_dim as f64 * self.partial_rotary_factor) as i32;
|
|
let (frequency, scaling) =
|
|
rope::inv_freq_and_scaling(dim, self.rope_theta, yarn.as_ref(), stream)?;
|
|
let axes = match &self.mrope_section {
|
|
Some(section) if section.iter().sum::<i32>() == dim / 2 => {
|
|
let axes = rope::build_mrope_axes(section, self.mrope_interleaved);
|
|
Some(Array::new(
|
|
&[axes.len() as i32],
|
|
Dtype::I32,
|
|
super::scalar_buffer(
|
|
&axes
|
|
.iter()
|
|
.flat_map(|v| v.to_le_bytes())
|
|
.collect::<Vec<_>>(),
|
|
)?,
|
|
)?)
|
|
}
|
|
_ => None,
|
|
};
|
|
Ok(Rope {
|
|
frequency,
|
|
scaling,
|
|
axes,
|
|
})
|
|
}
|
|
}
|
|
|
|
struct Rope {
|
|
frequency: Array,
|
|
scaling: f64,
|
|
axes: Option<Array>,
|
|
}
|
|
|
|
pub(super) struct Bank {
|
|
pub(super) parameters: weights::Parameters,
|
|
quantized: BTreeMap<String, QuantizedLinear>,
|
|
fusions: BTreeMap<String, weights::Fusion>,
|
|
rope: BTreeMap<String, Rope>,
|
|
config: Config,
|
|
has_mtp: bool,
|
|
}
|
|
|
|
// The executor can move this owner without rebuilding modules or losing their
|
|
// compile/staging caches. Reuse the already locked self_cell implementation:
|
|
// dependent modules are dropped before their weight bank, without local unsafe.
|
|
self_cell::self_cell!(
|
|
pub(super) struct BoundModel {
|
|
owner: Bank,
|
|
#[covariant]
|
|
dependent: TextModel,
|
|
}
|
|
);
|
|
|
|
impl Bank {
|
|
/// mx.eval(model.parameters()): one synchronous evaluation of the ordered
|
|
/// roots, not one eval per weight or a new alphabetical/shard grouping.
|
|
pub(super) fn materialize(
|
|
&self,
|
|
names: &[String],
|
|
streams: &Streams,
|
|
stream: Stream,
|
|
) -> Result<(), String> {
|
|
let roots = names
|
|
.iter()
|
|
.map(|name| {
|
|
self.parameters
|
|
.get(name)
|
|
.cloned()
|
|
.ok_or_else(|| format!("missing parameter {name}"))
|
|
})
|
|
.collect::<Result<Vec<_>, _>>()?;
|
|
ops::evaluate(streams, &roots, stream, false)
|
|
}
|
|
|
|
pub(super) fn load(
|
|
root: &Path,
|
|
options: weights::FusionOptions,
|
|
enable_mtp: bool,
|
|
streams: &Streams,
|
|
stream: Stream,
|
|
) -> Result<Self, String> {
|
|
let raw: serde_json::Value = serde_json::from_slice(
|
|
&std::fs::read(root.join("config.json")).map_err(|e| e.to_string())?,
|
|
)
|
|
.map_err(|e| e.to_string())?;
|
|
let config = Config::parse(&raw["text_config"])?;
|
|
let mut parameters = weights::load_shards(root, streams)?;
|
|
// This loader binds the converted checkpoint validated by DS4Server.
|
|
// Never silently apply converted norm conventions to raw HF weights.
|
|
if parameters.iter().any(|(k, v)| {
|
|
k.starts_with("model.language_model.")
|
|
|| (k.ends_with("linear_attn.conv1d.weight")
|
|
&& v.layout().shape().last() != Some(&1))
|
|
}) {
|
|
return Err("raw HF checkpoint requires raw sanitize before converted binding".into());
|
|
}
|
|
for name in parameters
|
|
.keys()
|
|
.filter(|k| k.ends_with("ple.conv1d.weight"))
|
|
.cloned()
|
|
.collect::<Vec<_>>()
|
|
{
|
|
let value = parameters.remove(&name).unwrap();
|
|
parameters.insert(name.replace("ple.conv1d.weight", "ple.conv_weight"), value);
|
|
}
|
|
let fusions = weights::fuse(
|
|
&mut parameters,
|
|
options,
|
|
&vec![config.indexer_n_heads > 0; config.num_hidden_layers],
|
|
stream,
|
|
)?;
|
|
let has_mtp = enable_mtp && root.join("mtp.safetensors").exists();
|
|
if has_mtp {
|
|
let head = super::load::safetensors(&root.join("mtp.safetensors"), streams)?;
|
|
parameters.extend(head.into_iter().filter(|(n, _)| n.starts_with("mtp.")));
|
|
for name in [
|
|
"mtp.pre_fc_norm_embedding.weight",
|
|
"mtp.pre_fc_norm_hidden.weight",
|
|
] {
|
|
if let Some(a) = parameters
|
|
.get(name)
|
|
.filter(|a| a.layout().shape().len() == 1)
|
|
{
|
|
let shifted = super::binary::binary(
|
|
&ops::astype(a, Dtype::F32, false, stream)?,
|
|
&Array::new(&[], Dtype::F32, super::scalar_buffer(&1f32.to_le_bytes())?)?,
|
|
super::binary::Binary::Add,
|
|
stream,
|
|
)?;
|
|
let dtype = a.layout().dtype();
|
|
let shifted = ops::astype(&shifted, dtype, false, stream)?;
|
|
parameters.insert(name.to_owned(), shifted);
|
|
}
|
|
}
|
|
}
|
|
let mut quantized = BTreeMap::new();
|
|
for (key, scales) in ¶meters {
|
|
let Some((name, weight_key, bias_key)) = key
|
|
.strip_suffix(".scales")
|
|
.map(|base| {
|
|
(
|
|
base.to_owned(),
|
|
format!("{base}.weight"),
|
|
format!("{base}.biases"),
|
|
)
|
|
})
|
|
.or_else(|| {
|
|
key.strip_suffix(".gu_scales").map(|base| {
|
|
(
|
|
format!("{base}.gu"),
|
|
format!("{base}.gu_weight"),
|
|
format!("{base}.gu_biases"),
|
|
)
|
|
})
|
|
})
|
|
else {
|
|
continue;
|
|
};
|
|
let weight = parameters
|
|
.get(&weight_key)
|
|
.ok_or_else(|| format!("missing {weight_key}"))?;
|
|
let biases = parameters
|
|
.get(&bias_key)
|
|
.ok_or_else(|| format!("missing {bias_key}"))?;
|
|
let recipe_name = name.strip_suffix(".gu").unwrap_or(&name);
|
|
let input_width = config.input_width(recipe_name);
|
|
let columns = weight.layout().dim(-1)?;
|
|
let groups = scales.layout().dim(-1)?;
|
|
if input_width <= 0 || columns <= 0 || groups <= 0 {
|
|
return Err(format!("invalid packed dimensions for {name}"));
|
|
}
|
|
let (bits, group) = if let Some(f) = fusions.get(recipe_name) {
|
|
(f.bits, f.group)
|
|
} else if name.starts_with("mtp.") {
|
|
(
|
|
u32::try_from(i64::from(columns) * 32 / i64::from(input_width))
|
|
.map_err(|e| e.to_string())?,
|
|
input_width as u32 / groups as u32,
|
|
)
|
|
} else {
|
|
let recipe = raw["quantization"]
|
|
.get(&name)
|
|
.unwrap_or(&raw["quantization"]);
|
|
if recipe
|
|
.get("mode")
|
|
.and_then(serde_json::Value::as_str)
|
|
.unwrap_or("affine")
|
|
!= "affine"
|
|
{
|
|
return Err(format!("unsupported quantization for {name}"));
|
|
}
|
|
(
|
|
recipe["bits"]
|
|
.as_u64()
|
|
.and_then(|n| u32::try_from(n).ok())
|
|
.ok_or_else(|| format!("missing bits for {name}"))?,
|
|
recipe["group_size"]
|
|
.as_u64()
|
|
.and_then(|n| u32::try_from(n).ok())
|
|
.ok_or_else(|| format!("missing group for {name}"))?,
|
|
)
|
|
};
|
|
if !matches!(bits, 2 | 3 | 4 | 5 | 6 | 8)
|
|
|| !matches!(group, 32 | 64 | 128)
|
|
|| weight.layout().dtype() != Dtype::U32
|
|
|| weight.layout().shape().len() < 2
|
|
|| scales.layout().shape() != biases.layout().shape()
|
|
|| i64::from(weight.layout().dim(-1)?) * 32
|
|
!= i64::from(input_width) * i64::from(bits)
|
|
|| i64::from(scales.layout().dim(-1)?) * i64::from(group) != i64::from(input_width)
|
|
{
|
|
return Err(format!("affine pack geometry mismatch for {name}"));
|
|
}
|
|
quantized.insert(
|
|
name,
|
|
QuantizedLinear {
|
|
weight: weight.clone(),
|
|
scales: scales.clone(),
|
|
biases: biases.clone(),
|
|
bits,
|
|
group,
|
|
},
|
|
);
|
|
}
|
|
let mut rope = BTreeMap::new();
|
|
let mut attention = config
|
|
.layer_types
|
|
.as_ref()
|
|
.unwrap()
|
|
.iter()
|
|
.enumerate()
|
|
.filter(|(_, t)| *t != "linear_attention")
|
|
.map(|(i, _)| format!("language_model.model.layers.{i}.self_attn"))
|
|
.collect::<Vec<_>>();
|
|
if has_mtp {
|
|
attention.push("mtp.layers.0.self_attn".into());
|
|
}
|
|
for name in attention {
|
|
rope.insert(name.clone(), config.rope(stream)?);
|
|
if config.indexer_n_heads > 0 {
|
|
rope.insert(format!("{name}.indexer"), config.rope(stream)?);
|
|
}
|
|
}
|
|
Ok(Self {
|
|
parameters,
|
|
quantized,
|
|
fusions,
|
|
rope,
|
|
config,
|
|
has_mtp,
|
|
})
|
|
}
|
|
|
|
pub(super) fn bind(
|
|
&self,
|
|
policy: model::CompilePolicy,
|
|
fused_moe_decode: bool,
|
|
fused_moe_verify: bool,
|
|
) -> Result<model::TextModel<'_>, String> {
|
|
let mut bind = Bind {
|
|
bank: self,
|
|
used: BTreeSet::new(),
|
|
order: Vec::new(),
|
|
fused_moe_decode,
|
|
fused_moe_verify,
|
|
};
|
|
let embedding = bind.linear("language_model.model.embed_tokens")?;
|
|
let layers = self
|
|
.config
|
|
.layer_types
|
|
.as_ref()
|
|
.unwrap()
|
|
.iter()
|
|
.enumerate()
|
|
.map(|(i, t)| {
|
|
bind.decoder(
|
|
&format!("language_model.model.layers.{i}"),
|
|
t == "linear_attention",
|
|
self.config.ple_layer_ids.contains(&(i + 1)),
|
|
)
|
|
})
|
|
.collect::<Result<Vec<_>, _>>()?;
|
|
let mixer = bind.hyper("language_model.model.hyper_connection_mixer", false)?;
|
|
let head = if self.config.tie_word_embeddings {
|
|
None
|
|
} else {
|
|
Some(bind.linear("language_model.lm_head")?)
|
|
};
|
|
let mut model = model::TextModel::new(embedding, layers, mixer, head, policy)?;
|
|
if self.has_mtp {
|
|
model.mtp = Some(mtp::Mtp {
|
|
norm_embedding: bind.array("mtp.pre_fc_norm_embedding.weight")?,
|
|
norm_hidden: bind.array("mtp.pre_fc_norm_hidden.weight")?,
|
|
fc_embedding: bind.linear("mtp.fc_embedding")?,
|
|
fc_hidden: bind.linear("mtp.fc_hidden")?,
|
|
eps: self.config.rms_norm_eps as f32,
|
|
layer: bind.decoder("mtp.layers.0", false, false)?,
|
|
mixer: bind.hyper("mtp.hyper_connection_mixer", false)?,
|
|
});
|
|
}
|
|
let unused = self
|
|
.parameters
|
|
.keys()
|
|
.filter(|n| !bind.used.contains(*n))
|
|
.cloned()
|
|
.collect::<Vec<_>>();
|
|
if !unused.is_empty() {
|
|
return Err(format!(
|
|
"unbound checkpoint parameters: {}",
|
|
unused.join(", ")
|
|
));
|
|
};
|
|
model.parameter_order = bind.order;
|
|
Ok(model)
|
|
}
|
|
}
|
|
|
|
struct Bind<'a> {
|
|
bank: &'a Bank,
|
|
used: BTreeSet<String>,
|
|
order: Vec<String>,
|
|
fused_moe_decode: bool,
|
|
fused_moe_verify: bool,
|
|
}
|
|
|
|
#[test]
|
|
#[ignore = "requires installed DS4_QWEN38_ARTIFACTS and SSD policy; only small norms and four table rows evaluated"]
|
|
fn mtplx_installed_module_binding_matches_reference() {
|
|
use super::super::gpu::Context;
|
|
use super::{allocator, configure_sources, stream::Device};
|
|
use sha2::{Digest, Sha256};
|
|
let root = std::path::PathBuf::from(
|
|
std::env::var_os("DS4_QWEN38_ARTIFACTS").expect("set installed artifact path"),
|
|
);
|
|
assert!(
|
|
!ngram::resident_policy(),
|
|
"this check must not load the resident table"
|
|
);
|
|
configure_sources().unwrap();
|
|
let _context = Context::open_qwen(0).unwrap();
|
|
let streams = Streams::new(Some(allocator::Allocator::new().unwrap()));
|
|
let gpu = streams.default_stream(Device::Gpu).unwrap();
|
|
let hash = |bytes: &[u8]| {
|
|
Sha256::digest(bytes)
|
|
.iter()
|
|
.map(|b| format!("{b:02x}"))
|
|
.collect::<String>()
|
|
};
|
|
let cases = include_str!("../../../../tests/fixtures/mtplx-binding.jsonl")
|
|
.lines()
|
|
.map(|s| serde_json::from_str::<serde_json::Value>(s).unwrap())
|
|
.collect::<Vec<_>>();
|
|
assert_eq!(cases.len(), 4);
|
|
let sidecars = include_str!("../../../../tests/fixtures/mtplx-binding-sidecar.jsonl")
|
|
.lines()
|
|
.map(|s| serde_json::from_str::<serde_json::Value>(s).unwrap())
|
|
.collect::<Vec<_>>();
|
|
assert_eq!(sidecars.len(), cases.len());
|
|
let orders = include_str!("../../../../tests/fixtures/mtplx-binding-order.jsonl")
|
|
.lines()
|
|
.map(|s| serde_json::from_str::<serde_json::Value>(s).unwrap())
|
|
.collect::<Vec<_>>();
|
|
assert_eq!(orders.len(), cases.len());
|
|
for ((case, sidecar_case), order) in cases.into_iter().zip(sidecars).zip(orders) {
|
|
assert_eq!(case["mask"], sidecar_case["mask"]);
|
|
assert_eq!(case["mtp"], sidecar_case["mtp"]);
|
|
let mask = case["mask"].as_u64().unwrap();
|
|
let mtp = case["mtp"].as_bool().unwrap();
|
|
eprintln!("installed binding mask={mask} mtp={mtp}");
|
|
let mut bank = Bank::load(
|
|
&root,
|
|
weights::FusionOptions {
|
|
gate_up: mask & 1 != 0,
|
|
gdn: mask & 2 != 0,
|
|
qsa: mask & 4 != 0,
|
|
},
|
|
mtp,
|
|
&streams,
|
|
gpu,
|
|
)
|
|
.unwrap();
|
|
assert!(
|
|
bank.parameters.values().all(|a| !a.has_data()),
|
|
"no eager parameter evaluation"
|
|
);
|
|
let mut model = bank
|
|
.bind(model::CompilePolicy::from_values(None, None), true, true)
|
|
.unwrap();
|
|
assert_eq!(case["mask"], order["mask"]);
|
|
assert_eq!(case["mtp"], order["mtp"]);
|
|
assert_eq!(
|
|
serde_json::json!(model.parameter_order.len()),
|
|
order["count"]
|
|
);
|
|
assert_eq!(
|
|
hash(&serde_json::to_vec(&model.parameter_order).unwrap()),
|
|
order["sha256"]
|
|
);
|
|
let linear = model
|
|
.layers
|
|
.iter()
|
|
.enumerate()
|
|
.filter(|(_, l)| matches!(l.attention, decoder::Attention::Linear(_)))
|
|
.map(|(i, _)| i)
|
|
.collect::<Vec<_>>();
|
|
let ple = model
|
|
.layers
|
|
.iter()
|
|
.enumerate()
|
|
.filter(|(_, l)| l.ple.is_some())
|
|
.map(|(i, _)| i)
|
|
.collect::<Vec<_>>();
|
|
assert_eq!(serde_json::json!(linear), case["linear"]);
|
|
assert_eq!(serde_json::json!(ple), case["ple"]);
|
|
assert_eq!(model.mtp.is_some(), mtp);
|
|
assert_eq!(model.layers.len(), 48);
|
|
assert_eq!(model.make_cache(4).len(), 48);
|
|
let metadata = bank
|
|
.parameters
|
|
.iter()
|
|
.map(|(name, a)| {
|
|
serde_json::json!([name, a.layout().shape(), a.layout().dtype().kernel_name()])
|
|
})
|
|
.collect::<Vec<_>>();
|
|
assert_eq!(
|
|
metadata.len() as u64,
|
|
case["parameter_count"].as_u64().unwrap()
|
|
);
|
|
assert_eq!(
|
|
hash(&serde_json::to_vec(&metadata).unwrap()),
|
|
case["metadata_sha256"]
|
|
);
|
|
let recipes = bank
|
|
.quantized
|
|
.iter()
|
|
.map(|(name, q)| {
|
|
(
|
|
name.clone(),
|
|
serde_json::json!([q.bits, q.group, q.weight.layout().shape()]),
|
|
)
|
|
})
|
|
.collect::<BTreeMap<_, _>>();
|
|
assert_eq!(serde_json::to_value(recipes).unwrap(), case["quantized"]);
|
|
if let Some(mtp) = &model.mtp {
|
|
// Exercise the same ordered-root loader on only the two small
|
|
// norms. Full-model materialization is intentionally not this test.
|
|
let norms = model
|
|
.parameter_order
|
|
.iter()
|
|
.filter(|name| name.starts_with("mtp.pre_fc_norm_"))
|
|
.cloned()
|
|
.collect::<Vec<_>>();
|
|
bank.materialize(&norms, &streams, gpu).unwrap();
|
|
for (name, norm) in [
|
|
("mtp.pre_fc_norm_hidden.weight", mtp.norm_hidden),
|
|
("mtp.pre_fc_norm_embedding.weight", mtp.norm_embedding),
|
|
] {
|
|
let out = ops::astype(norm, Dtype::F32, false, gpu).unwrap();
|
|
ops::evaluate(&streams, std::slice::from_ref(&out), gpu, false).unwrap();
|
|
let mut bytes = vec![0; out.layout().nbytes()];
|
|
out.buffer().read(out.offset(), &mut bytes).unwrap();
|
|
assert_eq!(hash(&bytes), case["norms"][name]);
|
|
}
|
|
}
|
|
model.post_weight_load(&root, &streams, gpu).unwrap();
|
|
let ids = Array::new(
|
|
&[4],
|
|
Dtype::I64,
|
|
super::scalar_buffer(
|
|
&[0i64, 1, 255, 1]
|
|
.iter()
|
|
.flat_map(|v| v.to_le_bytes())
|
|
.collect::<Vec<_>>(),
|
|
)
|
|
.unwrap(),
|
|
)
|
|
.unwrap();
|
|
let mut tables = Vec::new();
|
|
for (i, layer) in model.layers.iter_mut().enumerate() {
|
|
let Some(ple) = &mut layer.ple else { continue };
|
|
let table = &mut ple.embedding.table;
|
|
let out = table.gather(&ids, &streams, gpu).unwrap();
|
|
let out = ops::astype(&out, Dtype::F32, false, gpu).unwrap();
|
|
ops::evaluate(&streams, std::slice::from_ref(&out), gpu, false).unwrap();
|
|
let mut bytes = vec![0; out.layout().nbytes()];
|
|
out.buffer().read(out.offset(), &mut bytes).unwrap();
|
|
let sidecar = table.sidecar.as_ref().unwrap();
|
|
tables.push(
|
|
serde_json::json!({"layer":i, "resident":table.resident.is_some(),
|
|
"hot_mb":sidecar.hot_mebibytes(), "prefetch":sidecar.prefetch_enabled(),
|
|
"shape":out.layout().shape(), "sha256":hash(&bytes)}),
|
|
);
|
|
}
|
|
assert_eq!(serde_json::json!(tables), sidecar_case["tables"]);
|
|
assert_eq!(
|
|
model.set_ar_pipeline_mode(true),
|
|
sidecar_case["ar_ready"].as_bool().unwrap()
|
|
);
|
|
drop(model);
|
|
assert!(
|
|
bank.materialize(&["missing.weight".into()], &streams, gpu)
|
|
.unwrap_err()
|
|
.contains("missing parameter")
|
|
);
|
|
assert!(
|
|
bank.parameters
|
|
.iter()
|
|
.filter(|(n, _)| !n.starts_with("mtp.pre_fc_norm_"))
|
|
.all(|(_, a)| !a.has_data())
|
|
);
|
|
let mut owned = BoundModel::try_new(bank, |bank| {
|
|
bank.bind(model::CompilePolicy::from_values(None, None), true, true)
|
|
})
|
|
.unwrap();
|
|
let module_address = std::ptr::from_ref(owned.borrow_dependent()).addr();
|
|
owned.with_dependent_mut(|_, model| model.last_widened = Some(ids));
|
|
let mut moved = Box::new(owned);
|
|
assert_eq!(
|
|
module_address,
|
|
std::ptr::from_ref(moved.borrow_dependent()).addr()
|
|
);
|
|
moved.with_dependent_mut(|_, model| {
|
|
assert_eq!(model.last_widened.take().unwrap().layout().shape(), [4]);
|
|
});
|
|
bank = (*moved).into_owner();
|
|
// Strict binding rejects both missing and extraneous checkpoint entries.
|
|
let extra = bank.parameters["language_model.model.layers.0.linear_attn.A_log"].clone();
|
|
bank.parameters.insert("unexpected.weight".into(), extra);
|
|
assert!(
|
|
bank.bind(model::CompilePolicy::from_values(None, None), true, true)
|
|
.err()
|
|
.unwrap()
|
|
.contains("unbound checkpoint parameters")
|
|
);
|
|
bank.parameters.remove("unexpected.weight");
|
|
bank.parameters
|
|
.remove("language_model.model.layers.0.linear_attn.norm.weight");
|
|
assert!(
|
|
bank.bind(model::CompilePolicy::from_values(None, None), true, true)
|
|
.err()
|
|
.unwrap()
|
|
.contains("missing parameter")
|
|
);
|
|
eprintln!(
|
|
"binding metadata, module recipes, all-parameter coverage, MTP norms and SSD table lifecycle exact"
|
|
);
|
|
}
|
|
streams.clear_streams().unwrap();
|
|
}
|
|
|
|
impl<'a> Bind<'a> {
|
|
fn record(&mut self, name: &str) {
|
|
if self.used.insert(name.to_owned()) {
|
|
self.order.push(name.to_owned());
|
|
}
|
|
}
|
|
fn array(&mut self, name: &str) -> Result<&'a Array, String> {
|
|
let array = self
|
|
.bank
|
|
.parameters
|
|
.get(name)
|
|
.ok_or_else(|| format!("missing parameter {name}"))?;
|
|
self.record(name);
|
|
Ok(array)
|
|
}
|
|
fn optional(&mut self, name: &str) -> Result<Option<&'a Array>, String> {
|
|
if self.bank.parameters.contains_key(name) {
|
|
self.array(name).map(Some)
|
|
} else {
|
|
Ok(None)
|
|
}
|
|
}
|
|
fn quantized(&mut self, name: &str) -> Result<&'a QuantizedLinear, String> {
|
|
let layer = self
|
|
.bank
|
|
.quantized
|
|
.get(name)
|
|
.ok_or_else(|| format!("missing quantized module {name}"))?;
|
|
for part in ["weight", "scales", "biases"] {
|
|
let key = name.strip_suffix(".gu").map_or_else(
|
|
|| format!("{name}.{part}"),
|
|
|base| format!("{base}.gu_{part}"),
|
|
);
|
|
self.record(&key);
|
|
}
|
|
Ok(layer)
|
|
}
|
|
fn linear(&mut self, name: &str) -> Result<Linear<'a>, String> {
|
|
if self.bank.quantized.contains_key(name) {
|
|
self.quantized(name).map(Linear::Quantized)
|
|
} else {
|
|
self.array(&format!("{name}.weight")).map(Linear::Dense)
|
|
}
|
|
}
|
|
fn hyper(&mut self, p: &str, combine: bool) -> Result<hyper::GatedResidual<'a>, String> {
|
|
Ok(hyper::GatedResidual {
|
|
norm: self.array(&format!("{p}.hc_norm.weight"))?,
|
|
down: self.linear(&format!("{p}.input_mix_weight_down"))?,
|
|
up: self.linear(&format!("{p}.input_mix_weight_up"))?,
|
|
inject: if combine {
|
|
Some(self.linear(&format!("{p}.block_inject_weight"))?)
|
|
} else {
|
|
None
|
|
},
|
|
pack: None,
|
|
})
|
|
}
|
|
fn feed_forward(&mut self, p: &str) -> Result<moe::FeedForward, String> {
|
|
let mut down = None;
|
|
let input = if self.bank.quantized.contains_key(&format!("{p}.gu")) {
|
|
moe::GateUp::Fused(self.quantized(&format!("{p}.gu"))?.clone())
|
|
} else {
|
|
let gate = self.quantized(&format!("{p}.gate_proj"))?.clone();
|
|
// Qwen3NextMLP constructs gate/down/up; SwitchGLU gate/up/down.
|
|
if p.ends_with(".shared_expert") {
|
|
down = Some(self.quantized(&format!("{p}.down_proj"))?.clone());
|
|
}
|
|
moe::GateUp::Separate {
|
|
gate,
|
|
up: self.quantized(&format!("{p}.up_proj"))?.clone(),
|
|
}
|
|
};
|
|
Ok(moe::FeedForward {
|
|
input,
|
|
down: match down {
|
|
Some(down) => down,
|
|
None => self.quantized(&format!("{p}.down_proj"))?.clone(),
|
|
},
|
|
})
|
|
}
|
|
fn attention(&mut self, p: &str) -> Result<qsa_attention::Attention<'a>, String> {
|
|
let c = &self.bank.config;
|
|
let fused_name = format!("{p}.qkv_fused");
|
|
let input = if self.bank.fusions.contains_key(&fused_name) {
|
|
None // sanitize appends the fused module after the existing children.
|
|
} else {
|
|
let q = self.linear(&format!("{p}.q_proj"))?;
|
|
let qb = self.optional(&format!("{p}.q_proj.bias"))?;
|
|
let k = self.linear(&format!("{p}.k_proj"))?;
|
|
let kb = self.optional(&format!("{p}.k_proj.bias"))?;
|
|
let v = self.linear(&format!("{p}.v_proj"))?;
|
|
let vb = self.optional(&format!("{p}.v_proj.bias"))?;
|
|
Some(qsa_attention::Input::Separate([q, k, v], [qb, kb, vb]))
|
|
};
|
|
let output = self.linear(&format!("{p}.o_proj"))?;
|
|
let output_bias = self.optional(&format!("{p}.o_proj.bias"))?;
|
|
let q_norm = self.array(&format!("{p}.q_norm.weight"))?;
|
|
let k_norm = self.array(&format!("{p}.k_norm.weight"))?;
|
|
let r = self.bank.rope.get(p).ok_or("missing attention RoPE")?;
|
|
let indexer = if c.indexer_n_heads > 0 {
|
|
let ip = format!("{p}.indexer");
|
|
let r = self.bank.rope.get(&ip).ok_or("missing indexer RoPE")?;
|
|
Some(qsa_indexer::Indexer {
|
|
projection: if self
|
|
.bank
|
|
.fusions
|
|
.get(&fused_name)
|
|
.is_some_and(|f| f.splits.len() == 3)
|
|
{
|
|
None
|
|
} else {
|
|
Some(self.linear(&format!("{ip}.index_qk_proj"))?)
|
|
},
|
|
heads: c.indexer_n_heads,
|
|
kv_heads: c.indexer_kv_heads,
|
|
head_dim: c.indexer_head_dim,
|
|
block_topk: c.indexer_budget / c.indexer_compress_ratio,
|
|
ratio: c.indexer_compress_ratio,
|
|
q_norm: self.array(&format!("{ip}.q_layernorm.weight"))?,
|
|
k_norm: self.array(&format!("{ip}.k_layernorm.weight"))?,
|
|
inv_freq: &r.frequency,
|
|
eps: c.rms_norm_eps,
|
|
scaling: r.scaling,
|
|
scratch: 32 * 1024 * 1024,
|
|
compiled: None,
|
|
})
|
|
} else {
|
|
None
|
|
};
|
|
let input = match input {
|
|
Some(input) => input,
|
|
None => qsa_attention::Input::Fused(
|
|
self.linear(&fused_name)?,
|
|
&self.bank.fusions[&fused_name].splits,
|
|
),
|
|
};
|
|
Ok(qsa_attention::Attention {
|
|
input,
|
|
output,
|
|
output_bias,
|
|
heads: c.num_attention_heads,
|
|
kv_heads: c.num_key_value_heads,
|
|
dim: c.head_dim,
|
|
eps: c.rms_norm_eps as f32,
|
|
scale: (c.head_dim as f64).powf(-0.5),
|
|
q_norm,
|
|
k_norm,
|
|
inv_freq: &r.frequency,
|
|
rope_scaling: r.scaling,
|
|
mrope_axes: r.axes.as_ref(),
|
|
indexer,
|
|
})
|
|
}
|
|
fn decoder(
|
|
&mut self,
|
|
p: &str,
|
|
linear: bool,
|
|
has_ple: bool,
|
|
) -> Result<decoder::DecoderLayer<'a>, String> {
|
|
let c = &self.bank.config;
|
|
let attention = if linear {
|
|
let p = format!("{p}.linear_attn");
|
|
let conv = self.array(&format!("{p}.conv1d.weight"))?;
|
|
let f = self.bank.fusions.get(&format!("{p}.in_proj_fused"));
|
|
let input = if f.is_some() {
|
|
None
|
|
} else {
|
|
Some(gdn::InputProjections::Separate([
|
|
self.quantized(&format!("{p}.in_proj_qkv"))?,
|
|
self.quantized(&format!("{p}.in_proj_z"))?,
|
|
self.quantized(&format!("{p}.in_proj_b"))?,
|
|
self.quantized(&format!("{p}.in_proj_a"))?,
|
|
]))
|
|
};
|
|
let dt_bias = self.array(&format!("{p}.dt_bias"))?;
|
|
let a_log = self.array(&format!("{p}.A_log"))?;
|
|
let norm = self.array(&format!("{p}.norm.weight"))?;
|
|
let output = self.quantized(&format!("{p}.out_proj"))?;
|
|
let input = match input {
|
|
Some(input) => input,
|
|
None => gdn::InputProjections::Fused(
|
|
self.quantized(&format!("{p}.in_proj_fused"))?,
|
|
f.unwrap()
|
|
.splits
|
|
.as_slice()
|
|
.try_into()
|
|
.map_err(|_| "invalid GDN splits")?,
|
|
),
|
|
};
|
|
decoder::Attention::Linear(gdn::Weights {
|
|
input,
|
|
output,
|
|
conv,
|
|
a_log,
|
|
dt_bias,
|
|
norm,
|
|
})
|
|
} else {
|
|
decoder::Attention::Qsa(Box::new(self.attention(&format!("{p}.self_attn"))?))
|
|
};
|
|
let mlp = moe::SparseMoe {
|
|
router: self.quantized(&format!("{p}.mlp.gate"))?.clone(),
|
|
experts: self.feed_forward(&format!("{p}.mlp.switch_mlp"))?,
|
|
shared: self.feed_forward(&format!("{p}.mlp.shared_expert"))?,
|
|
shared_gate: self
|
|
.quantized(&format!("{p}.mlp.shared_expert_gate"))?
|
|
.clone(),
|
|
topk: c.num_experts_per_tok,
|
|
normalize: c.norm_topk_prob,
|
|
fused_decode: self.fused_moe_decode,
|
|
fused_verify: self.fused_moe_verify,
|
|
};
|
|
let attn_hyper = self.hyper(&format!("{p}.attn_hyper_connection"), true)?;
|
|
let mlp_hyper = self.hyper(&format!("{p}.mlp_hyper_connection"), true)?;
|
|
let ple = if has_ple {
|
|
let p = format!("{p}.ple");
|
|
let ep = format!("{p}.ple_embedding");
|
|
let eos = c
|
|
.eos_token_id
|
|
.as_i64()
|
|
.or_else(|| {
|
|
c.eos_token_id
|
|
.as_array()
|
|
.and_then(|a| a.first())
|
|
.and_then(serde_json::Value::as_i64)
|
|
})
|
|
.unwrap_or(0);
|
|
Some(ple::Ple {
|
|
embedding: ngram_stage::Embedding::new(
|
|
ngram::GpuHash {
|
|
multipliers: self.array(&format!("{ep}.layer_multipliers"))?,
|
|
sizes: self.array(&format!("{ep}.ngram_heads_vocab_sizes"))?,
|
|
offsets: self.array(&format!("{ep}.ngram_heads_offsets"))?,
|
|
eos,
|
|
},
|
|
ngram::Table {
|
|
resident: None,
|
|
sidecar: None,
|
|
weight: self
|
|
.optional(&format!("{ep}.ngram_embedding.weight"))?
|
|
.cloned(),
|
|
prefer_lazy: false,
|
|
sidecar_mode: c.ngram_sidecar,
|
|
dim: 160,
|
|
},
|
|
),
|
|
projection: ple::Projection {
|
|
key: self.linear(&format!("{p}.key_proj"))?,
|
|
value: self.linear(&format!("{p}.value_proj"))?,
|
|
norm_key: self.array(&format!("{p}.norm_key.weight"))?,
|
|
norm_query: self.array(&format!("{p}.norm_query.weight"))?,
|
|
norm_conv: self.array(&format!("{p}.norm_conv.weight"))?,
|
|
conv_weight: self.array(&format!("{p}.conv_weight"))?,
|
|
},
|
|
})
|
|
} else {
|
|
None
|
|
};
|
|
Ok(decoder::DecoderLayer {
|
|
attention,
|
|
mlp,
|
|
attn_hyper,
|
|
mlp_hyper,
|
|
ple,
|
|
})
|
|
}
|
|
}
|