model : support Gemma4 DSpark draft backbone (#29226)

* dspark: add Gemma 4 draft support

Add GGUF conversion and runtime support for full-attention and SWA Gemma 4
DSpark drafts, including tied output weights and boolean backbone metadata.

Assisted-by: Codex

* dflash: infer Gemma draft features from metadata
This commit is contained in:
Hrishith Thadicherla
2026-09-23 13:34:09 +03:00
committed by GitHub
parent 86b2daa730
commit 633733d0ae
5 changed files with 160 additions and 12 deletions
+1
View File
@@ -94,6 +94,7 @@ TEXT_MODEL_MAP: dict[str, str] = {
"Gemma3nForCausalLM": "gemma",
"Gemma3nForConditionalGeneration": "gemma",
"Gemma4AssistantForCausalLM": "gemma",
"Gemma4DSparkModel": "gemma",
"Gemma4ForConditionalGeneration": "gemma",
"Gemma4ForCausalLM": "gemma",
"Gemma4UnifiedForConditionalGeneration": "gemma",
+100
View File
@@ -11,6 +11,7 @@ if TYPE_CHECKING:
from torch import Tensor
from .base import MmprojModel, ModelBase, TextModel, gguf, logger
from .qwen import DFlashModel
@ModelBase.register("GemmaForCausalLM")
@@ -809,6 +810,105 @@ class Gemma4Model(Gemma3Model):
yield from super().modify_tensors(data_torch, name, bid)
@ModelBase.register("Gemma4DSparkModel")
class Gemma4DSparkModel(DFlashModel):
model_arch = gguf.MODEL_ARCH.DFLASH
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
if not self.hparams.get("attention_k_eq_v", False):
raise ValueError("Gemma4 DSpark currently requires attention_k_eq_v")
if self.hparams.get("layer_types") != ["full_attention"] * self.block_count:
raise ValueError("Gemma4 DSpark currently requires uniform full_attention layer types")
if self.hparams.get("hidden_activation", "gelu_pytorch_tanh") != "gelu_pytorch_tanh":
raise ValueError("Gemma4 DSpark currently requires hidden_activation=gelu_pytorch_tanh")
if self.hparams.get("attention_bias", False) or self.hparams.get("enable_moe_block", False):
raise ValueError("Gemma4 DSpark attention bias and MoE are not supported")
if (self.hparams.get("draft_vocab_size") or self.hparams["vocab_size"]) != self.hparams["vocab_size"]:
raise ValueError("Gemma4 DSpark currently requires a full draft vocabulary")
if "model.lm_head.weight" not in self.model_tensors and self.hparams.get("tie_word_embeddings") is not True:
raise ValueError("Gemma4 DSpark requires lm_head.weight unless tie_word_embeddings is true")
self.dflash_config = self.hparams.get("dflash_config", {})
markov_type = self.dflash_config.get("markov_head_type", self.hparams.get("markov_head_type", "vanilla"))
if markov_type != "vanilla":
raise ValueError("Gemma4 DSpark currently requires a vanilla Markov head")
# Gemma4TextConfig supplies these defaults when rope_parameters is absent.
rope = self.hparams.get("rope_parameters") or {
"full_attention": {"rope_type": "proportional", "partial_rotary_factor": 0.25, "rope_theta": 1000000.0},
}
self.rope_parameters = rope.get("full_attention", rope)
if self.rope_parameters.get("rope_type") not in ("default", "proportional"):
raise ValueError("Gemma4 DSpark requires default or proportional RoPE")
def set_vocab(self):
super().set_vocab()
mask_id = self.dflash_config.get("mask_token_id", self.hparams.get("mask_token_id"))
if mask_id is None:
raise ValueError("Gemma4 DSpark requires mask_token_id")
if "mask_token_id" not in self.dflash_config:
self.gguf_writer.add_mask_token_id(mask_id)
def set_gguf_parameters(self):
super().set_gguf_parameters()
head_dim = int(self.hparams["global_head_dim"])
self.gguf_writer.add_head_count_kv(self.hparams["num_global_key_value_heads"])
self.gguf_writer.add_key_length(head_dim)
self.gguf_writer.add_value_length(head_dim)
self.gguf_writer.add_rope_dimension_count(head_dim)
self.gguf_writer.add_embedding_scale(self.hparams["hidden_size"] ** 0.5)
self.gguf_writer.add_attention_scale(1.0)
self.gguf_writer.add_hidden_act("gelu_pytorch_tanh")
self.gguf_writer.add_sample_from_anchor(self.hparams.get("sample_from_anchor", True))
target_layers = self.dflash_config.get("target_layer_ids", self.hparams.get("target_layer_ids"))
if not target_layers:
raise ValueError("Gemma4 DSpark requires target_layer_ids")
self.gguf_writer.add_has_confidence_head(any("confidence_head.proj" in name for name in self.model_tensors))
if self.hparams.get("final_logit_softcapping"):
raise ValueError("Gemma4 DSpark logit softcapping is not supported")
# The top-level sliding_window is inert unless the draft enables SWA.
if self.dflash_config.get("use_swa", False):
window = self.dflash_config["swa_window_size"]
if window <= 0:
raise ValueError("Gemma4 DSpark swa_window_size must be positive")
@classmethod
def filter_tensors(cls, item: tuple[str, Callable[[], Tensor]]) -> tuple[str, Callable[[], Tensor]] | None:
name, gen = item
if not name.startswith("model."):
name = "model." + name
if name.endswith(".layer_scalar"):
name += ".weight"
name = name.replace("model.confidence_proj.", "model.confidence_head.proj.")
return super().filter_tensors((name, gen))
def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iterable[tuple[str, Tensor]]:
# The shared DFlash map assigns this name to Qwen's pre-FFN norm.
if name.endswith(".post_attention_layernorm.weight"):
name = self.format_tensor_name(gguf.MODEL_TENSOR.ATTN_POST_NORM, bid)
elif name.endswith(".pre_feedforward_layernorm.weight"):
name = self.format_tensor_name(gguf.MODEL_TENSOR.FFN_NORM, bid)
yield from super().modify_tensors(data_torch, name, bid)
def generate_extra_tensors(self) -> Iterable[tuple[str, Tensor]]:
if self.rope_parameters["rope_type"] == "proportional":
# Keep the unrotated dimensions in place, as in the Gemma4 converter.
head_dim = int(self.hparams["global_head_dim"])
fraction_value = self.rope_parameters.get("partial_rotary_factor", 0.25)
if not isinstance(fraction_value, (int, float)):
raise ValueError("Gemma4 DSpark partial_rotary_factor must be numeric")
fraction = float(fraction_value)
n_rot = int(head_dim * fraction / 2)
if not 0 < fraction <= 1 or head_dim * fraction != 2 * n_rot:
raise ValueError("Gemma4 DSpark rotary dimension count must be positive and even")
factors = torch.tensor([1.0] * n_rot + [1e30] * (head_dim // 2 - n_rot), dtype=torch.float32)
yield self.format_tensor_name(gguf.MODEL_TENSOR.ROPE_FREQS), factors
@ModelBase.register("Gemma4UnifiedForConditionalGeneration")
@ModelBase.example("hf-tiny-v2/tiny-random-Gemma4UnifiedForConditionalGeneration")
class Gemma4UnifiedModel(Gemma4Model):
+4 -3
View File
@@ -711,7 +711,7 @@ class DFlashModel(Qwen3Model):
if embedding_scale is not None:
self.gguf_writer.add_embedding_scale(float(embedding_scale))
target_layer_ids = dflash_config.get("target_layer_ids", [])
target_layer_ids = dflash_config.get("target_layer_ids", self.hparams.get("target_layer_ids", []))
if target_layer_ids:
extract_layer_ids = [i + 1 for i in target_layer_ids]
self.gguf_writer.add_target_layers(extract_layer_ids)
@@ -719,8 +719,9 @@ class DFlashModel(Qwen3Model):
use_sliding_window = self.hparams.get("use_sliding_window", False) or dflash_config.get("use_swa", False)
sliding_window = dflash_config.get("swa_window_size") or self.hparams.get("sliding_window")
layer_types = self.hparams.get("layer_types")
if use_sliding_window and sliding_window and layer_types:
is_swa = [lt == "sliding_attention" for lt in layer_types]
if use_sliding_window and sliding_window:
is_swa = ([True] * self.block_count if dflash_config.get("use_swa", False)
else [lt == "sliding_attention" for lt in layer_types or []])
self.gguf_writer.add_sliding_window(sliding_window)
self.gguf_writer.add_sliding_window_pattern(is_swa)
+4
View File
@@ -5207,6 +5207,10 @@ MODEL_TENSORS: dict[MODEL_ARCH, list[MODEL_TENSOR]] = {
MODEL_TENSOR.D2T,
],
MODEL_ARCH.DFLASH: [
MODEL_TENSOR.ATTN_POST_NORM,
MODEL_TENSOR.FFN_POST_NORM,
MODEL_TENSOR.LAYER_OUT_SCALE,
MODEL_TENSOR.ROPE_FREQS,
MODEL_TENSOR.TOKEN_EMBD,
MODEL_TENSOR.OUTPUT,
MODEL_TENSOR.OUTPUT_NORM,
+51 -9
View File
@@ -6,6 +6,19 @@
void llama_model_dflash::load_arch_hparams(llama_model_loader & ml) {
ml.get_key(LLM_KV_EMBEDDING_SCALE, hparams.f_embedding_scale, false);
ml.get_key(LLM_KV_ATTENTION_SCALE, hparams.f_attention_scale, false);
hparams.llm_ffn_op = LLM_FFN_SILU;
std::string hidden_act;
if (ml.get_key(LLM_KV_HIDDEN_ACT, hidden_act, false)) {
if (hidden_act == "gelu" || hidden_act == "gelu_pytorch_tanh") {
hparams.llm_ffn_op = LLM_FFN_GELU;
} else if (hidden_act != "silu") {
throw std::runtime_error("unsupported DFlash hidden activation: " + hidden_act);
}
}
ml.get_key(LLM_KV_ATTENTION_LAYERNORM_RMS_EPS, hparams.f_norm_rms_eps);
ml.get_key(LLM_KV_LOGIT_SCALE, hparams.f_logit_scale, false);
hparams.f_final_logit_softcapping = 0.0f;
@@ -108,9 +121,6 @@ void llama_model_dflash::load_arch_tensors(llama_model_loader &) {
}
// DSpark = DFlash + a semi-autoregressive Markov head and Confidence head
//
// TODO: only Qwen3-style backbones are supported for now; other backbones (e.g. Gemma4)
// need their own conversion path and graph tweaks
const struct ggml_tensor * markov_meta = ml->get_tensor_meta("markov_w1.weight");
if (markov_meta) {
const int64_t dspark_markov_rank = markov_meta->ne[0];
@@ -156,6 +166,9 @@ void llama_model_dflash::load_arch_tensors(llama_model_loader &) {
// optional: reduced-vocab drafts ship their own lm head, full-vocab drafts can share the target's via ctx_other
// a draft with its own embeddings + head references no target tensors and can run on devices the target does not use (e.g. -devd with a tensor-split target)
output = create_tensor(tn(LLM_TENSOR_OUTPUT, "weight"), { n_embd, n_vocab_draft }, TENSOR_NOT_REQUIRED);
if (output == nullptr && tok_embd != nullptr) {
output = create_tensor(tn(LLM_TENSOR_TOKEN_EMBD, "weight"), { n_embd, n_vocab_draft }, TENSOR_DUPLICATED);
}
if (hparams.dsv4_hc_mult > 0) {
const int64_t q_lora_rank = hparams.n_lora_q;
@@ -214,12 +227,17 @@ void llama_model_dflash::load_arch_tensors(llama_model_loader &) {
layer.wq = create_tensor(tn(LLM_TENSOR_ATTN_Q, "weight", i), { n_embd, n_embd_head_k * n_head }, 0);
layer.wk = create_tensor(tn(LLM_TENSOR_ATTN_K, "weight", i), { n_embd, n_embd_k_gqa }, 0);
layer.wv = create_tensor(tn(LLM_TENSOR_ATTN_V, "weight", i), { n_embd, n_embd_v_gqa }, 0);
layer.wv = create_tensor(tn(LLM_TENSOR_ATTN_V, "weight", i), { n_embd, n_embd_v_gqa }, TENSOR_NOT_REQUIRED);
layer.wo = create_tensor(tn(LLM_TENSOR_ATTN_OUT, "weight", i), { n_embd_head_k * n_head, n_embd }, 0);
layer.attn_q_norm = create_tensor(tn(LLM_TENSOR_ATTN_Q_NORM, "weight", i), { n_embd_head_k }, 0);
layer.attn_k_norm = create_tensor(tn(LLM_TENSOR_ATTN_K_NORM, "weight", i), { n_embd_head_k }, 0);
layer.attn_post_norm = create_tensor(tn(LLM_TENSOR_ATTN_POST_NORM, "weight", i), { n_embd }, TENSOR_NOT_REQUIRED);
layer.ffn_post_norm = create_tensor(tn(LLM_TENSOR_FFN_POST_NORM, "weight", i), { n_embd }, TENSOR_NOT_REQUIRED);
layer.out_scale = create_tensor(tn(LLM_TENSOR_LAYER_OUT_SCALE, "weight", i), { 1 }, TENSOR_NOT_REQUIRED);
layer.rope_freqs = create_tensor(tn(LLM_TENSOR_ROPE_FREQS, "weight", i), { n_embd_head_k/2 }, TENSOR_NOT_REQUIRED | (i > 0 ? TENSOR_DUPLICATED : 0));
// optional per-head attention sinks (e.g. Nemotron DSpark)
layer.attn_sinks = create_tensor(tn(LLM_TENSOR_ATTN_SINKS, "weight", i), { n_head }, TENSOR_NOT_REQUIRED);
@@ -571,7 +589,7 @@ llama_model_dflash::graph<false>::graph(const llama_model & model, const llm_gra
inp_attn = build_attn_inp_kv();
}
const float kq_scale = 1.0f/sqrtf(float(n_embd_head));
const float kq_scale = hparams.f_attention_scale != 0.0f ? hparams.f_attention_scale : 1.0f/sqrtf(float(n_embd_head));
// drafts for M-RoPE targets use degenerate sections (temporal dim only)
int sections[4];
@@ -582,7 +600,7 @@ llama_model_dflash::graph<false>::graph(const llama_model & model, const llm_gra
? ggml_rope_multi(ctx0, cur, pos, nullptr,
n_rot, sections, rope_type, n_ctx_orig, freq_base, freq_scale,
ext_factor, attn_factor, beta_fast, beta_slow)
: ggml_rope_ext(ctx0, cur, pos, nullptr,
: ggml_rope_ext(ctx0, cur, pos, model.layers[0].rope_freqs,
n_rot, rope_type, n_ctx_orig, freq_base, freq_scale,
ext_factor, attn_factor, beta_fast, beta_slow);
};
@@ -608,12 +626,16 @@ llama_model_dflash::graph<false>::graph(const llama_model & model, const llm_gra
const auto & layer = model.layers[il];
ggml_tensor * Kcur = build_lora_mm(layer.wk, inp_g, layer.wk_s);
ggml_tensor * Vcur = build_lora_mm(layer.wv, inp_g, layer.wv_s);
const bool shared_kv = layer.wv == nullptr;
ggml_tensor * Vcur = shared_kv ? Kcur : build_lora_mm(layer.wv, inp_g, layer.wv_s);
Kcur = ggml_reshape_3d(ctx0, Kcur, n_embd_head, n_head_kv, n_tokens);
Vcur = ggml_reshape_3d(ctx0, Vcur, n_embd_head, n_head_kv, n_tokens);
Kcur = build_norm(Kcur, layer.attn_k_norm, NULL, LLM_NORM_RMS, il);
if (shared_kv) {
Vcur = ggml_rms_norm(ctx0, Vcur, hparams.f_norm_rms_eps);
}
Kcur = build_rope(Kcur, inp_pos);
cb(Kcur, "Kcur_injected", il);
cb(Vcur, "Vcur_injected", il);
@@ -673,6 +695,9 @@ llama_model_dflash::graph<false>::graph(const llama_model & model, const llm_gra
ggml_tensor * inp_tokens = inp->tokens;
ggml_tensor * inpL = ggml_get_rows(ctx0, tok_embd, inp->tokens);
if (hparams.f_embedding_scale != 0.0f) {
inpL = ggml_scale(ctx0, inpL, hparams.f_embedding_scale);
}
cb(inpL, "inp_noise_embd", -1);
res->add_input(std::move(inp));
@@ -692,7 +717,8 @@ llama_model_dflash::graph<false>::graph(const llama_model & model, const llm_gra
ggml_tensor * Qcur = build_lora_mm(layer.wq, noise_norm, layer.wq_s);
ggml_tensor * Kcur = build_lora_mm(layer.wk, noise_norm, layer.wk_s);
ggml_tensor * Vcur = build_lora_mm(layer.wv, noise_norm, layer.wv_s);
const bool shared_kv = layer.wv == nullptr;
ggml_tensor * Vcur = shared_kv ? Kcur : build_lora_mm(layer.wv, noise_norm, layer.wv_s);
Qcur = ggml_reshape_3d(ctx0, Qcur, n_embd_head, n_head, n_tokens);
Kcur = ggml_reshape_3d(ctx0, Kcur, n_embd_head, n_head_kv, n_tokens);
@@ -700,6 +726,9 @@ llama_model_dflash::graph<false>::graph(const llama_model & model, const llm_gra
Qcur = build_norm(Qcur, layer.attn_q_norm, NULL, LLM_NORM_RMS, il);
Kcur = build_norm(Kcur, layer.attn_k_norm, NULL, LLM_NORM_RMS, il);
if (shared_kv) {
Vcur = ggml_rms_norm(ctx0, Vcur, hparams.f_norm_rms_eps);
}
Qcur = build_rope(Qcur, inp_pos);
Kcur = build_rope(Kcur, inp_pos);
@@ -717,6 +746,11 @@ llama_model_dflash::graph<false>::graph(const llama_model & model, const llm_gra
cb(cur, "attn_conv_out", il);
}
if (layer.attn_post_norm) {
cur = build_norm(cur, layer.attn_post_norm, NULL, LLM_NORM_RMS, il);
cb(cur, "attn_post_norm", il);
}
ggml_tensor * ffn_inp = ggml_add(ctx0, cur, inpL);
cb(ffn_inp, "ffn_inp", il);
@@ -735,7 +769,7 @@ llama_model_dflash::graph<false>::graph(const llama_model & model, const llm_gra
layer.ffn_gate, NULL, layer.ffn_gate_s,
layer.ffn_down, NULL, layer.ffn_down_s,
NULL,
LLM_FFN_SILU, LLM_FFN_PAR, il);
hparams.llm_ffn_op, LLM_FFN_PAR, il);
cb(cur, "ffn_out", il);
if (ffn_dynamic) {
@@ -743,7 +777,15 @@ llama_model_dflash::graph<false>::graph(const llama_model & model, const llm_gra
cb(cur, "ffn_conv_out", il);
}
if (layer.ffn_post_norm) {
cur = build_norm(cur, layer.ffn_post_norm, NULL, LLM_NORM_RMS, il);
cb(cur, "ffn_post_norm", il);
}
cur = ggml_add(ctx0, cur, ffn_inp);
if (layer.out_scale) {
cur = ggml_mul(ctx0, cur, layer.out_scale);
}
cb(cur, "l_out", il);
inpL = cur;