convert : update to support dflash

This commit is contained in:
Georgi Gerganov
2026-09-28 22:24:36 +03:00
parent 57b557cb95
commit 119282476d
4 changed files with 80 additions and 2 deletions
+1 -1
View File
@@ -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")
+63
View File
@@ -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,63 @@ 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")
with open(self.target_model_dir / "config.json", encoding="utf-8") as f:
target_config = json.load(f)
target_config = target_config.get("text_config", target_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
+6
View File
@@ -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);
+10 -1
View File
@@ -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);