mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-09-27 21:46:57 +02:00
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:
@@ -94,6 +94,7 @@ TEXT_MODEL_MAP: dict[str, str] = {
|
||||
"Gemma3nForCausalLM": "gemma",
|
||||
"Gemma3nForConditionalGeneration": "gemma",
|
||||
"Gemma4AssistantForCausalLM": "gemma",
|
||||
"Gemma4DSparkModel": "gemma",
|
||||
"Gemma4ForConditionalGeneration": "gemma",
|
||||
"Gemma4ForCausalLM": "gemma",
|
||||
"Gemma4UnifiedForConditionalGeneration": "gemma",
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
|
||||
@@ -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
@@ -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;
|
||||
|
||||
Reference in New Issue
Block a user