Commit 22bdcc4cd for llama.cpp

commit 22bdcc4cdd54e590a3ba1da1e5b0d3864bbdda2a
Author: Georgi Gerganov <ggerganov@gmail.com>
Date:   Wed Sep 30 17:04:35 2026 +0300

    mimo : support dflash (convert + feature extraction) (#29650)

    * convert : update to support dflash

    * cont : fix

    Co-authored-by: Sigbjørn Skjæret <sigbjorn.skjaeret@huggingface.co>

    ---------

    Co-authored-by: Sigbjørn Skjæret <sigbjorn.skjaeret@huggingface.co>

diff --git a/conversion/base.py b/conversion/base.py
index 221aa8093..35f564e40 100644
--- a/conversion/base.py
+++ b/conversion/base.py
@@ -234,7 +234,7 @@ class ModelBase:

         prefix = "model" if not self.is_mistral_format else "consolidated"
         part_names: list[str] = ModelBase.get_model_part_names(self.dir_model, prefix, ".safetensors")
-        is_safetensors: bool = len(part_names) > 0
+        is_safetensors: bool = len(part_names) > 0 or (not self.is_mistral_format and (self.dir_model / "model.safetensors.index.json").is_file())
         if not is_safetensors:
             part_names = ModelBase.get_model_part_names(self.dir_model, "pytorch_model", ".bin")

diff --git a/conversion/qwen.py b/conversion/qwen.py
index 64d606176..6b87ff25e 100644
--- a/conversion/qwen.py
+++ b/conversion/qwen.py
@@ -686,6 +686,12 @@ class DFlashModel(Qwen3Model):
         super().set_gguf_parameters()

         dflash_config = self.hparams.get("dflash_config", {})
+        if (partial_rotary_factor := self.rope_parameters.get("partial_rotary_factor")) is not None:
+            head_dim = self.hparams.get("head_dim") or self.hparams["hidden_size"] // self.hparams["num_attention_heads"]
+            self.gguf_writer.add_rope_dimension_count(int(head_dim * partial_rotary_factor))
+        if (value_scale := dflash_config.get("attention_value_scale")) is not None:
+            self.gguf_writer.add_attn_value_scale(float(value_scale))
+
         block_size = dflash_config.get("block_size", self.hparams.get("block_size", 16))
         self.gguf_writer.add_block_size(block_size)

@@ -737,6 +743,62 @@ class DFlashModel(Qwen3Model):
             head_dim = self.hparams.get("head_dim") or self.hparams["hidden_size"] // self.hparams["num_attention_heads"]
             self.gguf_writer.add_rope_dimension_sections([head_dim // 2, 0, 0, 0])

+    def generate_extra_tensors(self) -> Iterable[tuple[str, Tensor]]:
+        yield from super().generate_extra_tensors()
+
+        mask_path = self.dir_model / "mask_embedding.pt"
+        if not mask_path.is_file():
+            return
+
+        mask = torch.load(mask_path, map_location="cpu", weights_only=True)
+        mask_id = self.hparams.get("dflash_config", {}).get("mask_token_id")
+        if mask_id is None or mask["mask_token_id"] != mask_id:
+            raise ValueError("mask_embedding.pt mask_token_id does not match dflash_config")
+        if tuple(mask["embedding"].shape) != (self.hparams["hidden_size"],):
+            raise ValueError("mask_embedding.pt has an unexpected embedding shape")
+        if not 0 <= mask_id < self.hparams["vocab_size"]:
+            raise ValueError("mask_embedding.pt mask_token_id is outside the vocabulary")
+
+        def target_tensor(name: str) -> Tensor:
+            if self.target_model_dir is None:
+                raise ValueError("mask_embedding.pt requires --target-model-dir with the target embeddings and output head")
+            index_path = self.target_model_dir / "model.safetensors.index.json"
+            if index_path.is_file():
+                with open(index_path, encoding="utf-8") as f:
+                    weight_map = json.load(f)["weight_map"]
+                part_names = [weight_map[name]]
+            else:
+                part_names = self.get_model_part_names(self.target_model_dir, "model", ".safetensors")
+
+            for part_name in part_names:
+                with gguf.utility.SafetensorsLocal(self.target_model_dir / part_name) as part:
+                    if name in part:
+                        return LazyTorchTensor.from_local_tensor(part[name])
+            raise ValueError(f"Target tensor {name!r} was not found in safetensors")
+
+        embedding_name = "model.embed_tokens.weight"
+        if embedding_name in self.model_tensors:
+            embeddings = self.model_tensors.pop(embedding_name)()
+        else:
+            embeddings = target_tensor(embedding_name)
+
+        if "model.lm_head.weight" not in self.model_tensors:
+            if self.target_model_dir is None:
+                raise ValueError("mask_embedding.pt requires --target-model-dir to obtain the output head")
+            target_config = ModelBase.load_hparams(self.target_model_dir, False)
+            target_config = {**target_config, **target_config.get("text_config", {})}
+            head_name = embedding_name if target_config.get("tie_word_embeddings", False) else "lm_head.weight"
+            # Keep the output head separate from the patched input embedding table.
+            yield "model.lm_head.weight", target_tensor(head_name)
+
+        embeddings = LazyTorchTensor.to_eager(embeddings).clone()
+        if tuple(embeddings.shape) != (self.hparams["vocab_size"], self.hparams["hidden_size"]):
+            raise ValueError("Target token embedding shape does not match the DFlash draft")
+        # MiMo's target mask row is untrained; the draft provides its own vector.
+        embeddings[mask_id] = mask["embedding"].to(embeddings.dtype)
+        self.hparams["has_embed_tokens"] = True
+        yield embedding_name, embeddings
+
     def _target_uses_mrope(self) -> bool:
         if self.target_model_dir is None:
             return False
diff --git a/src/models/dflash.cpp b/src/models/dflash.cpp
index 1e8881c0c..c448e63f3 100644
--- a/src/models/dflash.cpp
+++ b/src/models/dflash.cpp
@@ -8,6 +8,7 @@ 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);
+    ml.get_key(LLM_KV_ATTENTION_VALUE_SCALE, hparams.f_attn_value_scale, false);

     hparams.llm_ffn_op = LLM_FFN_SILU;
     std::string hidden_act;
@@ -738,6 +739,11 @@ llama_model_dflash::graph<false>::graph(const llama_model & model, const llm_gra
             ? build_attn(inp_attn_iswa, layer.wo, NULL, layer.wo_s, Qcur, Kcur, Vcur, nullptr, layer.attn_sinks, nullptr, kq_scale, il)
             : build_attn(inp_attn,      layer.wo, NULL, layer.wo_s, Qcur, Kcur, Vcur, nullptr, layer.attn_sinks, nullptr, kq_scale, il);

+        if (hparams.f_attn_value_scale != 0.0f) {
+            cur = ggml_scale(ctx0, cur, hparams.f_attn_value_scale);
+            cb(cur, "attn_out_scaled", il);
+        }
+
         if (attn_dynamic) {
             cur = build_dflash2_conv(*this, cur, attn_dynamic, layer.dflash_attn_conv_base, 1);
             cb(cur, "attn_conv_out", il);
diff --git a/src/models/mimo2.cpp b/src/models/mimo2.cpp
index b6d7aceda..ce315f956 100644
--- a/src/models/mimo2.cpp
+++ b/src/models/mimo2.cpp
@@ -102,9 +102,12 @@ llama_model_mimo2::graph::graph(const llama_model & model, const llm_graph_param

     const float v_scale = hparams.f_attn_value_scale;
     const bool emit_h_nextn = cparams.embeddings_nextn;
-    const bool crop_last_layer = inp_out_ids && (!emit_h_nextn || cparams.embeddings_nextn_masked);
+    const bool extract_final_inp = (size_t) n_layer < cparams.embeddings_layer_inp.size() && cparams.embeddings_layer_inp[n_layer];
+    const bool crop_last_layer = inp_out_ids && (!emit_h_nextn || cparams.embeddings_nextn_masked) && !extract_final_inp;

     for (int il = 0; il < n_layer; ++il) {
+        res->t_layer_inp[il] = inpL;
+
         ggml_tensor * inpSA = inpL;

         uint32_t n_head_l    = hparams.n_head(il);
@@ -231,6 +234,12 @@ llama_model_mimo2::graph::graph(const llama_model & model, const llm_graph_param
     }

     cur = inpL;
+    if (extract_final_inp) {
+        res->t_layer_inp[n_layer] = cur;
+        if (inp_out_ids && (!emit_h_nextn || cparams.embeddings_nextn_masked)) {
+            cur = ggml_get_rows(ctx0, cur, inp_out_ids);
+        }
+    }

     if (emit_h_nextn) {
         cb(cur, "h_nextn", -1);