Dflash 2 speculative decoding (#2345)

* Add DFlash2 speculative decoding support

* Fix DFlash2 requantized output selection

* Fix legacy DFlash output row contract

* return llm_build_norm to apply also for dflash 2

* Fix DFlash2 GGUF conversion

* capture dflash states in any type

* remove duplicated layer_rows_raw
This commit is contained in:
Samuel Oliveira Alves 2026-08-26 12:09:34 -03:00 committed by GitHub
parent 73ad16269b
commit 28fbe34ce9
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
28 changed files with 1100 additions and 46 deletions

View File

@ -226,6 +226,10 @@ struct common_params_speculative {
int32_t mtp_heads = 1; // MTP heads to use; 1 is the default, while >1 and 0 (all model heads) are experimental
int32_t dflash_cross_ctx = 512; // target-feature context window for DFlash
// Samplers for DFlash2
float draft_temperature = 0.0f;
uint32_t draft_seed = LLAMA_DEFAULT_SEED;
float p_split = 0.1f; // speculative decoding split probability
float p_min = 0.75f; // minimum speculative decoding probability (greedy)

View File

@ -2,10 +2,15 @@
#include "sampling.h"
#include "llama-vocab.h"
#include "common.h"
#include "speculative.h"
#include "reasoning-budget.cpp"
#include <limits>
#include <random>
#include <unordered_map>
// Keep dflash randomness independent from target
static constexpr uint32_t COMMON_SPECULATIVE_VERIFIER_SEED_XOR = 0x9e3779b9U;
#if defined(__GNUC__) && (defined(__x86_64__) || defined(__i386__))
#include <immintrin.h>
#endif
@ -220,6 +225,7 @@ void common_sampler_reset(common_sampler * ctx) {
llama_free_adaptive_p(ctx->adapt_p_ctx);
ctx->adapt_p_ctx = llama_init_adaptive_p(ctx->params.adaptive_target, ctx->params.adaptive_decay, ctx->params.adaptive_updt_w_cur, ctx->rng());
ctx->speculative_rng.seed(ctx->speculative_seed);
}
void common_sampler_review(common_sampler * ctx, const size_t n_unsent, const bool rewind_status) {
@ -234,6 +240,8 @@ void llama_sampling_set_rng_seed(struct common_sampler * ctx, uint32_t seed) {
seed = std::random_device{}();
}
ctx->rng.seed(seed);
ctx->speculative_seed = seed ^ COMMON_SPECULATIVE_VERIFIER_SEED_XOR;
ctx->speculative_rng.seed(ctx->speculative_seed);
}
void common_sampler_clone(common_sampler * src, common_sampler * dst) {
@ -241,6 +249,8 @@ void common_sampler_clone(common_sampler * src, common_sampler * dst) {
dst->mirostat_mu = src->mirostat_mu;
dst->n_valid = src->n_valid;
dst->rng = src->rng;
dst->speculative_seed = src->speculative_seed;
dst->speculative_rng = src->speculative_rng;
dst->server_biases = src->server_biases;
if (dst->grammar) {
@ -811,6 +821,80 @@ std::vector<llama_token> common_sampler_sample_and_accept_n(struct common_sample
return result;
}
std::vector<llama_token> common_sampler_sample_and_accept_n(
struct common_sampler * gsmpl,
struct llama_context * ctx,
const std::vector<int> & idxs,
const std::vector<llama_token> & draft,
const std::vector<common_speculative_token_dist> & dists,
bool grammar_first) {
GGML_ASSERT(idxs.size() == draft.size() + 1);
GGML_ASSERT(dists.size() == draft.size());
std::vector<llama_token> result;
result.reserve(idxs.size());
std::uniform_real_distribution<float> uniform(0.0f, 1.0f);
const auto emit = [&](llama_token id) {
gsmpl->drafted_text += common_token_to_piece(ctx, id, true);
common_sampler_accept(gsmpl, ctx, id, true);
result.push_back(id);
};
size_t i = 0;
for (; i < draft.size(); ++i) {
const llama_token fallback = common_sampler_sample(gsmpl, ctx, idxs[i], true);
const auto & q = dists[i];
GGML_ASSERT(q.ids.size() == q.probs.size());
std::unordered_map<llama_token, float> q_probs;
q_probs.reserve(q.ids.size());
for (size_t j = 0; j < q.ids.size(); ++j) {
q_probs[q.ids[j]] += q.probs[j];
}
const auto q_prob = [&q_probs](llama_token id) {
const auto it = q_probs.find(id);
return it == q_probs.end() ? 0.0f : it->second;
};
const auto * p = common_sampler_get_candidates(gsmpl, false);
float p_draft = 0.0f;
for (size_t j = 0; j < p->size; ++j) {
if (p->data[j].id == draft[i]) {
p_draft = p->data[j].p;
break;
}
}
const float q_draft = q_prob(draft[i]);
if (q_draft > 0.0f && uniform(gsmpl->speculative_rng) * q_draft <= p_draft) {
emit(draft[i]);
continue;
}
std::vector<float> residual(p->size);
float residual_sum = 0.0f;
for (size_t j = 0; j < p->size; ++j) {
residual[j] = std::max(0.0f, p->data[j].p - q_prob(p->data[j].id));
residual_sum += residual[j];
}
llama_token id = fallback;
if (residual_sum > 0.0f) {
std::discrete_distribution<size_t> sample(residual.begin(), residual.end());
id = p->data[sample(gsmpl->speculative_rng)].id;
}
emit(id);
break;
}
if (i == draft.size()) {
const llama_token id = common_sampler_sample(gsmpl, ctx, idxs[i], grammar_first);
emit(id);
}
return result;
}
static void elb_print(common_params_sampling& sparams, const common_params_sampling::elb_param::elb_entry& entry) {
#undef X

View File

@ -11,6 +11,8 @@
#define A_DOT_B(a, b) a.b
struct common_speculative_token_dist;
// sampler types
enum class llama_sampler_type : char {
DRY = 'd',
@ -236,6 +238,8 @@ struct common_sampler {
llama_token_data_array cur_p; // current candidates
std::mt19937 rng;
uint32_t speculative_seed = LLAMA_DEFAULT_SEED;
std::mt19937 speculative_rng;
std::vector<float>* server_biases;
@ -357,6 +361,14 @@ std::vector<llama_token> llama_sampling_sample_and_accept_n(struct common_sample
std::vector<llama_token> common_sampler_sample_and_accept_n(struct common_sampler * gsmpl, struct llama_context * ctx, const std::vector<int> & idxs, const std::vector<llama_token> & draft, bool grammar_first = false);
std::vector<llama_token> common_sampler_sample_and_accept_n(
struct common_sampler * gsmpl,
struct llama_context * ctx,
const std::vector<int> & idxs,
const std::vector<llama_token> & draft,
const std::vector<common_speculative_token_dist> & dists,
bool grammar_first = false);
// Greedy argmax sampling for speculative drafting
llama_token common_sampler_sample_speculative(struct common_sampler * gsmpl, struct llama_context * ctx, int idx, float * out_prob = nullptr);

View File

@ -1,8 +1,12 @@
#pragma once
#include <algorithm>
#include <cmath>
#include <cstddef>
#include <cstring>
#include <cstdlib>
#include <limits>
#include <random>
#include <vector>
static bool common_speculative_are_dflash_compatible(
@ -79,6 +83,8 @@ static void dflash_materialize_target_window_features(common_speculative_state_d
// DFlash runtime state and draft path.
struct common_speculative_state_dflash : public common_speculative_state {
// Separated seed for dflash 2 againts target samplers
static constexpr uint32_t SELECTOR_SEED_XOR = 0x85ebca6bU;
llama_context * ctx_tgt;
llama_context * ctx_dft;
@ -91,8 +97,14 @@ struct common_speculative_state_dflash : public common_speculative_state {
int32_t cross_ctx = 0;
bool is_dspark = false;
bool is_dsv4_dspark = false;
bool is_dflash2 = false;
bool ready = false;
std::vector<common_speculative_token_dist> proposal_dists;
std::mt19937 selector_rng;
uint32_t selector_seed = LLAMA_DEFAULT_SEED;
bool selector_rng_initialized = false;
std::vector<int32_t> target_layer_ids;
std::vector<float> target_window;
std::vector<llama_pos> target_window_pos;
@ -127,6 +139,10 @@ struct common_speculative_state_dflash : public common_speculative_state {
is_dspark = type == COMMON_SPECULATIVE_TYPE_DSPARK;
is_dsv4_dspark = is_dspark && llama_model_is_deepseek4(model_tgt);
char selector_top_k[32] = {};
if (llama_model_meta_val_str(model_dft, "dflash.selector_top_k", selector_top_k, sizeof(selector_top_k)) >= 0) {
is_dflash2 = std::atoi(selector_top_k) > 0;
}
const bool has_dspark_head = llama_model_dflash_has_dspark_head(model_dft);
if (is_dspark != has_dspark_head) {
LOG_ERR("%s: %s stage requires %s DSpark Markov tensors\n", __func__,
@ -264,6 +280,8 @@ struct common_speculative_state_dflash : public common_speculative_state {
GGML_UNUSED(prompt);
llama_kv_cache_clear(ctx_dft);
llama_reset_dflash_kv_cache_state(ctx_dft);
proposal_dists.clear();
selector_rng_initialized = false;
}
void draft(
@ -274,6 +292,7 @@ struct common_speculative_state_dflash : public common_speculative_state {
GGML_UNUSED(prompt_tgt);
result.clear();
proposal_dists.clear();
if (!ready || target_window_rows <= 0) {
return;
}
@ -321,13 +340,15 @@ struct common_speculative_state_dflash : public common_speculative_state {
llama_kv_cache_clear(ctx_dft);
batch.n_tokens = 0;
const int32_t batch_len = is_dspark ? n_keep : n_keep + 1;
const bool output_seed_row = is_dspark;
const bool output_mask_rows = !is_dflash2;
// id_last's true position is one past the newest committed feature row
// (last_target_pos): seed there, masks follow. Mirrors mainline's
// [id_last @ n_past, mask @ n_past+1, ...] block geometry.
const llama_pos draft_pos_base = last_target_pos >= 0 ? last_target_pos + 1 : (llama_pos) target_window_rows;
common_batch_add(batch, id_last, draft_pos_base, { 0 }, is_dspark);
common_batch_add(batch, id_last, draft_pos_base, { 0 }, output_seed_row);
for (int32_t i = 1; i < batch_len; ++i) {
common_batch_add(batch, mask_token_id, draft_pos_base + i, { 0 }, true);
common_batch_add(batch, mask_token_id, draft_pos_base + i, { 0 }, output_mask_rows);
}
if (llama_decode(ctx_dft, batch) != 0) {
@ -335,6 +356,70 @@ struct common_speculative_state_dflash : public common_speculative_state {
batch.n_tokens = 0;
return;
}
const int32_t selector_top_k = llama_get_dflash_draft_lattice_top_k(ctx_dft);
if (selector_top_k > 0) {
const int32_t n_positions = llama_get_dflash_draft_lattice_n_positions(ctx_dft);
const int32_t n_positions_used = std::min(n_positions, n_keep + 1);
if (n_positions_used <= 1) {
batch.n_tokens = 0;
return;
}
std::vector<float> scores((size_t) selector_top_k * selector_top_k * n_positions_used);
std::vector<int32_t> ids((size_t) selector_top_k * n_positions_used);
if (!llama_copy_dflash_draft_lattice(ctx_dft, scores.data(), scores.size(), ids.data(), ids.size())) {
LOG_ERR("%s: failed to copy DFlash2 selector lattice\n", __func__);
batch.n_tokens = 0;
return;
}
const float temperature = params.draft_temperature;
if (!selector_rng_initialized || selector_seed != params.draft_seed) {
selector_seed = params.draft_seed == LLAMA_DEFAULT_SEED
? std::random_device{}()
: params.draft_seed;
selector_rng.seed(selector_seed ^ SELECTOR_SEED_XOR);
selector_rng_initialized = true;
}
int32_t predecessor = 0;
for (int32_t pos = 1; pos < n_positions_used; ++pos) {
const float * row = scores.data() + (size_t) pos * selector_top_k * selector_top_k;
const float * path = row + (size_t) predecessor * selector_top_k;
common_speculative_token_dist dist;
if (temperature > 0.0f) {
dist.ids.resize(selector_top_k);
dist.probs.resize(selector_top_k);
const float max_score = *std::max_element(path, path + selector_top_k);
float sum = 0.0f;
for (int32_t k = 0; k < selector_top_k; ++k) {
dist.ids[(size_t) k] = (llama_token) ids[(size_t) pos * selector_top_k + k];
dist.probs[(size_t) k] = std::exp((path[k] - max_score) / temperature);
sum += dist.probs[(size_t) k];
}
if (!(sum > 0.0f) || !std::isfinite(sum)) {
result.clear();
proposal_dists.clear();
batch.n_tokens = 0;
return;
}
for (float & probability : dist.probs) {
probability /= sum;
}
std::discrete_distribution<int32_t> sample(dist.probs.begin(), dist.probs.end());
predecessor = sample(selector_rng);
result.push_back(dist.ids[(size_t) predecessor]);
proposal_dists.push_back(std::move(dist));
} else {
predecessor = (int32_t) std::distance(path,
std::max_element(path, path + selector_top_k));
result.push_back((llama_token) ids[(size_t) pos * selector_top_k + predecessor]);
}
}
batch.n_tokens = 0;
return;
}
result.reserve((size_t) n_keep);
for (int32_t i = 0; i < n_keep; ++i) {

View File

@ -17,6 +17,7 @@
#include <iomanip>
#include <limits>
#include <map>
#include <sstream>
#include <unordered_map>
#define SPEC_VOCAB_MAX_SIZE_DIFFERENCE 128
@ -1098,6 +1099,8 @@ struct common_speculative {
int last_n_drafted = 0;
int64_t t_step_start_us = 0;
bool last_step_target_only = false;
float draft_temperature = 0.0f;
uint32_t draft_seed = LLAMA_DEFAULT_SEED;
};
static bool common_speculative_stage_chain_matches(
@ -1127,6 +1130,8 @@ static common_params_speculative common_speculative_get_runtime_params(
result.n_min = stage.has_n_min_override() ? stage.n_min : params.n_min;
result.p_min = stage.has_p_min_override() ? stage.p_min : params.p_min;
result.mtp_heads = stage.has_mtp_heads_override() ? stage.mtp_heads : params.mtp_heads;
result.draft_temperature = params.draft_temperature;
result.draft_seed = params.draft_seed;
if (config.type == COMMON_SPECULATIVE_TYPE_SUFFIX) {
result.suffix_min_match_len = stage.has_suffix_min_match_len_override()
@ -1169,6 +1174,8 @@ void common_speculative_prepare_request(common_speculative * spec, common_params
const auto & runtime_stage = use_runtime_stage_overrides ? runtime_stages[i] : spec->configs[i].stage;
common_params_speculative impl_params = common_speculative_get_runtime_params(spec->configs[i], params, runtime_stage);
impl_params.draft_temperature = spec->draft_temperature;
impl_params.draft_seed = spec->draft_seed;
mtp_state->mtp_heads_active = std::max<int32_t>(0, impl_params.mtp_heads);
}
}
@ -1585,6 +1592,8 @@ llama_tokens common_speculative_draft(
auto & impl = spec->impls[i];
const auto & runtime_stage = use_runtime_stage_overrides ? runtime_stages[i] : spec->configs[i].stage;
common_params_speculative impl_params = common_speculative_get_runtime_params(spec->configs[i], params, runtime_stage);
impl_params.draft_temperature = spec->draft_temperature;
impl_params.draft_seed = spec->draft_seed;
if (spec->tuner && spec->tuner->enabled && impl->type == COMMON_SPECULATIVE_TYPE_DFLASH) {
impl_params.n_max = params.n_max;
}
@ -1910,9 +1919,15 @@ common_speculative_draft_result common_speculative_draft_ex(
const llama_tokens & prompt_tgt,
llama_token id_last,
llama_pos draft_base_pos,
llama_seq_id draft_seq_id) {
llama_seq_id draft_seq_id,
const common_params_sampling * sampling) {
common_speculative_draft_result result = {};
if (spec != nullptr) {
spec->draft_temperature = sampling != nullptr ? sampling->temp : 0.0f;
spec->draft_seed = sampling != nullptr ? sampling->seed : LLAMA_DEFAULT_SEED;
}
if (common_speculative_has_type(spec, COMMON_SPECULATIVE_TYPE_MTP)) {
if (!common_speculative_ensure_sequence_hidden(spec, ctx, draft_seq_id, draft_base_pos - 1)) {
LOG_ERR("%s: seq_id=%d MTP hidden state is empty during speculation\n",
@ -1933,6 +1948,13 @@ common_speculative_draft_result common_speculative_draft_ex(
: COMMON_SPECULATIVE_TYPE_NONE;
result.target_only = spec != nullptr && spec->last_step_target_only;
if (spec != nullptr && spec->curr_impl != nullptr &&
spec->curr_impl->type == COMMON_SPECULATIVE_TYPE_DFLASH) {
if (auto * dflash_state = common_speculative_get_dflash_state(spec); dflash_state != nullptr) {
result.proposal_dists = dflash_state->proposal_dists;
}
}
return result;
}
@ -3330,6 +3352,7 @@ int32_t mtp_update_kv_cache(struct llama_context * ctx, const llama_batch& batch
llama_set_mtp_op_type(ctx, MTP_OP_NONE);
return ret;
}
common_speculative_round_result common_speculative_run_round(
common_speculative * spec,
llama_model * model,
@ -3401,14 +3424,21 @@ common_speculative_round_result common_speculative_run_round(
draft_history,
result.sampled_before,
n_past,
seq_id);
seq_id,
&sparams);
auto & draft = draft_result.tokens;
auto & proposal_dists = draft_result.proposal_dists;
const int min_usable_draft = params.get_min_usable_stage_n_min();
if ((int) draft.size() < min_usable_draft || (draft.empty() && !draft_result.target_only)) {
return result;
}
if (!proposal_dists.empty() && proposal_dists.size() != draft.size()) {
result.failed = true;
result.error = "DFlash2 proposal distribution count does not match draft";
return result;
}
if (common_speculative_needs_checkpoint(model)) {
if (!common_speculative_before_draft(
spec,
@ -3446,10 +3476,11 @@ common_speculative_round_result common_speculative_run_round(
result.error = "speculative verify decode failed";
return result;
}
std::vector<llama_token> ids;
try {
ids = common_sampler_sample_and_accept_n(sampler, ctx, verify_indices, draft);
ids = proposal_dists.empty()
? common_sampler_sample_and_accept_n(sampler, ctx, verify_indices, draft)
: common_sampler_sample_and_accept_n(sampler, ctx, verify_indices, draft, proposal_dists);
} catch (const std::exception & e) {
llama_batch_free(verify_batch);
result.failed = true;

View File

@ -24,6 +24,11 @@ using common_speculative_feature_view = llama_spec_feature_view;
static constexpr common_speculative_feature_kind COMMON_SPECULATIVE_FEATURE_NONE = LLAMA_SPEC_FEATURE_NONE;
static constexpr common_speculative_feature_kind COMMON_SPECULATIVE_FEATURE_HIDDEN_STATE = LLAMA_SPEC_FEATURE_HIDDEN_STATE;
struct common_speculative_token_dist {
llama_tokens ids;
std::vector<float> probs;
};
struct common_speculative_checkpoint {
bool valid = false;
int mode = LLAMA_SPEC_CKPT_NONE;
@ -36,6 +41,7 @@ struct common_speculative_checkpoint {
struct common_speculative_draft_result {
llama_tokens tokens;
std::vector<common_speculative_token_dist> proposal_dists; // Sparse proposal distributions populated by stochastic DFlash2
common_speculative_type type = COMMON_SPECULATIVE_TYPE_NONE;
bool target_only = false;
};
@ -134,7 +140,8 @@ common_speculative_draft_result common_speculative_draft_ex(
const llama_tokens & prompt,
llama_token id_last,
llama_pos draft_base_pos = -1,
llama_seq_id draft_seq_id = 0);
llama_seq_id draft_seq_id = 0,
const common_params_sampling * sampling = nullptr);
int common_speculative_get_configured_n_max(const common_speculative * spec);

View File

@ -313,7 +313,9 @@ class Model:
gguf.MODEL_TENSOR.TOKEN_TYPES,
)
)
or not name.endswith(".weight")
or (not name.endswith(".weight") and not (
self.model_arch == gguf.MODEL_ARCH.DFLASH2 and new_name.endswith(".weight")
))
):
data_qtype = gguf.GGMLQuantizationType.F32
@ -644,6 +646,9 @@ class Model:
if chkhsh == "d30d75d9059f1aa2c19359de71047b3ae408c70875e8a3ccf8c5fba56c9d8af4":
# ref: https://huggingface.co/Qwen/Qwen3.5-9B-Instruct
res = "qwen35"
if chkhsh == "d353350c764d8c3b39c763113960e4fb4919bea5fbf208a0e3b22e8469dc7406":
# ref: https://huggingface.co/meta-llama/Llama-4-Scout-17B-16E-Instruct
res = "llama4"
if chkhsh == "99cc61242f7106804ce24fdf3a6451e4a55251078dffd5453c806e11b2310db3":
# ref: https://huggingface.co/Qwen/Qwen3.5-27B
res = "qwen35"
@ -2623,6 +2628,157 @@ class DFlashDraftModel(Qwen3Model):
return tensors
@Model.register("DFlash2DraftModel")
class DFlash2DraftModel(DFlashDraftModel):
"""DFlash 2 sidecar with dynamic convolution and candidate selector tensors."""
model_arch = gguf.MODEL_ARCH.DFLASH2
def _set_vocab_tokenizer_json(self, dir_model: Path, vocab_size: int) -> None:
tokenizer_path = dir_model / "tokenizer.json"
with open(tokenizer_path, "r", encoding="utf-8") as f:
tokenizer_json = json.load(f)
from tokenizers import Tokenizer
tokenizer = Tokenizer.from_file(str(tokenizer_path))
class TokenizerShim:
def encode(self, text: str) -> list[int]:
return tokenizer.encode(text).ids
vocab: dict[str, int] = tokenizer_json["model"]["vocab"]
reverse_vocab = {id_: token for token, id_ in vocab.items()}
added_vocab = {
item["id"]: item
for item in tokenizer_json.get("added_tokens", [])
if isinstance(item.get("id"), int) and isinstance(item.get("content"), str)
}
reverse_vocab.update({id_: item["content"] for id_, item in added_vocab.items()})
assert max(reverse_vocab) < vocab_size
tokpre = self.get_vocab_base_pre(TokenizerShim())
tokens: list[str] = []
toktypes: list[int] = []
for i in range(vocab_size):
token = reverse_vocab.get(i)
if token is None:
tokens.append(f"[PAD{i}]")
toktypes.append(gguf.TokenType.UNUSED)
continue
added_token = added_vocab.get(i)
if added_token is not None:
if not added_token.get("normalized", True):
token = tokenizer.decode(tokenizer.encode(token, add_special_tokens=False).ids)
if added_token.get("special", False) or self.does_token_look_special(token):
toktypes.append(gguf.TokenType.CONTROL)
else:
token = token.replace("\u2581", " ")
toktypes.append(gguf.TokenType.USER_DEFINED)
else:
toktypes.append(gguf.TokenType.NORMAL)
tokens.append(token)
self.gguf_writer.add_tokenizer_model("gpt2")
self.gguf_writer.add_tokenizer_pre(tokpre)
self.gguf_writer.add_token_list(tokens)
self.gguf_writer.add_token_types(toktypes)
special_vocab = gguf.SpecialVocab(dir_model, load_merges=True)
special_vocab.add_to_gguf(self.gguf_writer)
eot_id = next(
(id_ for id_, item in added_vocab.items() if item["content"] == "<|eot|>"),
None,
)
if eot_id is not None:
self.gguf_writer.add_eot_token_id(eot_id)
def set_vocab(self):
target_dir = self._require_target_model_dir()
target_hparams = self._get_target_hparams()
target_raw_hparams = self._get_target_raw_hparams()
target_architectures = target_raw_hparams.get("architectures", [])
if "MuseGlimmerForConditionalGeneration" in target_architectures:
self._set_vocab_tokenizer_json(target_dir, int(target_hparams["vocab_size"]))
else:
super().set_vocab()
if (bos_token_id := target_hparams.get("bos_token_id")) is not None:
self.gguf_writer.add_bos_token_id(int(bos_token_id))
def set_gguf_parameters(self):
dflash_cfg = self.hparams.get("dflash_config")
dflash_cfg = dflash_cfg if isinstance(dflash_cfg, dict) else {}
Qwen3Model.set_gguf_parameters(self)
self.gguf_writer.add_causal_attention(self._causal_attention())
rope_parameters = self.hparams.get("rope_parameters")
if isinstance(rope_parameters, dict) and (rope_theta := rope_parameters.get("rope_theta")) is not None:
self.gguf_writer.add_rope_freq_base(float(rope_theta))
def required(name: str) -> int:
value = dflash_cfg.get(name, self.hparams.get(name))
if value is None:
raise ValueError(f"DFlash2DraftModel conversion requires explicit {name} metadata")
return int(value)
block_size = required("block_size")
mask_token_id = required("mask_token_id")
target_layer_ids = dflash_cfg.get("target_layer_ids")
if target_layer_ids is None:
raise ValueError("DFlash2DraftModel conversion requires target_layer_ids metadata")
target_layers = [int(value) + 1 for value in target_layer_ids]
if not target_layers or any(value <= 0 for value in target_layers):
raise ValueError("DFlash2DraftModel conversion requires target_layer_ids metadata")
if len(set(target_layers)) != len(target_layers):
raise ValueError("DFlash2DraftModel conversion requires unique target_layer_ids metadata")
self.gguf_writer.add_uint32(f"{self.gguf_writer.arch}.block_size", block_size)
self.gguf_writer.add_mask_token_id(mask_token_id)
self.gguf_writer.add_array(f"{self.gguf_writer.arch}.target_layers", target_layers)
self.gguf_writer.add_conv_kernel_size(required("conv_kernel_size"))
self.gguf_writer.add_conv_group_size(required("conv_group_size"))
self.gguf_writer.add_selector_rank(required("selector_rank"))
self.gguf_writer.add_selector_top_k(required("selector_top_k"))
for name, method in (
("output_multiplier", self.gguf_writer.add_logit_scale),
("final_logit_softcapping", self.gguf_writer.add_final_logit_softcapping),
("input_embedding_scale", self.gguf_writer.add_embedding_scale),
):
value = dflash_cfg.get(name, self.hparams.get(name))
if value is not None and (name != "final_logit_softcapping" or float(value) > 0):
method(float(value))
use_sliding_window = self.hparams.get("use_sliding_window")
sliding_window = self.hparams.get("sliding_window")
if use_sliding_window and sliding_window:
layer_types = self.hparams.get("layer_types")
swa_pattern = ([str(value) == "sliding_attention" for value in layer_types]
if layer_types else [True] * self.block_count)
self.gguf_writer.add_sliding_window(int(sliding_window))
self.gguf_writer.add_sliding_window_pattern(swa_pattern)
@staticmethod
def normalize_tensor_name(name: str) -> str:
if name.startswith("candidate_selector."):
name = f"model.{name}"
if name in (
"model.candidate_selector.predecessor_codebook",
"model.candidate_selector.successor_codebook",
):
name += ".weight"
return name
def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iterable[tuple[str, Tensor]]:
top_level_name = name[6:] if name.startswith("model.") else name
if top_level_name == "fc.weight":
return [("fc.weight", data_torch)]
if top_level_name == "hidden_norm.weight":
return [("enc.output_norm.weight", data_torch)]
return super().modify_tensors(data_torch, self.normalize_tensor_name(name), bid)
@Model.register("Qwen3DSparkModel")
class DSparkModel(DFlashDraftModel):
"""Qwen3 DSpark sidecar: DFlash backbone plus a Markov head."""

View File

@ -973,9 +973,11 @@ int main(int argc, char ** argv) {
draft_history,
sampled_before,
n_past,
0);
0,
&sparams);
auto & draft = draft_result.tokens;
auto & proposal_dists = draft_result.proposal_dists;
int max_usable_draft = (int) draft.size();
if (n_predict_budget >= 0 && n_predict_budget != std::numeric_limits<int>::max()) {
max_usable_draft = std::min(max_usable_draft, std::max(0, n_predict_budget - 2));
@ -984,6 +986,9 @@ int main(int argc, char ** argv) {
max_usable_draft = std::min(max_usable_draft, std::max(0, (int) llama_n_batch(ctx) - 1));
if ((int) draft.size() > max_usable_draft) {
draft.resize(max_usable_draft);
if (!proposal_dists.empty()) {
proposal_dists.resize(max_usable_draft);
}
}
const int min_usable_draft = params.speculative.get_min_usable_stage_n_min();
@ -1022,16 +1027,16 @@ int main(int argc, char ** argv) {
LOG_TEE("%s : failed to eval speculative batch\n", __func__);
return 1;
}
std::vector<llama_token> ids;
try {
ids = common_sampler_sample_and_accept_n(ctx_sampling, ctx, verify_indices, draft);
ids = proposal_dists.empty()
? common_sampler_sample_and_accept_n(ctx_sampling, ctx, verify_indices, draft)
: common_sampler_sample_and_accept_n(ctx_sampling, ctx, verify_indices, draft, proposal_dists);
} catch (const std::exception & e) {
llama_batch_free(verify_batch);
LOG_TEE("%s: speculative sampling failed: %s\n", __func__, e.what());
return 1;
}
std::vector<int32_t> accepted_output_indices;
if (!ids.empty()) {
accepted_output_indices.assign(verify_indices.begin(), verify_indices.begin() + ids.size());
@ -1051,7 +1056,6 @@ int main(int argc, char ** argv) {
LOG_TEE("%s: speculative checkpoint restore/commit failed\n", __func__);
return 1;
}
llama_batch_free(verify_batch);
if (!ids.empty()) {

View File

@ -499,6 +499,7 @@ void server_slot::reset() {
prompt_batch_i1 = -1;
n_sent_text = 0;
drafted.clear();
draft_proposal_dists.clear();
spec_target_only = false;
i_batch_dft.clear();
spec_prompt_warmup_failed = false;
@ -3542,8 +3543,10 @@ void server_context::add_sampled_tokens() {
cached_text_tokens,
slot.sampled,
draft_base_pos,
slot.id);
slot.id,
&slot.sparams);
llama_tokens & draft = draft_result.tokens;
auto & proposal_dists = draft_result.proposal_dists;
slot.spec_target_only = draft_result.target_only;
const int n_draft_max = slot.get_n_draft_max();
@ -3556,6 +3559,15 @@ void server_context::add_sampled_tokens() {
SLT_WRN(slot, "draft size %d exceeds max %d, truncating\n", (int)draft.size(), n_draft_max);
}
draft.resize(n_draft_max);
if (!proposal_dists.empty()) {
proposal_dists.resize(n_draft_max);
}
}
if (!proposal_dists.empty() && proposal_dists.size() != draft.size()) {
SLT_WRN(slot, "discarding mismatched DFlash2 proposal distributions (%d != %d)\n",
(int) proposal_dists.size(), (int) draft.size());
proposal_dists.clear();
}
// add the sampled token to the batch
@ -3569,6 +3581,7 @@ void server_context::add_sampled_tokens() {
// fallback to normal decoding
slot.i_batch = slot.i_batch_dft[0];
slot.drafted.clear();
slot.draft_proposal_dists.clear();
slot.i_batch_dft.clear();
} else {
if (slot.spec_target_only) {
@ -3590,6 +3603,7 @@ void server_context::add_sampled_tokens() {
slot.cache_tokens.push_back(draft[i]);
}
slot.drafted = std::move(draft);
slot.draft_proposal_dists = std::move(proposal_dists);
}
}
else {
@ -4260,7 +4274,9 @@ void server_context::speculative_decoding_accept() {
// the accepted tokens from the speculation
std::vector<llama_token> ids;
try {
ids = common_sampler_sample_and_accept_n(slot.ctx_sampling, ctx, slot.i_batch_dft, slot.drafted);
ids = slot.draft_proposal_dists.empty()
? common_sampler_sample_and_accept_n(slot.ctx_sampling, ctx, slot.i_batch_dft, slot.drafted)
: common_sampler_sample_and_accept_n(slot.ctx_sampling, ctx, slot.i_batch_dft, slot.drafted, slot.draft_proposal_dists);
} catch (const std::exception & e) {
LOG_ERROR("speculative sampling failed, releasing slot", {
{"id_slot", slot.id},
@ -4272,6 +4288,7 @@ void server_context::speculative_decoding_accept() {
slot.i_batch = -1;
slot.i_batch_dft.clear();
slot.drafted.clear();
slot.draft_proposal_dists.clear();
continue;
}
@ -4284,6 +4301,7 @@ void server_context::speculative_decoding_accept() {
slot.i_batch_dft.clear();
slot.drafted.clear();
slot.draft_proposal_dists.clear();
slot.n_past += ids.size();
const int64_t t_current = ggml_time_us();
@ -4329,6 +4347,7 @@ void server_context::speculative_decoding_accept() {
slot.i_batch = -1;
slot.i_batch_dft.clear();
slot.drafted.clear();
slot.draft_proposal_dists.clear();
continue;
}
slot.spec_target_only = false;
@ -4958,6 +4977,7 @@ void server_context::update_slots() {
}
slot.cache_tokens.keep_first(slot.cache_tokens.n_tokens() - (int32_t) slot.drafted.size());
slot.drafted.clear();
slot.draft_proposal_dists.clear();
slot.i_batch_dft.clear();
slot.n_past = slot.cache_tokens.n_tokens();
slot.spec_target_only = false;

View File

@ -126,6 +126,7 @@ struct server_slot {
// sampling
llama_token sampled; // in speculative mode, this is the last accepted token
llama_tokens drafted;
std::vector<common_speculative_token_dist> draft_proposal_dists;
bool spec_target_only = false;
json json_schema;

View File

@ -101,6 +101,10 @@ class Keys:
ATTN_LOGIT_SOFTCAPPING = "{arch}.attn_logit_softcapping"
FINAL_LOGIT_SOFTCAPPING = "{arch}.final_logit_softcapping"
ROUTER_LOGIT_SOFTCAPPING = "{arch}.router_logit_softcapping"
CONV_KERNEL_SIZE = "{arch}.conv_kernel_size"
CONV_GROUP_SIZE = "{arch}.conv_group_size"
SELECTOR_RANK = "{arch}.selector_rank"
SELECTOR_TOP_K = "{arch}.selector_top_k"
class Attention:
HEAD_COUNT = "{arch}.attention.head_count"
@ -255,6 +259,7 @@ class MODEL_ARCH(IntEnum):
GEMMA4 = auto()
GEMMA4_MTP = auto()
DFLASH = auto()
DFLASH2 = auto()
DFLASH_DRAFT = auto()
STARCODER2 = auto()
MAMBA = auto()
@ -406,6 +411,13 @@ class MODEL_TENSOR(IntEnum):
DSPARK_MARKOV_W1 = auto() # DSpark Markov lookup matrix
DSPARK_MARKOV_W2 = auto() # DSpark Markov projection matrix
DSPARK_CONF_PROJ = auto() # DSpark confidence projection
DFLASH_ATTN_CONV_BASE = auto()
DFLASH_ATTN_CONV_PROJ = auto()
DFLASH_FFN_CONV_BASE = auto()
DFLASH_FFN_CONV_PROJ = auto()
DFLASH_SELECTOR_PREV = auto()
DFLASH_SELECTOR_NEXT = auto()
DFLASH_SELECTOR_HIDDEN = auto()
ATTN_KV = auto()
ATTN_KV_NORM = auto()
ATTN_OUT_A = auto()
@ -476,6 +488,7 @@ MODEL_ARCH_NAMES: dict[MODEL_ARCH, str] = {
MODEL_ARCH.GEMMA4: "gemma4",
MODEL_ARCH.GEMMA4_MTP: "gemma4_mtp",
MODEL_ARCH.DFLASH: "dflash",
MODEL_ARCH.DFLASH2: "dflash",
MODEL_ARCH.DFLASH_DRAFT: "dflash-draft",
MODEL_ARCH.STARCODER2: "starcoder2",
MODEL_ARCH.MAMBA: "mamba",
@ -638,6 +651,13 @@ TENSOR_NAMES: dict[MODEL_TENSOR, str] = {
MODEL_TENSOR.DSPARK_MARKOV_W1: "markov_w1",
MODEL_TENSOR.DSPARK_MARKOV_W2: "markov_w2",
MODEL_TENSOR.DSPARK_CONF_PROJ: "conf_proj",
MODEL_TENSOR.DFLASH_ATTN_CONV_BASE: "blk.{bid}.attn_conv_base",
MODEL_TENSOR.DFLASH_ATTN_CONV_PROJ: "blk.{bid}.attn_conv_proj",
MODEL_TENSOR.DFLASH_FFN_CONV_BASE: "blk.{bid}.ffn_conv_base",
MODEL_TENSOR.DFLASH_FFN_CONV_PROJ: "blk.{bid}.ffn_conv_proj",
MODEL_TENSOR.DFLASH_SELECTOR_PREV: "selector_predecessor",
MODEL_TENSOR.DFLASH_SELECTOR_NEXT: "selector_successor",
MODEL_TENSOR.DFLASH_SELECTOR_HIDDEN: "selector_hidden",
# openPangu-2.0
MODEL_TENSOR.INDEXER_K_NORM: "blk.{bid}.attn_indexer_k_norm",
MODEL_TENSOR.INDEXER_PROJ: "blk.{bid}.attn_indexer_weights_proj",
@ -1571,6 +1591,31 @@ MODEL_TENSORS: dict[MODEL_ARCH, list[MODEL_TENSOR]] = {
MODEL_TENSOR.DSPARK_MARKOV_W2,
MODEL_TENSOR.DSPARK_CONF_PROJ,
],
MODEL_ARCH.DFLASH2: [
MODEL_TENSOR.TOKEN_EMBD,
MODEL_TENSOR.OUTPUT_NORM,
MODEL_TENSOR.OUTPUT,
MODEL_TENSOR.ATTN_NORM,
MODEL_TENSOR.ATTN_Q,
MODEL_TENSOR.ATTN_Q_NORM,
MODEL_TENSOR.ATTN_K,
MODEL_TENSOR.ATTN_K_NORM,
MODEL_TENSOR.ATTN_V,
MODEL_TENSOR.ATTN_OUT,
MODEL_TENSOR.FFN_NORM,
MODEL_TENSOR.FFN_GATE,
MODEL_TENSOR.FFN_DOWN,
MODEL_TENSOR.FFN_UP,
MODEL_TENSOR.DFLASH_FC,
MODEL_TENSOR.DFLASH_HIDDEN_NORM,
MODEL_TENSOR.DFLASH_ATTN_CONV_BASE,
MODEL_TENSOR.DFLASH_ATTN_CONV_PROJ,
MODEL_TENSOR.DFLASH_FFN_CONV_BASE,
MODEL_TENSOR.DFLASH_FFN_CONV_PROJ,
MODEL_TENSOR.DFLASH_SELECTOR_PREV,
MODEL_TENSOR.DFLASH_SELECTOR_NEXT,
MODEL_TENSOR.DFLASH_SELECTOR_HIDDEN,
],
MODEL_ARCH.BITNET: [
MODEL_TENSOR.ATTN_Q,
MODEL_TENSOR.ATTN_K,

View File

@ -749,6 +749,18 @@ class GGUFWriter:
def add_final_logit_softcapping(self, value: float) -> None:
self.add_float32(Keys.LLM.FINAL_LOGIT_SOFTCAPPING.format(arch=self.arch), value)
def add_conv_kernel_size(self, value: int) -> None:
self.add_uint32(Keys.LLM.CONV_KERNEL_SIZE.format(arch=self.arch), value)
def add_conv_group_size(self, value: int) -> None:
self.add_uint32(Keys.LLM.CONV_GROUP_SIZE.format(arch=self.arch), value)
def add_selector_rank(self, value: int) -> None:
self.add_uint32(Keys.LLM.SELECTOR_RANK.format(arch=self.arch), value)
def add_selector_top_k(self, value: int) -> None:
self.add_uint32(Keys.LLM.SELECTOR_TOP_K.format(arch=self.arch), value)
def add_expert_count(self, count: int) -> None:
self.add_uint32(Keys.LLM.EXPERT_COUNT.format(arch=self.arch), count)

View File

@ -101,6 +101,15 @@ class TensorNameMap:
MODEL_TENSOR.MHC_MERGE_GAMMA: (
"model.merge_mhc_module.norm_gamma",
),
MODEL_TENSOR.DFLASH_SELECTOR_PREV: (
"model.candidate_selector.predecessor_codebook",
),
MODEL_TENSOR.DFLASH_SELECTOR_NEXT: (
"model.candidate_selector.successor_codebook",
),
MODEL_TENSOR.DFLASH_SELECTOR_HIDDEN: (
"model.candidate_selector.hidden_projection",
),
}
block_mappings_cfg: dict[MODEL_TENSOR, tuple[str, ...]] = {
@ -441,6 +450,19 @@ class TensorNameMap:
"model.layers.{bid}.self_attn.attention_sink_bias", # MiMo DFlash
),
MODEL_TENSOR.DFLASH_ATTN_CONV_BASE: (
"model.layers.{bid}.attention_conv.base_kernel",
),
MODEL_TENSOR.DFLASH_ATTN_CONV_PROJ: (
"model.layers.{bid}.attention_conv.kernel_projection",
),
MODEL_TENSOR.DFLASH_FFN_CONV_BASE: (
"model.layers.{bid}.mlp_conv.base_kernel",
),
MODEL_TENSOR.DFLASH_FFN_CONV_PROJ: (
"model.layers.{bid}.mlp_conv.kernel_projection",
),
MODEL_TENSOR.ROPE_FREQS: (
"language_model.encoder.layers.{bid}.self_attention.rotary_emb.inv_freq", # persimmon
),

View File

@ -1144,6 +1144,15 @@ extern "C" {
// Returns LLAMA_TOKEN_NULL if argmax is not available (falls back to logits path).
LLAMA_API llama_token llama_get_dflash_draft_token_ith(struct llama_context * ctx, int32_t i);
// Copy DFlash2 selector lattice after a draft decode with top_k*top_k scores
// and top_k candidate IDs per position.
LLAMA_API int32_t llama_get_dflash_draft_lattice_top_k(struct llama_context * ctx);
LLAMA_API int32_t llama_get_dflash_draft_lattice_n_positions(struct llama_context * ctx);
LLAMA_API bool llama_copy_dflash_draft_lattice(
struct llama_context * ctx,
float * scores, size_t score_count,
int32_t * ids, size_t id_count);
// Get all output token embeddings.
// when pooling_type == LLAMA_POOLING_TYPE_NONE or when using a generative model,
// the embeddings for which llama_batch.logits[i] != 0 are stored contiguously

View File

@ -4,6 +4,65 @@
#include <cmath>
static ggml_tensor * build_dflash2_conv(
llm_build_context & g,
ggml_tensor * hidden,
ggml_tensor * dynamic,
ggml_tensor * base,
int side) {
const int64_t hidden_size = hidden->ne[0];
const int64_t n_tokens = hidden->ne[1];
const int64_t kernel_size = g.hparams.dflash_conv_kernel_size;
const int64_t group_size = g.hparams.dflash_conv_group_size;
const int64_t n_groups = hidden_size / group_size;
const int64_t n_blocks = 1;
const int64_t block_size = n_tokens / n_blocks;
GGML_ASSERT(n_tokens > 0 && n_tokens % n_blocks == 0);
GGML_ASSERT(hidden_size % group_size == 0);
GGML_ASSERT(dynamic != nullptr && base != nullptr);
GGML_ASSERT(side >= 0 && side < 2);
ggml_context * ctx0 = g.ctx0;
hidden = ggml_cont_2d(ctx0, hidden, hidden_size, n_tokens);
dynamic = ggml_cont_2d(ctx0, dynamic, dynamic->ne[0], n_tokens);
ggml_tensor * blocks = ggml_reshape_3d(ctx0, hidden, hidden_size, block_size, n_blocks);
ggml_tensor * grouped = ggml_reshape_3d(ctx0, hidden, group_size, n_groups, n_tokens);
ggml_tensor * coeffs = ggml_reshape_4d(ctx0, dynamic, n_groups, kernel_size, 2, n_tokens);
ggml_tensor * coeffs_side = ggml_view_3d(ctx0, coeffs, n_groups, kernel_size, n_tokens,
coeffs->nb[1], coeffs->nb[3], side * coeffs->nb[2]);
ggml_tensor * result = nullptr;
for (int64_t tap = 0; tap < kernel_size; ++tap) {
ggml_tensor * values = blocks;
if (tap > 0) {
ggml_tensor * zeros = ggml_fill(ctx0,
ggml_new_tensor_3d(ctx0, hidden->type, hidden_size, std::min(tap, block_size), n_blocks), 0.0f);
if (tap < block_size) {
ggml_tensor * previous = ggml_view_3d(ctx0, blocks, hidden_size, block_size - tap, n_blocks,
blocks->nb[1], blocks->nb[2], 0);
values = ggml_concat(ctx0, zeros, previous, 1);
} else {
values = zeros;
}
}
values = ggml_reshape_2d(ctx0, values, hidden_size, n_tokens);
ggml_tensor * coeff = ggml_view_2d(ctx0, coeffs_side, n_groups, n_tokens,
coeffs_side->nb[2], tap * coeffs_side->nb[1]);
coeff = ggml_cont(ctx0, coeff);
coeff = ggml_reshape_3d(ctx0, coeff, 1, n_groups, n_tokens);
coeff = ggml_reshape_2d(ctx0, ggml_repeat(ctx0, coeff, grouped), hidden_size, n_tokens);
ggml_tensor * base_tap = ggml_view_1d(ctx0, base, hidden_size,
tap * base->nb[1] + side * base->nb[2]);
ggml_tensor * weight = ggml_add(ctx0, coeff, ggml_repeat(ctx0, base_tap, hidden));
ggml_tensor * term = ggml_mul(ctx0, weight, values);
result = result ? ggml_add(ctx0, result, term) : term;
}
return result;
}
ggml_tensor * llm_build_context::build_dspark_logits(
llm_build_context & llm,
ggml_tensor * base_logits,
@ -119,7 +178,8 @@ ggml_cgraph * llm_build_context::build_dflash_kv_cache() {
}
ggml_tensor * fused_target = llm_build_lora_mm(lctx, ctx0, model.dflash_fc, target_features);
fused_target = llm_build_norm(ctx0, fused_target, hparams, model.dflash_hidden_norm, nullptr, LLM_NORM_RMS, cb, -1);
fused_target = llm_build_norm(ctx0, fused_target, hparams,
model.dflash_hidden_norm, nullptr, LLM_NORM_RMS, cb, -1);
cb(fused_target, "dflash_kv_fused_target", -1);
if (hparams.dflash_dsv4) {
@ -302,6 +362,7 @@ ggml_cgraph * llm_build_context::build_dflash() {
const int64_t n_embd_head_k = hparams.n_embd_head_k(0);
const int64_t n_embd_head_v = hparams.n_embd_head_v(0);
const int64_t n_target_features = hparams.dflash_n_target_features;
const bool is_dflash2 = model.arch == LLM_ARCH_DFLASH2;
const int64_t ctx_len = lctx.dflash.visible_cross_ctx > 0
? (int64_t) lctx.dflash.visible_cross_ctx
: std::max<int64_t>(1, (int64_t) cparams.n_ctx - (int64_t) hparams.dflash_block_size);
@ -367,7 +428,8 @@ ggml_cgraph * llm_build_context::build_dflash() {
ggml_tensor * inpL = llm_build_inp_embd(ctx0, lctx, hparams, batch, tok_embd, cb);
ggml_tensor * inp_pos = build_inp_pos();
ggml_tensor * inp_out_ids = (n_tokens > 1 && n_outputs < n_tokens) ? build_inp_out_ids() : nullptr;
ggml_tensor * inp_out_ids = (!llm_arch_requires_all_graph_output_rows(model.arch) &&
n_tokens > 1 && n_outputs < n_tokens) ? build_inp_out_ids() : nullptr;
const float kq_scale = 1.0f / std::sqrt((float) n_embd_head_k);
@ -378,6 +440,13 @@ ggml_cgraph * llm_build_context::build_dflash() {
cb(cur, "attn_norm", il);
ggml_tensor * input_normed = cur;
ggml_tensor * attn_dynamic = nullptr;
if (is_dflash2) {
attn_dynamic = llm_build_lora_mm(lctx, ctx0, model.layers[il].dflash_attn_conv_proj, cur);
cur = build_dflash2_conv(*this, cur, attn_dynamic, model.layers[il].dflash_attn_conv_base, 0);
cb(cur, "dflash2_attn_conv_in", il);
}
ggml_tensor * Qcur = llm_build_lora_mm(lctx, ctx0, model.layers[il].wq, cur);
ggml_tensor * Kcur_noise = llm_build_lora_mm(lctx, ctx0, model.layers[il].wk, cur);
ggml_tensor * Vcur_noise = llm_build_lora_mm(lctx, ctx0, model.layers[il].wv, cur);
@ -502,6 +571,11 @@ ggml_cgraph * llm_build_context::build_dflash() {
if (model.layers[il].bo) { cur = ggml_add(ctx0, cur, model.layers[il].bo); }
cb(cur, "kqv_out", il);
if (is_dflash2) {
cur = build_dflash2_conv(*this, cur, attn_dynamic, model.layers[il].dflash_attn_conv_base, 1);
cb(cur, "dflash2_attn_conv_out", il);
}
cur = ggml_add(ctx0, cur, inpSA);
cb(cur, "attn_residual", il);
@ -511,8 +585,16 @@ ggml_cgraph * llm_build_context::build_dflash() {
}
ggml_tensor * ffn_residual = cur;
cur = llm_build_norm(ctx0, cur, hparams, model.layers[il].attn_post_norm, nullptr, LLM_NORM_RMS, cb, il);
cb(cur, "attn_post_norm", il);
ggml_tensor * ffn_dynamic = nullptr;
if (is_dflash2) {
cur = llm_build_norm(ctx0, cur, hparams, model.layers[il].ffn_norm, nullptr, LLM_NORM_RMS, cb, il);
ffn_dynamic = llm_build_lora_mm(lctx, ctx0, model.layers[il].dflash_ffn_conv_proj, cur);
cur = build_dflash2_conv(*this, cur, ffn_dynamic, model.layers[il].dflash_ffn_conv_base, 0);
cb(cur, "dflash2_ffn_conv_in", il);
} else {
cur = llm_build_norm(ctx0, cur, hparams, model.layers[il].attn_post_norm, nullptr, LLM_NORM_RMS, cb, il);
cb(cur, "attn_post_norm", il);
}
cur = llm_build_ffn(ctx0, lctx, nullptr, cur,
model.layers[il].ffn_up, nullptr, nullptr,
@ -522,6 +604,11 @@ ggml_cgraph * llm_build_context::build_dflash() {
LLM_FFN_SILU, LLM_FFN_PAR, cb, il, gf, false, false);
cb(cur, "ffn_out", il);
if (is_dflash2) {
cur = build_dflash2_conv(*this, cur, ffn_dynamic, model.layers[il].dflash_ffn_conv_base, 1);
cb(cur, "dflash2_ffn_conv_out", il);
}
cur = ggml_add(ctx0, cur, ffn_residual);
cb(cur, "l_out", il);
@ -530,6 +617,16 @@ ggml_cgraph * llm_build_context::build_dflash() {
GGML_ASSERT(model.output_mtp != nullptr);
ggml_tensor * result = build_output(lctx, ctx0, inpL, model.output_mtp, model.output_norm, cb);
if (is_dflash2) {
if (hparams.f_logit_scale != 0.0f) {
result = ggml_scale(ctx0, result, hparams.f_logit_scale);
}
if (hparams.f_final_logit_softcapping > 0.0f) {
result = ggml_softcap(ctx0, result,
1.0f / hparams.f_final_logit_softcapping,
hparams.f_final_logit_softcapping);
}
}
if (lctx.dflash.dspark) {
cb(result, "dflash_base_result_output", -1);
} else {
@ -538,6 +635,70 @@ ggml_cgraph * llm_build_context::build_dflash() {
ggml_build_forward_expand(gf, result);
lctx.dflash.draft_tokens_tensor = nullptr;
lctx.dflash.draft_lattice_tensor = nullptr;
lctx.dflash.draft_lattice_ids_tensor = nullptr;
lctx.dflash.draft_lattice_top_k = 0;
if (is_dflash2) {
const int64_t top_k = hparams.dflash_selector_top_k;
const int64_t rank = hparams.dflash_selector_rank;
const int64_t lattice_width = top_k * top_k;
const int64_t result_tokens = result->ne[1];
GGML_ASSERT(lctx.inp_tokens != nullptr);
ggml_tensor * candidates = ggml_top_k(ctx0, result, top_k);
candidates = ggml_cont_2d(ctx0, candidates, top_k, result_tokens);
ggml_tensor * logits_rows = ggml_reshape_3d(ctx0, result, 1, result->ne[0], result_tokens);
ggml_tensor * unary_rows = ggml_get_rows(ctx0, logits_rows, candidates);
ggml_tensor * unary = ggml_reshape_2d(ctx0,
unary_rows, top_k, result_tokens);
ggml_tensor * hidden = llm_build_norm(ctx0, inpL, hparams, model.output_norm, nullptr, LLM_NORM_RMS, cb, -1);
hidden = llm_build_lora_mm(lctx, ctx0, model.dflash_selector_hidden, hidden);
ggml_tensor * packed = ggml_fill(ctx0,
ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, lattice_width, 1), 0.0f);
const int64_t block_size = std::min<int64_t>(result_tokens, hparams.dflash_block_size);
ggml_tensor * anchor = ggml_view_1d(ctx0, lctx.inp_tokens, 1, 0);
for (int64_t pos = 1; pos < block_size; ++pos) {
ggml_tensor * ids = ggml_cont_2d(ctx0,
ggml_view_2d(ctx0, candidates, top_k, 1, candidates->nb[1], pos * candidates->nb[1]),
top_k, 1);
ggml_tensor * unary_pos = ggml_cont_2d(ctx0,
ggml_view_2d(ctx0, unary, top_k, 1, unary->nb[1], pos * unary->nb[1]),
top_k, 1);
ggml_tensor * successor = ggml_get_rows(ctx0, model.dflash_selector_next, ids);
ggml_tensor * hidden_pos = ggml_cont_2d(ctx0,
ggml_view_2d(ctx0, hidden, rank, 1, hidden->nb[1], pos * hidden->nb[1]), rank, 1);
ggml_tensor * predecessor = pos == 1
? ggml_get_rows(ctx0, model.dflash_selector_prev, anchor)
: ggml_get_rows(ctx0, model.dflash_selector_prev,
ggml_view_2d(ctx0, candidates, top_k, 1, candidates->nb[1], (pos - 1) * candidates->nb[1]));
ggml_tensor * hidden_repeat = ggml_repeat(ctx0, hidden_pos, predecessor);
ggml_tensor * conditioned = ggml_mul(ctx0, predecessor, hidden_repeat);
ggml_tensor * scores = ggml_mul_mat(ctx0, successor, conditioned);
if (pos == 1) {
scores = ggml_repeat_4d(ctx0, scores, top_k, top_k, 1, 1);
}
ggml_tensor * unary_repeat = ggml_repeat(ctx0,
ggml_reshape_3d(ctx0, unary_pos, top_k, 1, 1), scores);
scores = ggml_add(ctx0, scores, unary_repeat);
ggml_tensor * row = ggml_reshape_2d(ctx0, scores, lattice_width, 1);
packed = ggml_concat(ctx0, packed, row, 1);
}
cb(packed, "dflash2_lattice", -1);
ggml_build_forward_expand(gf, packed);
cb(candidates, "dflash2_lattice_ids", -1);
ggml_set_output(candidates);
ggml_build_forward_expand(gf, candidates);
lctx.dflash.draft_lattice_tensor = packed;
lctx.dflash.draft_lattice_ids_tensor = candidates;
lctx.dflash.draft_lattice_top_k = (int32_t) top_k;
}
ggml_tensor * draft_tokens = nullptr;
if (lctx.dflash.dspark) {
result = build_dspark_logits(*this, result, lctx.inp_tokens, &draft_tokens);

View File

@ -87,6 +87,7 @@ static const std::map<llm_arch, const char *> LLM_ARCH_NAMES = {
{ LLM_ARCH_GEMMA4, "gemma4" },
{ LLM_ARCH_GEMMA4_MTP, "gemma4_mtp" },
{ LLM_ARCH_DFLASH, "dflash" },
{ LLM_ARCH_DFLASH2, "dflash" },
{ LLM_ARCH_DFLASH_DRAFT, "dflash-draft" },
{ LLM_ARCH_GEMMA4_ASSISTANT,"gemma4-assistant" },
{ LLM_ARCH_OPENPANGU, "openpangu" },
@ -174,6 +175,10 @@ static const std::map<llm_kv, const char *> LLM_KV_NAMES = {
{ LLM_KV_DFLASH_N_TARGET_FEATURES, "%s.dflash.n_target_features" },
{ LLM_KV_DFLASH_BACKBONE_ROTARY_BASE, "%s.dflash.backbone_rotary_base" },
{ LLM_KV_DFLASH_LAGUNA, "%s.dflash.laguna" },
{ LLM_KV_DFLASH_CONV_KERNEL_SIZE, "%s.conv_kernel_size" },
{ LLM_KV_DFLASH_CONV_GROUP_SIZE, "%s.conv_group_size" },
{ LLM_KV_DFLASH_SELECTOR_RANK, "%s.selector_rank" },
{ LLM_KV_DFLASH_SELECTOR_TOP_K, "%s.selector_top_k" },
{ LLM_KV_ATTENTION_HEAD_COUNT, "%s.attention.head_count" },
{ LLM_KV_ATTENTION_HEAD_COUNT_KV, "%s.attention.head_count_kv" },
@ -328,5 +333,9 @@ bool llm_arch_is_hybrid(const llm_arch & arch) {
}
bool llm_arch_is_dflash_family(const llm_arch & arch) {
return arch == LLM_ARCH_DFLASH || arch == LLM_ARCH_DFLASH_DRAFT;
return arch == LLM_ARCH_DFLASH || arch == LLM_ARCH_DFLASH2 || arch == LLM_ARCH_DFLASH_DRAFT;
}
bool llm_arch_requires_all_graph_output_rows(const llm_arch & arch) {
return arch == LLM_ARCH_DFLASH2;
}

View File

@ -85,6 +85,7 @@ enum llm_arch {
LLM_ARCH_GEMMA4,
LLM_ARCH_GEMMA4_MTP,
LLM_ARCH_DFLASH,
LLM_ARCH_DFLASH2,
LLM_ARCH_DFLASH_DRAFT,
LLM_ARCH_GEMMA4_ASSISTANT,
LLM_ARCH_OPENPANGU,
@ -157,6 +158,10 @@ enum llm_kv {
LLM_KV_DFLASH_N_TARGET_FEATURES,
LLM_KV_DFLASH_BACKBONE_ROTARY_BASE,
LLM_KV_DFLASH_LAGUNA,
LLM_KV_DFLASH_CONV_KERNEL_SIZE,
LLM_KV_DFLASH_CONV_GROUP_SIZE,
LLM_KV_DFLASH_SELECTOR_RANK,
LLM_KV_DFLASH_SELECTOR_TOP_K,
LLM_KV_ATTENTION_HEAD_COUNT,
LLM_KV_ATTENTION_HEAD_COUNT_KV,
@ -438,6 +443,13 @@ enum llm_tensor {
LLM_TENSOR_DSPARK_MARKOV_W1,
LLM_TENSOR_DSPARK_MARKOV_W2,
LLM_TENSOR_DSPARK_CONF_PROJ,
LLM_TENSOR_DFLASH_ATTN_CONV_BASE,
LLM_TENSOR_DFLASH_ATTN_CONV_PROJ,
LLM_TENSOR_DFLASH_FFN_CONV_BASE,
LLM_TENSOR_DFLASH_FFN_CONV_PROJ,
LLM_TENSOR_DFLASH_SELECTOR_PREV,
LLM_TENSOR_DFLASH_SELECTOR_NEXT,
LLM_TENSOR_DFLASH_SELECTOR_HIDDEN,
// openPangu-2.0
LLM_TENSOR_ATTN_QA_CONV, // MoME causal conv on q-lora latent
@ -469,5 +481,7 @@ const char * llama_model_arch_name(llm_arch arch);
bool llm_arch_is_recurrent(const llm_arch & arch);
bool llm_arch_is_hybrid(const llm_arch & arch);
bool llm_arch_is_dflash_family(const llm_arch & arch);
// DFlash2 consumes every anchor/mask row to construct its selector lattice
bool llm_arch_requires_all_graph_output_rows(const llm_arch & arch);
llm_tensor llm_tensor_type(llm_arch arch, const std::string & tensor_name, int il);

View File

@ -2884,6 +2884,7 @@ ggml_cgraph * llm_build_context::llama_build_graph(
result = llm.build_gemma4_mtp();
} break;
case LLM_ARCH_DFLASH:
case LLM_ARCH_DFLASH2:
case LLM_ARCH_DFLASH_DRAFT:
{
result = llm.build_dflash();

View File

@ -448,6 +448,13 @@ struct llama_context {
struct capture_state {
std::vector<int32_t> layer_ids;
std::vector<std::vector<float>> layer_rows;
struct capture_chunk {
int32_t row_offset = 0;
int32_t row_count = 0;
size_t byte_offset = 0;
ggml_type type = GGML_TYPE_COUNT;
};
std::vector<std::vector<capture_chunk>> layer_chunks;
std::vector<int32_t> layer_rows_written;
int32_t row_count = 0;
int32_t row_width = 0;
@ -477,6 +484,11 @@ struct llama_context {
std::vector<llama_token> draft_tokens;
struct ggml_tensor * draft_tokens_tensor = nullptr;
struct ggml_tensor * draft_lattice_tensor = nullptr;
struct ggml_tensor * draft_lattice_ids_tensor = nullptr;
std::vector<float> draft_lattice;
std::vector<int32_t> draft_lattice_ids;
int32_t draft_lattice_top_k = 0;
};
dflash_runtime dflash;
using dflash_capture_state = dflash_runtime::capture_state;

View File

@ -380,7 +380,6 @@ bool llama_prepare_dflash_graph_inputs(
const int32_t cross_ctx = lctx.dflash.visible_cross_ctx > 0
? lctx.dflash.visible_cross_ctx
: std::max<int32_t>(1, (int32_t) lctx.cparams.n_ctx - (int32_t) lctx.model.hparams.dflash_block_size);
const bool is_dsv4 = lctx.model.hparams.dflash_dsv4;
ggml_tensor * kq_mask = lctx.dflash.kv.kq_mask_tensor;
ggml_tensor * kq_mask_swa = lctx.dflash.kv.kq_mask_swa_tensor;
@ -695,8 +694,10 @@ bool llama_prepare_dflash_graph_inputs(
for (int32_t k = cross_ctx; k < cross_ctx + (int32_t) n_tokens; ++k) {
const int32_t block_k = k - cross_ctx;
// DSV4 Dspark rows see the complete current block, standard DFlash is causal.
if ((is_dsv4 || block_k <= (int32_t) j) && ((int32_t) j - block_k) < swa_window) {
// Follow the draft model's attention contract. DFlash2 proposal blocks are
// non-causal unless the model metadata explicitly requests causal attention.
if ((!lctx.cparams.causal_attn || block_k <= (int32_t) j) &&
((int32_t) j - block_k) < swa_window) {
row[k] = h_zero;
}
}
@ -720,10 +721,11 @@ bool llama_prepare_dflash_graph_inputs(
for (int32_t k = cross_ctx; k < cross_ctx + (int32_t) n_tokens; ++k) {
const int32_t block_k = k - cross_ctx;
// intra-block draft tokens are contiguous from draft_pos_base, so the
// SWA distance is (j - block_k); apply the same window bound as the
// cross-context section above (causal AND within n_swa).
if ((is_dsv4 || block_k <= (int32_t) j) && ((int32_t) j - block_k) < swa_window) {
// Intra-block draft tokens are contiguous from draft_pos_base, so the
// SWA distance is (j - block_k); apply the model's causal setting and
// the same window bound as the cross-context section above.
if ((!lctx.cparams.causal_attn || block_k <= (int32_t) j) &&
((int32_t) j - block_k) < swa_window) {
row[k] = 0.0f;
}
}

View File

@ -1737,6 +1737,42 @@ void llm_load_hparams(
default: model.type = e_model::MODEL_UNKNOWN;
}
} break;
case LLM_ARCH_DFLASH2:
{
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;
ml.get_key(LLM_KV_FINAL_LOGIT_SOFTCAPPING, hparams.f_final_logit_softcapping, false);
ml.get_key(LLM_KV_EMBEDDING_SCALE, hparams.f_embedding_scale, false);
ml.get_key("dflash.block_size", hparams.dflash_block_size);
ml.get_key("dflash.conv_kernel_size", hparams.dflash_conv_kernel_size);
ml.get_key("dflash.conv_group_size", hparams.dflash_conv_group_size);
ml.get_key("dflash.selector_rank", hparams.dflash_selector_rank);
ml.get_key("dflash.selector_top_k", hparams.dflash_selector_top_k);
ml.get_key(LLM_KV_TOKENIZER_MASK_ID, hparams.dflash_mask_token_id);
ml.get_key(LLM_KV_ATTENTION_CAUSAL, hparams.causal_attn, false);
load_dflash_target_layer_ids(
ml,
LLM_KV(model.arch)(LLM_KV_DFLASH_TARGET_LAYERS),
hparams,
true);
for (uint32_t i = 0; i < hparams.dflash_n_target_layers; ++i) {
if (hparams.dflash_target_layer_ids[i] == 0) {
throw std::runtime_error("dflash2: target_layers must use one-based IDs");
}
--hparams.dflash_target_layer_ids[i];
}
hparams.dflash_n_target_features = hparams.n_embd * hparams.dflash_n_target_layers;
if (hparams.dflash_selector_top_k == 0 ||
hparams.n_embd < hparams.dflash_selector_top_k * (hparams.dflash_selector_top_k + 1)) {
throw std::runtime_error("dflash2: hidden size is too small for selector top-k lattice");
}
ml.get_key(LLM_KV_ATTENTION_SLIDING_WINDOW, hparams.n_swa, false);
ml.get_key_or_arr(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, hparams.swa_layers, hparams.n_layer, false);
hparams.n_layer_kv_from_start = hparams.n_layer;
hparams.dflash_dsv4 = false;
model.type = e_model::MODEL_UNKNOWN;
} break;
case LLM_ARCH_DFLASH:
case LLM_ARCH_DEEPSEEK4:
case LLM_ARCH_GLM_DSA:

View File

@ -170,6 +170,10 @@ struct llama_hparams {
uint32_t dflash_n_target_features = 0;
uint32_t dflash_n_target_layers = 0;
uint32_t dflash_target_layer_ids[8] = {};
uint32_t dflash_conv_kernel_size = 0;
uint32_t dflash_conv_group_size = 0;
uint32_t dflash_selector_rank = 0;
uint32_t dflash_selector_top_k = 0;
float dflash_backbone_rotary_base = 0.0f;
bool dflash_laguna = false;
bool dflash_dsv4 = false;
@ -197,6 +201,10 @@ struct llama_hparams {
if (this->dflash_mask_token_id != other.dflash_mask_token_id) return true;
if (this->dflash_n_target_features != other.dflash_n_target_features) return true;
if (this->dflash_n_target_layers != other.dflash_n_target_layers) return true;
if (this->dflash_conv_kernel_size != other.dflash_conv_kernel_size) return true;
if (this->dflash_conv_group_size != other.dflash_conv_group_size) return true;
if (this->dflash_selector_rank != other.dflash_selector_rank) return true;
if (this->dflash_selector_top_k != other.dflash_selector_top_k) return true;
if (this->dflash_laguna != other.dflash_laguna) return true;
if (this->dflash_dsv4 != other.dflash_dsv4) return true;
if (this->n_layer != other.n_layer) return true;

View File

@ -105,6 +105,8 @@ struct create_tensors_helper : public create_tensors_helper_interface {
bool create_dflash_tensors(const LLM_TN & tn);
bool create_dflash2_tensors(const LLM_TN & tn);
bool create_dflash_dsv4_tensors(const LLM_TN & tn);
bool create_starcoder2_tensors(const LLM_TN & tn);
@ -2467,6 +2469,65 @@ bool create_tensors_helper::create_dflash_tensors(const LLM_TN & tn) {
return use_mmap_buffer;
}
bool create_tensors_helper::create_dflash2_tensors(const LLM_TN & tn) {
LOADING_PRELUDE
const bool use_split_ctx = model.split_mode == LLAMA_SPLIT_MODE_GRAPH || model.split_mode == LLAMA_SPLIT_MODE_ATTN;
const int64_t kernel_size = hparams.dflash_conv_kernel_size;
const int64_t group_size = hparams.dflash_conv_group_size;
const int64_t rank = hparams.dflash_selector_rank;
const int64_t top_k = hparams.dflash_selector_top_k;
if (kernel_size <= 0 || group_size <= 0 || rank <= 0 || top_k <= 0 ||
n_embd % group_size != 0 || n_embd < top_k * (top_k + 1)) {
throw std::runtime_error("invalid DFlash2 convolution or selector dimensions");
}
model.tok_embd = nullptr;
model.output = nullptr;
model.output_mtp = nullptr;
model.dflash_fc = create_tensor(ctx_output, tn(LLM_TENSOR_DFLASH_FC, "weight"),
{(int64_t) hparams.dflash_n_target_features, n_embd}, 0);
model.output_norm = create_tensor(ctx_output, tn(LLM_TENSOR_OUTPUT_NORM, "weight"), {n_embd}, 0);
model.dflash_hidden_norm = create_tensor(ctx_output, tn(LLM_TENSOR_DFLASH_HIDDEN_NORM, "weight"), {n_embd}, 0);
model.dflash_selector_prev = create_tensor(ctx_output,
tn(LLM_TENSOR_DFLASH_SELECTOR_PREV, "weight"), {rank, n_vocab}, 0);
model.dflash_selector_next = create_tensor(ctx_output,
tn(LLM_TENSOR_DFLASH_SELECTOR_NEXT, "weight"), {rank, n_vocab}, 0);
model.dflash_selector_hidden = create_tensor(ctx_output,
tn(LLM_TENSOR_DFLASH_SELECTOR_HIDDEN, "weight"), {n_embd, rank}, 0);
const int64_t n_conv_groups = n_embd / group_size;
const int64_t n_conv_proj = 2 * kernel_size * n_conv_groups;
for (int i = 0; i < n_layer; ++i) {
ggml_context * ctx_split = use_split_ctx ? ctx_for_layer_split(i) : ctx_for_layer(i);
auto & layer = model.layers[i];
layer.attn_norm = create_tensor(ctx_split, tn(LLM_TENSOR_ATTN_NORM, "weight", i), {n_embd}, 0);
layer.wq = create_tensor(ctx_split, tn(LLM_TENSOR_ATTN_Q, "weight", i), {n_embd, n_embd_head_k * n_head}, 0);
layer.wk = create_tensor(ctx_split, tn(LLM_TENSOR_ATTN_K, "weight", i), {n_embd, n_embd_k_gqa}, 0);
layer.wv = create_tensor(ctx_split, tn(LLM_TENSOR_ATTN_V, "weight", i), {n_embd, n_embd_v_gqa}, 0);
layer.wo = create_tensor(ctx_split, tn(LLM_TENSOR_ATTN_OUT, "weight", i), {n_embd_head_v * n_head, n_embd}, 0);
layer.attn_q_norm = create_tensor(ctx_split, tn(LLM_TENSOR_ATTN_Q_NORM, "weight", i), {n_embd_head_k}, 0);
layer.attn_k_norm = create_tensor(ctx_split, tn(LLM_TENSOR_ATTN_K_NORM, "weight", i), {n_embd_head_k}, 0);
layer.ffn_norm = create_tensor(ctx_split, tn(LLM_TENSOR_FFN_NORM, "weight", i), {n_embd}, 0);
layer.ffn_gate = create_tensor(ctx_split, tn(LLM_TENSOR_FFN_GATE, "weight", i), {n_embd, n_ff}, 0);
layer.ffn_down = create_tensor(ctx_split, tn(LLM_TENSOR_FFN_DOWN, "weight", i), {n_ff, n_embd}, 0);
layer.ffn_up = create_tensor(ctx_split, tn(LLM_TENSOR_FFN_UP, "weight", i), {n_embd, n_ff}, 0);
layer.dflash_attn_conv_base = create_tensor(ctx_split,
tn(LLM_TENSOR_DFLASH_ATTN_CONV_BASE, i), {n_embd, kernel_size, 2}, 0);
layer.dflash_attn_conv_proj = create_tensor(ctx_split,
tn(LLM_TENSOR_DFLASH_ATTN_CONV_PROJ, "weight", i), {n_embd, n_conv_proj}, 0);
layer.dflash_ffn_conv_base = create_tensor(ctx_split,
tn(LLM_TENSOR_DFLASH_FFN_CONV_BASE, i), {n_embd, kernel_size, 2}, 0);
layer.dflash_ffn_conv_proj = create_tensor(ctx_split,
tn(LLM_TENSOR_DFLASH_FFN_CONV_PROJ, "weight", i), {n_embd, n_conv_proj}, 0);
}
return use_mmap_buffer;
}
bool create_tensors_helper::create_dflash_dsv4_tensors(const LLM_TN & tn) {
LOADING_PRELUDE
@ -5241,6 +5302,8 @@ bool create_tensors_helper::create_tensors() {
use_mmap_buffer = create_gemma4_mtp_tensors(tn); break;
case LLM_ARCH_DFLASH:
use_mmap_buffer = create_dflash_dsv4_tensors(tn); break;
case LLM_ARCH_DFLASH2:
use_mmap_buffer = create_dflash2_tensors(tn); break;
case LLM_ARCH_DFLASH_DRAFT:
use_mmap_buffer = create_dflash_tensors(tn); break;
case LLM_ARCH_STARCODER2:

View File

@ -88,6 +88,8 @@ struct llama_model_loader {
std::string arch_name;
LLM_KV llm_kv = LLM_KV(LLM_ARCH_UNKNOWN);
mutable bool arch_resolved = false;
mutable llm_arch resolved_arch = LLM_ARCH_UNKNOWN;
llama_expert_tensor_index expert_tensor_index;
llama_model_loader(const std::string & fname, int ncmoe, bool use_mmap, bool check_tensors, bool repack_tensors, bool use_thp,
@ -131,7 +133,16 @@ struct llama_model_loader {
const std::string& get_arch_name() const { return arch_name; }
enum llm_arch get_arch() const { return llm_kv.arch; }
enum llm_arch get_arch() const {
if (!arch_resolved) {
resolved_arch = llm_kv.arch;
if (resolved_arch == LLM_ARCH_DFLASH && get_tensor_meta("selector_hidden.weight") != nullptr) {
resolved_arch = LLM_ARCH_DFLASH2;
}
arch_resolved = true;
}
return resolved_arch;
}
const char * get_tensor_name(int i) const;

View File

@ -937,6 +937,34 @@ static const std::map<llm_arch, std::map<llm_tensor, std::string>> LLM_TENSOR_NA
{ LLM_TENSOR_DSPARK_CONF_PROJ, "conf_proj" },
},
},
{
LLM_ARCH_DFLASH2,
{
{ LLM_TENSOR_TOKEN_EMBD, "token_embd" },
{ LLM_TENSOR_OUTPUT_NORM, "output_norm" },
{ LLM_TENSOR_OUTPUT, "output" },
{ LLM_TENSOR_ATTN_NORM, "blk.%d.attn_norm" },
{ LLM_TENSOR_ATTN_Q, "blk.%d.attn_q" },
{ LLM_TENSOR_ATTN_Q_NORM, "blk.%d.attn_q_norm" },
{ LLM_TENSOR_ATTN_K, "blk.%d.attn_k" },
{ LLM_TENSOR_ATTN_K_NORM, "blk.%d.attn_k_norm" },
{ LLM_TENSOR_ATTN_V, "blk.%d.attn_v" },
{ LLM_TENSOR_ATTN_OUT, "blk.%d.attn_output" },
{ LLM_TENSOR_FFN_NORM, "blk.%d.ffn_norm" },
{ LLM_TENSOR_FFN_GATE, "blk.%d.ffn_gate" },
{ LLM_TENSOR_FFN_DOWN, "blk.%d.ffn_down" },
{ LLM_TENSOR_FFN_UP, "blk.%d.ffn_up" },
{ LLM_TENSOR_DFLASH_FC, "fc" },
{ LLM_TENSOR_DFLASH_HIDDEN_NORM, "enc.output_norm" },
{ LLM_TENSOR_DFLASH_ATTN_CONV_BASE, "blk.%d.attn_conv_base" },
{ LLM_TENSOR_DFLASH_ATTN_CONV_PROJ, "blk.%d.attn_conv_proj" },
{ LLM_TENSOR_DFLASH_FFN_CONV_BASE, "blk.%d.ffn_conv_base" },
{ LLM_TENSOR_DFLASH_FFN_CONV_PROJ, "blk.%d.ffn_conv_proj" },
{ LLM_TENSOR_DFLASH_SELECTOR_PREV, "selector_predecessor" },
{ LLM_TENSOR_DFLASH_SELECTOR_NEXT, "selector_successor" },
{ LLM_TENSOR_DFLASH_SELECTOR_HIDDEN, "selector_hidden" },
},
},
{
LLM_ARCH_GEMMA4_ASSISTANT,
{

View File

@ -198,6 +198,12 @@ struct llama_layer {
struct ggml_tensor * wo_enc = nullptr;
struct ggml_tensor * attn_sinks = nullptr;
// DFlash2 dynamic grouped-convolution parameters.
struct ggml_tensor * dflash_attn_conv_base = nullptr;
struct ggml_tensor * dflash_attn_conv_proj = nullptr;
struct ggml_tensor * dflash_ffn_conv_base = nullptr;
struct ggml_tensor * dflash_ffn_conv_proj = nullptr;
// attention bias
struct ggml_tensor * bq = nullptr;
struct ggml_tensor * bk = nullptr;
@ -483,6 +489,9 @@ struct llama_model {
struct ggml_tensor * mtp_centroids = nullptr;
struct ggml_tensor * dflash_fc = nullptr;
struct ggml_tensor * dflash_hidden_norm = nullptr;
struct ggml_tensor * dflash_selector_prev = nullptr;
struct ggml_tensor * dflash_selector_next = nullptr;
struct ggml_tensor * dflash_selector_hidden = nullptr;
std::vector<struct ggml_tensor *> dflash_aux_hidden_norms;
struct ggml_tensor * dspark_markov_w1 = nullptr;
struct ggml_tensor * dspark_markov_w2 = nullptr;

View File

@ -128,11 +128,20 @@ bool llama_model_dflash_has_dspark_head(const struct llama_model * model) {
}
static const ggml_tensor * llama_dflash_output_tensor(
const struct llama_model * model) {
const struct llama_model * model,
bool dflash2) {
if (model == nullptr) {
return nullptr;
}
if (dflash2) {
if (model->output_mtp_ptr != nullptr &&
model->output_mtp == model->output_mtp_ptr.get()) {
return model->output_mtp;
}
return model->output != nullptr ? model->output : model->tok_embd;
}
if (model->output_mtp != nullptr) {
return model->output_mtp;
}
@ -151,8 +160,9 @@ int32_t llama_model_dflash_io_mode(
return LLAMA_DFLASH_IO_MODE_INVALID;
}
const ggml_tensor * draft_output = llama_dflash_output_tensor(draft_model);
const ggml_tensor * target_output = llama_dflash_output_tensor(target_model);
const bool dflash2 = draft_model->arch == LLM_ARCH_DFLASH2;
const ggml_tensor * draft_output = llama_dflash_output_tensor(draft_model, dflash2);
const ggml_tensor * target_output = llama_dflash_output_tensor(target_model, dflash2);
if (draft_model->tok_embd == nullptr || draft_output == nullptr || target_model->tok_embd == nullptr || target_output == nullptr) {
return LLAMA_DFLASH_IO_MODE_INVALID;
}
@ -213,7 +223,8 @@ bool llama_model_dflash_io_tensors_match(
const struct llama_model * draft_model,
int32_t n_embd,
int32_t n_vocab) {
const ggml_tensor * output = llama_dflash_output_tensor(draft_model);
const ggml_tensor * output = llama_dflash_output_tensor(
draft_model, draft_model != nullptr && draft_model->arch == LLM_ARCH_DFLASH2);
if (draft_model == nullptr || draft_model->tok_embd == nullptr || output == nullptr || n_embd <= 0 || n_vocab <= 0) {
return false;
}
@ -235,23 +246,39 @@ bool llama_model_share_dflash_io_tensors(
return true;
}
const bool dflash2 = draft_model->arch == LLM_ARCH_DFLASH2;
const ggml_tensor * target_output_const = llama_dflash_output_tensor(target_model, dflash2);
ggml_tensor * target_output = const_cast<ggml_tensor *>(target_output_const);
if (dflash2 && target_output != nullptr) {
const bool uses_requantized_primary =
target_model->output_mtp_ptr != nullptr &&
target_model->output_mtp == target_model->output_mtp_ptr.get() &&
target_output == target_model->output_mtp;
LLAMA_LOG_INFO("%s: DFlash2 target output = %s (%s)\n",
__func__,
uses_requantized_primary ? "requantized primary" : "primary",
ggml_type_name(target_output->type));
}
if (draft_model->tok_embd == nullptr) {
draft_model->tok_embd = target_model->tok_embd;
}
if (draft_model->output == nullptr) {
draft_model->output = target_model->output ? target_model->output : target_model->tok_embd;
draft_model->output = target_output ? target_output : target_model->tok_embd;
if (draft_model->output == nullptr) {
draft_model->output = draft_model->tok_embd;
}
}
const bool uses_shared_tok = draft_model->tok_embd == target_model->tok_embd;
const bool uses_shared_output = draft_model->output == target_model->output ||
const bool uses_shared_output = draft_model->output == target_output ||
draft_model->output == target_model->tok_embd;
if (draft_model->output_mtp == nullptr) {
if (target_model->output_mtp != nullptr && uses_shared_tok && uses_shared_output) {
if (draft_model->arch != LLM_ARCH_DFLASH2 &&
target_model->output_mtp != nullptr && uses_shared_tok && uses_shared_output) {
draft_model->output_mtp = target_model->output_mtp;
} else if (draft_model->output != nullptr) {
draft_model->output_mtp = draft_model->output;
@ -262,9 +289,10 @@ bool llama_model_share_dflash_io_tensors(
const bool output_mtp_aliases_output = draft_model->output_mtp == draft_model->output;
const bool tok_embd_is_shared = draft_model->tok_embd == target_model->tok_embd;
const bool output_is_shared = draft_model->output == target_model->output ||
const bool output_is_shared = draft_model->output == target_output ||
draft_model->output == target_model->tok_embd;
const bool output_mtp_is_shared = draft_model->output_mtp == target_model->output_mtp ||
const bool output_mtp_is_shared = draft_model->output_mtp == target_output ||
(!dflash2 && draft_model->output_mtp == target_model->output_mtp) ||
draft_model->output_mtp == target_model->output ||
draft_model->output_mtp == target_model->tok_embd;
const bool isolate_shared_io =
@ -306,7 +334,7 @@ bool llama_model_share_dflash_io_tensors(
}
}
const struct ggml_tensor * output = llama_dflash_output_tensor(draft_model);
const struct ggml_tensor * output = llama_dflash_output_tensor(draft_model, dflash2);
return draft_model->tok_embd != nullptr && output != nullptr;
}
@ -493,6 +521,14 @@ static int llama_dflash_capture_eval_callback(struct ggml_tensor * tensor, bool
return 0;
}
const ggml_type capture_type = tensor->type;
const ggml_type_traits_t capture_traits = ggml_internal_get_type_traits(capture_type);
if (capture_type != GGML_TYPE_F32 && capture_traits.to_float == nullptr) {
ctx->dflash.capture->invalid = true;
LLAMA_LOG_WARN("%s: unsupported DFlash capture type %s\n", __func__, ggml_type_name(capture_type));
return 2;
}
auto & capture = *ctx->dflash.capture;
if (capture.capture_batch_id == 0) {
capture.capture_batch_id = 1;
@ -514,6 +550,7 @@ static int llama_dflash_capture_eval_callback(struct ggml_tensor * tensor, bool
}
auto & rows = capture.layer_rows[(size_t) layer_idx];
auto & chunks = capture.layer_chunks[(size_t) layer_idx];
auto & rows_written = capture.layer_rows_written[(size_t) layer_idx];
if (rows_written + row_count > capture.expected_rows) {
capture.invalid = true;
@ -527,9 +564,34 @@ static int llama_dflash_capture_eval_callback(struct ggml_tensor * tensor, bool
}
auto backend = ggml_backend_sched_get_tensor_backend(ctx->sched, tensor);
GGML_ASSERT(backend);
ggml_backend_tensor_get_async(backend, tensor,
rows.data() + (size_t) rows_written * (size_t) row_width,
0, (size_t) row_count * (size_t) row_width * sizeof(float));
const size_t raw_row_stride = (size_t) row_width * sizeof(float);
const size_t byte_offset = (size_t) rows_written * raw_row_stride;
void * readback_dst = nullptr;
size_t readback_bytes = 0;
if (capture_type == GGML_TYPE_F32) {
readback_bytes = (size_t) row_count * (size_t) row_width * sizeof(float);
} else {
const size_t row_bytes = ggml_row_size(capture_type, row_width);
readback_bytes = (size_t) row_count * row_bytes;
}
const size_t rows_bytes = rows.size() * sizeof(float);
if (byte_offset > rows_bytes || readback_bytes > rows_bytes - byte_offset) {
capture.invalid = true;
LLAMA_LOG_WARN("%s: DFlash capture readback exceeds row storage for layer %d: offset=%zu size=%zu capacity=%zu\n",
__func__, layer_id, byte_offset, readback_bytes, rows_bytes);
return 2;
}
readback_dst = capture_type == GGML_TYPE_F32
? static_cast<void *>(rows.data() + (size_t) rows_written * (size_t) row_width)
: static_cast<void *>(reinterpret_cast<uint8_t *>(rows.data()) + byte_offset);
ggml_backend_tensor_get_async(backend, tensor, readback_dst, 0, readback_bytes);
chunks.push_back({ rows_written, row_count, byte_offset, capture_type });
rows_written += row_count;
capture.row_width = row_width;
capture.row_count = std::max(capture.row_count, rows_written);
@ -553,6 +615,7 @@ bool llama_set_dflash_capture_layers(
auto capture = std::make_unique<llama_context::dflash_runtime::capture_state>();
capture->layer_ids.assign(layer_ids, layer_ids + n_layers);
capture->layer_rows.resize((size_t) n_layers);
capture->layer_chunks.resize((size_t) n_layers);
capture->layer_rows_written.assign((size_t) n_layers, 0);
capture->layer_seen_batch_id.assign((size_t) n_layers, 0);
capture->prev_cb_eval = ctx->cparams.cb_eval;
@ -613,6 +676,9 @@ void llama_begin_dflash_capture_batch(struct llama_context * ctx, int32_t expect
capture.invalid = expected_rows <= 0;
std::fill(capture.layer_rows_written.begin(), capture.layer_rows_written.end(), 0);
std::fill(capture.layer_seen_batch_id.begin(), capture.layer_seen_batch_id.end(), 0);
for (auto & chunks : capture.layer_chunks) {
chunks.clear();
}
}
void llama_finish_dflash_capture_batch(
@ -648,6 +714,7 @@ static bool llama_spec_prepare_dflash_capture(
n_layers = (int32_t) capture.layer_ids.size();
if (capture.invalid || row_count <= 0 || row_width <= 0 || n_layers <= 0 ||
capture.expected_rows <= 0 || capture.layer_rows.size() != (size_t) n_layers ||
capture.layer_chunks.size() != (size_t) n_layers ||
capture.layer_rows_written.size() != (size_t) n_layers) {
return false;
}
@ -667,6 +734,46 @@ static bool llama_spec_prepare_dflash_capture(
}
for (int32_t layer_idx = 0; layer_idx < n_layers; ++layer_idx) {
auto & rows = capture.layer_rows[(size_t) layer_idx];
if (rows.size() != (size_t) row_count * (size_t) row_width) {
return false;
}
const auto & chunks = capture.layer_chunks[(size_t) layer_idx];
const size_t rows_bytes = rows.size() * sizeof(float);
std::vector<uint8_t> chunk_buffer;
for (const auto & chunk : chunks) {
if (chunk.row_offset < 0 || chunk.row_count <= 0 ||
chunk.row_offset + chunk.row_count > row_count) {
return false;
}
if (chunk.type == GGML_TYPE_F32) {
continue;
}
const ggml_type_traits_t traits = ggml_internal_get_type_traits(chunk.type);
if (traits.to_float == nullptr) {
return false;
}
const size_t row_bytes = ggml_row_size(chunk.type, row_width);
const size_t chunk_bytes = (size_t) chunk.row_count * row_bytes;
if (chunk.byte_offset > rows_bytes ||
chunk_bytes > rows_bytes - chunk.byte_offset) {
return false;
}
chunk_buffer.resize(chunk_bytes);
std::memcpy(chunk_buffer.data(),
reinterpret_cast<const uint8_t *>(rows.data()) + chunk.byte_offset,
chunk_bytes);
traits.to_float(
chunk_buffer.data(),
rows.data() + (size_t) chunk.row_offset * (size_t) row_width,
(int64_t) chunk.row_count * row_width);
}
if (capture.layer_seen_batch_id[(size_t) layer_idx] != capture.capture_batch_id) {
LLAMA_LOG_WARN("%s: DFlash capture is stale for layer %d (seen_batch=%llu current_batch=%llu rows=%d width=%d)\n",
__func__,
@ -678,7 +785,6 @@ static bool llama_spec_prepare_dflash_capture(
return false;
}
const auto & rows = capture.layer_rows[(size_t) layer_idx];
if (capture.layer_rows_written[(size_t) layer_idx] != row_count ||
rows.size() != (size_t) row_count * (size_t) row_width) {
LLAMA_LOG_WARN("%s: DFlash capture rows mismatch for layer %d: got=%d/%zu expected=%d/%zu (rows=%d width=%d)\n",

View File

@ -4010,6 +4010,10 @@ static std::pair<std::vector<double>, double> get_layer_sizes(const llama_model_
continue;
}
if (name == "dflash_fc.weight" || name == "dflash_hidden_norm.weight" ||
(model.arch == LLM_ARCH_DFLASH2 &&
(name == "fc.weight" || name == "output_norm.weight" ||
name == "selector_predecessor.weight" || name == "selector_successor.weight" ||
name == "selector_hidden.weight" || name == "enc.output_norm.weight")) ||
(model.arch == LLM_ARCH_DFLASH &&
(name == "fc.weight" || name == "enc.output_norm.weight")) ||
name.rfind("dflash_aux_hidden_norm.", 0) == 0 ||
@ -5242,7 +5246,8 @@ static void llama_set_inputs(llama_context & lctx, const llama_batch & batch) {
auto tim1 = ggml_time_us();
#endif
const int64_t n_tokens = batch.n_tokens;
if (n_tokens > 1 && !cparams.mtp && lctx.n_outputs < n_tokens) {
if (n_tokens > 1 && !cparams.mtp && lctx.n_outputs < n_tokens &&
!llm_arch_requires_all_graph_output_rows(lctx.model.arch)) {
GGML_ASSERT(lctx.inp_out_ids && "every model that can must skip unused outputs");
}
@ -6589,6 +6594,8 @@ static int llama_decode_internal(
//}
lctx.dflash.draft_tokens.clear();
lctx.dflash.draft_lattice.clear();
lctx.dflash.draft_lattice_ids.clear();
if (lctx.dflash.draft_tokens_tensor != nullptr) {
ggml_backend_t backend_argmax = ggml_backend_sched_get_tensor_backend(
lctx.sched, lctx.dflash.draft_tokens_tensor);
@ -6602,6 +6609,53 @@ static int llama_decode_internal(
}
}
if (lctx.model.arch == LLM_ARCH_DFLASH2 && lctx.dflash.draft_lattice_tensor != nullptr &&
lctx.dflash.draft_lattice_ids_tensor != nullptr) {
ggml_backend_t backend_lattice = ggml_backend_sched_get_tensor_backend(
lctx.sched, lctx.dflash.draft_lattice_tensor);
ggml_backend_t backend_lattice_ids = ggml_backend_sched_get_tensor_backend(
lctx.sched, lctx.dflash.draft_lattice_ids_tensor);
if (backend_lattice != nullptr && backend_lattice_ids != nullptr) {
const size_t n_values = (size_t) lctx.dflash.draft_lattice_tensor->ne[0] *
(size_t) lctx.dflash.draft_lattice_tensor->ne[1];
lctx.dflash.draft_lattice.resize(n_values);
const int32_t top_k = lctx.dflash.draft_lattice_top_k;
const int32_t width = top_k * top_k;
const int32_t n_positions = (int32_t) lctx.dflash.draft_lattice_tensor->ne[1];
const size_t n_ids = (size_t) top_k * (size_t) n_positions;
lctx.dflash.draft_lattice_ids.resize(n_ids);
ggml_backend_tensor_get_async(backend_lattice,
lctx.dflash.draft_lattice_tensor,
lctx.dflash.draft_lattice.data(), 0,
n_values * sizeof(float));
ggml_backend_tensor_get_async(backend_lattice_ids,
lctx.dflash.draft_lattice_ids_tensor,
lctx.dflash.draft_lattice_ids.data(), 0,
n_ids * sizeof(int32_t));
llama_synchronize(&lctx);
std::vector<llama_token> selected;
selected.reserve(std::max(0, n_positions - 1));
int32_t previous = 0;
for (int32_t pos = 1; pos < n_positions; ++pos) {
const float * row = lctx.dflash.draft_lattice.data() + (size_t) pos * width;
int32_t best = 0;
float best_score = -INFINITY;
for (int32_t current = 0; current < top_k; ++current) {
const float score = row[current + top_k * previous];
if (score > best_score) {
best_score = score;
best = current;
}
}
selected.push_back((llama_token) lctx.dflash.draft_lattice_ids[(size_t) pos * top_k + best]);
previous = best;
}
if (!selected.empty()) {
lctx.dflash.draft_tokens = std::move(selected);
}
}
}
// extract logits
{
const bool dflash_skip_logits = (llm_arch_is_dflash_family(lctx.model.arch)
@ -6760,10 +6814,14 @@ static int llama_decode_internal(
#if IK_PRINT_TIMING
auto tim1 = ggml_time_us();
#endif
if (lctx.cparams.mtp_op_type == MTP_OP_NONE && !lctx.prev) {
// Keep scheduler alive in case someone dont run graph reuse so
// speculative checkpoint restoration can be completed
const bool speculative_checkpoint_active =
lctx.kv_self.ckpt.selected_spec_mode != LLAMA_SPEC_CKPT_NONE;
if (lctx.cparams.mtp_op_type == MTP_OP_NONE && !lctx.prev && !speculative_checkpoint_active) {
ggml_backend_sched_reset(lctx.sched);
}
else if (lctx.cparams.mtp_op_type != MTP_OP_NONE && !lctx.prev_mtp) {
else if (lctx.cparams.mtp_op_type != MTP_OP_NONE && !lctx.prev_mtp && !speculative_checkpoint_active) {
ggml_backend_sched_reset(lctx.sched);
}
#if IK_PRINT_TIMING
@ -8870,6 +8928,7 @@ enum llama_rope_type llama_rope_type(const struct llama_model * model) {
case LLM_ARCH_GEMMA4:
case LLM_ARCH_GEMMA4_MTP:
case LLM_ARCH_DFLASH_DRAFT:
case LLM_ARCH_DFLASH2:
case LLM_ARCH_GEMMA4_ASSISTANT:
return LLAMA_ROPE_TYPE_NEOX;
@ -11737,6 +11796,49 @@ llama_token llama_get_dflash_draft_token_ith(struct llama_context * ctx, int32_t
return ctx->dflash.draft_tokens[(size_t) i];
}
int32_t llama_get_dflash_draft_lattice_top_k(struct llama_context * ctx) {
if (ctx == nullptr) {
return 0;
}
llama_synchronize(ctx);
return ctx->dflash.draft_lattice_top_k;
}
int32_t llama_get_dflash_draft_lattice_n_positions(struct llama_context * ctx) {
if (ctx == nullptr || ctx->dflash.draft_lattice_top_k <= 0) {
return 0;
}
llama_synchronize(ctx);
const size_t top_k = (size_t) ctx->dflash.draft_lattice_top_k;
return (int32_t) (ctx->dflash.draft_lattice_ids.size() / top_k);
}
bool llama_copy_dflash_draft_lattice(
struct llama_context * ctx,
float * scores, size_t score_count,
int32_t * ids, size_t id_count) {
if (ctx == nullptr || scores == nullptr || ids == nullptr || ctx->dflash.draft_lattice_top_k <= 0) {
return false;
}
llama_synchronize(ctx);
const size_t top_k = (size_t) ctx->dflash.draft_lattice_top_k;
const size_t n_positions = ctx->dflash.draft_lattice_ids.size() / top_k;
const size_t n_scores = top_k * top_k * n_positions;
const size_t n_ids = top_k * n_positions;
if (ctx->dflash.draft_lattice.size() != n_scores || score_count < n_scores || id_count < n_ids) {
return false;
}
const int32_t n_vocab = (int32_t) ctx->model.vocab.n_tokens();
if (std::any_of(ctx->dflash.draft_lattice_ids.begin(),
ctx->dflash.draft_lattice_ids.begin() + n_ids,
[n_vocab](int32_t token) { return token < 0 || token >= n_vocab; })) {
return false;
}
std::memcpy(scores, ctx->dflash.draft_lattice.data(), n_scores * sizeof(float));
std::memcpy(ids, ctx->dflash.draft_lattice_ids.data(), n_ids * sizeof(int32_t));
return true;
}
float * llama_get_embeddings(struct llama_context * ctx) {
llama_synchronize(ctx);