diff --git a/common/common.h b/common/common.h index 3bb51367..ebadab7c 100644 --- a/common/common.h +++ b/common/common.h @@ -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) diff --git a/common/sampling.cpp b/common/sampling.cpp index 3899167e..ab98fe62 100644 --- a/common/sampling.cpp +++ b/common/sampling.cpp @@ -2,10 +2,15 @@ #include "sampling.h" #include "llama-vocab.h" #include "common.h" +#include "speculative.h" #include "reasoning-budget.cpp" #include #include +#include + +// 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 #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 common_sampler_sample_and_accept_n(struct common_sample return result; } +std::vector common_sampler_sample_and_accept_n( + struct common_sampler * gsmpl, + struct llama_context * ctx, + const std::vector & idxs, + const std::vector & draft, + const std::vector & dists, + bool grammar_first) { + GGML_ASSERT(idxs.size() == draft.size() + 1); + GGML_ASSERT(dists.size() == draft.size()); + + std::vector result; + result.reserve(idxs.size()); + + std::uniform_real_distribution 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 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 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 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 diff --git a/common/sampling.h b/common/sampling.h index c36aff4f..32081b11 100644 --- a/common/sampling.h +++ b/common/sampling.h @@ -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* server_biases; @@ -357,6 +361,14 @@ std::vector llama_sampling_sample_and_accept_n(struct common_sample std::vector common_sampler_sample_and_accept_n(struct common_sampler * gsmpl, struct llama_context * ctx, const std::vector & idxs, const std::vector & draft, bool grammar_first = false); +std::vector common_sampler_sample_and_accept_n( + struct common_sampler * gsmpl, + struct llama_context * ctx, + const std::vector & idxs, + const std::vector & draft, + const std::vector & 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); diff --git a/common/speculative-dflash-impl.h b/common/speculative-dflash-impl.h index e0de29d9..bb437ae9 100644 --- a/common/speculative-dflash-impl.h +++ b/common/speculative-dflash-impl.h @@ -1,8 +1,12 @@ #pragma once #include +#include #include #include +#include +#include +#include #include 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 proposal_dists; + std::mt19937 selector_rng; + uint32_t selector_seed = LLAMA_DEFAULT_SEED; + bool selector_rng_initialized = false; + std::vector target_layer_ids; std::vector target_window; std::vector 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 scores((size_t) selector_top_k * selector_top_k * n_positions_used); + std::vector 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 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) { diff --git a/common/speculative.cpp b/common/speculative.cpp index bf777391..71e566af 100644 --- a/common/speculative.cpp +++ b/common/speculative.cpp @@ -17,6 +17,7 @@ #include #include #include +#include #include #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(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 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; diff --git a/common/speculative.h b/common/speculative.h index 05ff28ae..ea1beb93 100644 --- a/common/speculative.h +++ b/common/speculative.h @@ -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 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 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); diff --git a/convert_hf_to_gguf.py b/convert_hf_to_gguf.py index 55f879e9..c6b560f3 100644 --- a/convert_hf_to_gguf.py +++ b/convert_hf_to_gguf.py @@ -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.""" diff --git a/examples/main/main.cpp b/examples/main/main.cpp index e4e70071..bbe58b3e 100644 --- a/examples/main/main.cpp +++ b/examples/main/main.cpp @@ -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::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 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 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()) { diff --git a/examples/server/server-context.cpp b/examples/server/server-context.cpp index 58855240..2b1c151b 100644 --- a/examples/server/server-context.cpp +++ b/examples/server/server-context.cpp @@ -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 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; diff --git a/examples/server/server-context.h b/examples/server/server-context.h index 551b7d13..02b996ab 100644 --- a/examples/server/server-context.h +++ b/examples/server/server-context.h @@ -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 draft_proposal_dists; bool spec_target_only = false; json json_schema; diff --git a/gguf-py/gguf/constants.py b/gguf-py/gguf/constants.py index e26e8f04..02b8323b 100644 --- a/gguf-py/gguf/constants.py +++ b/gguf-py/gguf/constants.py @@ -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, diff --git a/gguf-py/gguf/gguf_writer.py b/gguf-py/gguf/gguf_writer.py index 84a6a965..7e85ba75 100644 --- a/gguf-py/gguf/gguf_writer.py +++ b/gguf-py/gguf/gguf_writer.py @@ -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) diff --git a/gguf-py/gguf/tensor_mapping.py b/gguf-py/gguf/tensor_mapping.py index 2b80d249..1b15f2eb 100644 --- a/gguf-py/gguf/tensor_mapping.py +++ b/gguf-py/gguf/tensor_mapping.py @@ -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 ), diff --git a/include/llama.h b/include/llama.h index b1e40e73..82f81680 100644 --- a/include/llama.h +++ b/include/llama.h @@ -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 diff --git a/src/graphs/build_dflash.cpp b/src/graphs/build_dflash.cpp index 87030f65..94222f8a 100644 --- a/src/graphs/build_dflash.cpp +++ b/src/graphs/build_dflash.cpp @@ -4,6 +4,65 @@ #include +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(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(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); diff --git a/src/llama-arch.cpp b/src/llama-arch.cpp index dd00a254..4a33ccd5 100644 --- a/src/llama-arch.cpp +++ b/src/llama-arch.cpp @@ -87,6 +87,7 @@ static const std::map 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_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; } diff --git a/src/llama-arch.h b/src/llama-arch.h index 7b47e1f5..fe3cc0f7 100644 --- a/src/llama-arch.h +++ b/src/llama-arch.h @@ -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); diff --git a/src/llama-build-context.cpp b/src/llama-build-context.cpp index 18a4c905..669c92b4 100644 --- a/src/llama-build-context.cpp +++ b/src/llama-build-context.cpp @@ -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(); diff --git a/src/llama-context.h b/src/llama-context.h index 5b296f61..dc7dd136 100644 --- a/src/llama-context.h +++ b/src/llama-context.h @@ -448,6 +448,13 @@ struct llama_context { struct capture_state { std::vector layer_ids; std::vector> 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> layer_chunks; std::vector layer_rows_written; int32_t row_count = 0; int32_t row_width = 0; @@ -477,6 +484,11 @@ struct llama_context { std::vector 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 draft_lattice; + std::vector draft_lattice_ids; + int32_t draft_lattice_top_k = 0; }; dflash_runtime dflash; using dflash_capture_state = dflash_runtime::capture_state; diff --git a/src/llama-dflash.cpp b/src/llama-dflash.cpp index 3704b115..5ac52219 100644 --- a/src/llama-dflash.cpp +++ b/src/llama-dflash.cpp @@ -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(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; } } diff --git a/src/llama-hparams.cpp b/src/llama-hparams.cpp index 84eb3b58..4af1743b 100644 --- a/src/llama-hparams.cpp +++ b/src/llama-hparams.cpp @@ -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: diff --git a/src/llama-hparams.h b/src/llama-hparams.h index b68a7384..ac70d591 100644 --- a/src/llama-hparams.h +++ b/src/llama-hparams.h @@ -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; diff --git a/src/llama-load-tensors.cpp b/src/llama-load-tensors.cpp index c1e2f3b3..2ca2a212 100644 --- a/src/llama-load-tensors.cpp +++ b/src/llama-load-tensors.cpp @@ -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: diff --git a/src/llama-model-loader.h b/src/llama-model-loader.h index 30390e42..a07eefb8 100644 --- a/src/llama-model-loader.h +++ b/src/llama-model-loader.h @@ -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; diff --git a/src/llama-model.cpp b/src/llama-model.cpp index 378f026b..4e38a3b5 100644 --- a/src/llama-model.cpp +++ b/src/llama-model.cpp @@ -937,6 +937,34 @@ static const std::map> 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, { diff --git a/src/llama-model.h b/src/llama-model.h index f86fd30d..88446bc2 100644 --- a/src/llama-model.h +++ b/src/llama-model.h @@ -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 dflash_aux_hidden_norms; struct ggml_tensor * dspark_markov_w1 = nullptr; struct ggml_tensor * dspark_markov_w2 = nullptr; diff --git a/src/llama-spec-features-dflash.cpp b/src/llama-spec-features-dflash.cpp index deed94a2..451bc769 100644 --- a/src/llama-spec-features-dflash.cpp +++ b/src/llama-spec-features-dflash.cpp @@ -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(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(rows.data() + (size_t) rows_written * (size_t) row_width) + : static_cast(reinterpret_cast(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(); 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 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(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", diff --git a/src/llama.cpp b/src/llama.cpp index 46a17d74..d670c5f8 100644 --- a/src/llama.cpp +++ b/src/llama.cpp @@ -4010,6 +4010,10 @@ static std::pair, 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 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);