Integrate DS4 execution parity in Rust

This commit is contained in:
Georg Bauer
2026-07-26 17:58:05 +02:00
parent c9f0c3661c
commit 4420b81117
20 changed files with 11643 additions and 358 deletions

View File

@@ -7,7 +7,7 @@ pub(crate) fn validate_model_artifact(
) -> Result<(), String> {
if support {
let model = Gguf::open(path)?;
validate_dspark(&model, &FLASH)
validate_support(&model, &FLASH).map(|_| ())
} else {
let model = Model::open_main(path, expected)?;
let summary = model.summary();
@@ -29,6 +29,192 @@ pub(crate) fn validate_model_artifact(
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(super) enum SupportKind {
LegacyMtp,
DSpark,
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub(super) struct DsparkConfig {
pub(super) block_size: u32,
pub(super) markov_rank: u32,
pub(super) noise_token: u32,
pub(super) target_layers: Vec<u32>,
pub(super) stages: u32,
}
pub(super) fn dspark_config(model: &Gguf) -> Result<DsparkConfig, String> {
let block_size = first_u32(
model,
&[
"deepseek4.dspark.block_size",
"deepseek4.dspark_block_size",
"dspark.block_size",
],
)?;
let markov_rank = first_u32(
model,
&[
"deepseek4.dspark.markov_rank",
"deepseek4.dspark_markov_rank",
"dspark.markov_rank",
],
)?;
let noise_token = first_u32(
model,
&[
"deepseek4.dspark.noise_token_id",
"deepseek4.dspark_noise_token_id",
"dspark.noise_token_id",
],
)?;
let target_layers = first_u32s(
model,
&[
"deepseek4.dspark.target_layer_ids",
"deepseek4.dspark_target_layer_ids",
"dspark.target_layer_ids",
],
)?;
let stages = model
.tensors
.keys()
.filter_map(|name| {
name.strip_prefix("mtp.")?
.split('.')
.next()?
.parse::<u32>()
.ok()
})
.max()
.map_or(0, |stage| stage + 1);
Ok(DsparkConfig {
block_size,
markov_rank,
noise_token,
target_layers: target_layers.to_vec(),
stages,
})
}
pub(super) fn validate_support(model: &Gguf, shape: &Shape) -> Result<SupportKind, String> {
if model.tensors.contains_key("mtp.0.e_proj.weight")
&& model.tensors.contains_key("mtp.0.h_proj.weight")
&& model.tensors.contains_key("mtp.0.hc_head_base.weight")
{
validate_legacy_mtp(model, shape)?;
Ok(SupportKind::LegacyMtp)
} else if model.metadata.contains_key("deepseek4.dspark.block_size")
|| model.metadata.contains_key("deepseek4.dspark_block_size")
|| model.metadata.contains_key("dspark.block_size")
{
validate_dspark(model, shape)?;
Ok(SupportKind::DSpark)
} else {
Err("support GGUF is neither legacy MTP nor DSpark".into())
}
}
fn validate_legacy_mtp(model: &Gguf, shape: &Shape) -> Result<(), String> {
if shape.model != ModelChoice::DeepSeekV4Flash {
return Err("legacy MTP support is available only for DeepSeek V4 Flash".into());
}
let prefix = "mtp.0";
let hc_dim = shape.embd * shape.hc;
let hc_mix = 2 * shape.hc + shape.hc * shape.hc;
let q_dim = shape.heads * shape.head_dim;
let output_low = shape.out_groups * shape.lora_o;
for (suffix, types, dims) in [
("hc_head_base.weight", &[F32][..], vec![shape.hc]),
("hc_head_fn.weight", PLAIN, vec![hc_dim, shape.hc]),
("hc_head_scale.weight", &[F32][..], vec![1]),
("e_proj.weight", &[Q8_0][..], vec![shape.embd, shape.embd]),
("h_proj.weight", &[Q8_0][..], vec![shape.embd, shape.embd]),
("enorm.weight", &[F32][..], vec![shape.embd]),
("hnorm.weight", &[F32][..], vec![shape.embd]),
("norm.weight", &[F32][..], vec![shape.embd]),
("hc_attn_fn.weight", PLAIN, vec![hc_dim, hc_mix]),
("hc_attn_scale.weight", &[F32][..], vec![3]),
("hc_attn_base.weight", &[F32][..], vec![hc_mix]),
("attn_norm.weight", &[F32][..], vec![shape.embd]),
(
"attn_q_a.weight",
&[Q8_0][..],
vec![shape.embd, shape.lora_q],
),
("attn_q_a_norm.weight", &[F32][..], vec![shape.lora_q]),
("attn_q_b.weight", &[Q8_0][..], vec![shape.lora_q, q_dim]),
(
"attn_kv.weight",
&[Q8_0][..],
vec![shape.embd, shape.head_dim],
),
("attn_kv_a_norm.weight", &[F32][..], vec![shape.head_dim]),
("attn_sinks.weight", &[F32][..], vec![shape.heads]),
(
"attn_output_a.weight",
&[Q8_0][..],
vec![
shape.head_dim * (shape.heads / shape.out_groups),
output_low,
],
),
(
"attn_output_b.weight",
&[Q8_0][..],
vec![output_low, shape.embd],
),
("hc_ffn_fn.weight", PLAIN, vec![hc_dim, hc_mix]),
("hc_ffn_scale.weight", &[F32][..], vec![3]),
("hc_ffn_base.weight", &[F32][..], vec![hc_mix]),
("ffn_norm.weight", &[F32][..], vec![shape.embd]),
(
"ffn_gate_inp.weight",
PLAIN,
vec![shape.embd, shape.experts],
),
("exp_probs_b.bias", &[F32][..], vec![shape.experts]),
(
"ffn_gate_exps.weight",
ROUTED,
vec![shape.embd, shape.ff_expert, shape.experts],
),
(
"ffn_up_exps.weight",
ROUTED,
vec![shape.embd, shape.ff_expert, shape.experts],
),
(
"ffn_down_exps.weight",
ROUTED,
vec![shape.ff_expert, shape.embd, shape.experts],
),
(
"ffn_gate_shexp.weight",
&[Q8_0][..],
vec![shape.embd, shape.ff_expert],
),
(
"ffn_up_shexp.weight",
&[Q8_0][..],
vec![shape.embd, shape.ff_expert],
),
(
"ffn_down_shexp.weight",
&[Q8_0][..],
vec![shape.ff_expert, shape.embd],
),
] {
expect(model, &format!("{prefix}.{suffix}"), types, &dims)?;
}
same_type(
model,
"mtp.0.ffn_gate_exps.weight",
"mtp.0.ffn_up_exps.weight",
)
}
pub(super) fn validate_main(model: &Gguf, expected: ModelChoice) -> Result<Shape, String> {
let family = if model.bytes("general.architecture").ok() == Some(b"glm-dsa") {
ModelFamily::Glm
@@ -575,38 +761,13 @@ pub(super) fn validate_dspark(model: &Gguf, shape: &Shape) -> Result<(), String>
if shape.model != ModelChoice::DeepSeekV4Flash {
return Err("DSpark support is available only for DeepSeek V4 Flash".into());
}
let block_size = first_u32(
model,
&[
"deepseek4.dspark.block_size",
"deepseek4.dspark_block_size",
"dspark.block_size",
],
)?;
let markov_rank = first_u32(
model,
&[
"deepseek4.dspark.markov_rank",
"deepseek4.dspark_markov_rank",
"dspark.markov_rank",
],
)?;
let noise_token = first_u32(
model,
&[
"deepseek4.dspark.noise_token_id",
"deepseek4.dspark_noise_token_id",
"dspark.noise_token_id",
],
)?;
let targets = first_u32s(
model,
&[
"deepseek4.dspark.target_layer_ids",
"deepseek4.dspark_target_layer_ids",
"dspark.target_layer_ids",
],
)?;
let DsparkConfig {
block_size,
markov_rank,
noise_token,
target_layers: targets,
stages,
} = dspark_config(model)?;
if !(1..=16).contains(&block_size) || markov_rank == 0 || noise_token >= shape.vocab as u32 {
return Err("invalid DSpark block, Markov, or noise-token metadata".into());
}
@@ -617,18 +778,6 @@ pub(super) fn validate_dspark(model: &Gguf, shape: &Shape) -> Result<(), String>
{
return Err("invalid DSpark target-layer metadata".into());
}
let stages = model
.tensors
.keys()
.filter_map(|name| {
name.strip_prefix("mtp.")?
.split('.')
.next()?
.parse::<u32>()
.ok()
})
.max()
.map_or(0, |stage| stage + 1);
if !(1..=8).contains(&stages) {
return Err(format!("invalid DSpark stage count: {stages}"));
}
@@ -1079,4 +1228,15 @@ mod tests {
validate_model_artifact(path, ModelChoice::DeepSeekV4Flash, true).unwrap();
}
}
#[test]
fn installed_legacy_mtp_fixture_passes_the_target_layout() {
let path = Path::new("../ds4/gguf/DeepSeek-V4-Flash-MTP-Q4K-Q8_0-F32.gguf");
if path.exists() {
assert_eq!(
validate_support(&Gguf::open(path).unwrap(), &FLASH).unwrap(),
SupportKind::LegacyMtp
);
}
}
}