diff --git a/BUILD.bazel b/BUILD.bazel index 47730efc..5b2a47db 100644 --- a/BUILD.bazel +++ b/BUILD.bazel @@ -690,6 +690,36 @@ cc_library( ], ) +cc_test( + name = "deepseek_test", + size = "small", + timeout = "long", + srcs = ["deepseek/deepseek_test.cc"], + linkstatic = True, + local_defines = ["HWY_IS_TEST"], + deps = [ + ":activations", + ":allocator", + ":basics", + ":configs", + ":gemma_args", + ":gemma_lib", + ":kv_cache", + ":mat", + ":matmul_env", + ":matmul_static", + ":ops", + ":test_util", + ":threading_context", + ":weights", + "@googletest//:gtest_main", # buildcleaner: keep + "//compression:types", + "@highway//:hwy", + "@highway//:hwy_test_util", + "@highway//:nanobenchmark", # buildcleaner: keep + ], +) + cc_library( name = "activations", hdrs = ["gemma/activations.h"], @@ -854,6 +884,7 @@ cc_library( "//io", "//io:blob_store", "//paligemma:image", + "@highway//:algo", "@highway//:bit_set", "@highway//:hwy", "@highway//:nanobenchmark", # timer @@ -1134,3 +1165,60 @@ cc_binary( "@nlohmann_json//:json", ], ) + +cc_binary( + name = "convert_dsv4", + srcs = ["deepseek/convert_dsv4.cc"], + deps = [ + ":args", + ":basics", + ":configs", + ":gemma_lib", + ":mat", + ":model_store", + ":tensor_info", + ":threading_context", + ":tokenizer", + ":weights", + "//compression:compress", + "//compression:types", + "//io", + "//io:blob_store", + "@highway//:abort_header_only", + "@highway//:hwy", + "@nlohmann_json//:json", + ], +) + +cc_library( + name = "dsv4_tokenizer", + srcs = ["deepseek/dsv4_tokenizer.cc"], + hdrs = [ + "deepseek/dsv4_tokenizer.h", + "deepseek/dsv4_unicode_ranges.inc", + ], + deps = [ + "@highway//:hwy", + "@nlohmann_json//:json", + ], +) + +cc_binary( + name = "run_dsv4", + srcs = ["deepseek/run_dsv4.cc"], + deps = [ + ":args", + ":basics", + ":configs", + ":dsv4_tokenizer", + ":gemma_args", + ":gemma_lib", + ":kv_cache", + ":mat", + ":matmul_env", + ":ops", + ":threading_context", + "@highway//:hwy", + "@highway//:profiler", + ], +) diff --git a/deepseek/convert_dsv4.cc b/deepseek/convert_dsv4.cc index 12c54b70..cfeeac3e 100644 --- a/deepseek/convert_dsv4.cc +++ b/deepseek/convert_dsv4.cc @@ -305,6 +305,7 @@ struct ConvertArgs { bool verify_only = false; // With verify_only: check only the MTP tensors (fast pre-flight). bool mtp_only = false; + bool dspark = false; }; class Converter { @@ -463,7 +464,7 @@ class Converter { expert = -1; std::string base = name; // Up to two numeric suffixes. - for (int pass = 0; pass < 2; ++pass) { + for (size_t pass = 0; pass < 2; ++pass) { const size_t us = base.rfind('_'); if (us == std::string::npos || us + 1 >= base.size()) break; bool numeric = true; @@ -487,8 +488,14 @@ class Converter { } void Run() { - const ModelConfig config(Model::DEEPSEEK4_FLASH, Type::kSFP, - PromptWrapping::GEMMA_IT); + const bool is_dspark = + args_.dspark || + (checkpoint_.Find("mtp.0.main_proj.weight") != nullptr); + ModelConfig config(Model::DEEPSEEK4_FLASH, Type::kSFP, + PromptWrapping::GEMMA_IT); + if (is_dspark) { + config.num_mtp_layers = 3; + } WeightsPtrs weights(config); const bool has_mtp = config.num_mtp_layers > 0; const LayerConfig mtp_lc = @@ -508,22 +515,27 @@ class Converter { } const size_t rows = mat.Rows(), cols = mat.Cols(); const size_t num = mat.Extents().Area(); - // The MTP block is registered as layer index `num_layers`; its source - // tensors live under "mtp.0." instead of "layers.N.". - const bool is_mtp = layer >= static_cast(config.num_layers); + // The MTP block is registered as layer indices + // `num_layers`..`num_layers + num_mtp_layers - 1`; its source tensors + // live under "mtp.0.".."mtp.k.". + const bool is_mtp = + layer >= 0 && static_cast(layer) >= config.num_layers; if (args_.mtp_only && !is_mtp && base.rfind("mtp_", 0) != 0) return; const LayerConfig* lc = layer < 0 ? nullptr : (is_mtp ? &mtp_lc : &config.layer_configs[layer]); + const size_t mtp_idx = + is_mtp ? (static_cast(layer) - config.num_layers) : 0; const std::string P = layer < 0 ? "" - : (is_mtp ? "mtp.0." : "layers." + std::to_string(layer) + "."); + : (is_mtp ? "mtp." + std::to_string(mtp_idx) + "." + : "layers." + std::to_string(layer) + "."); std::string src; // Transform: 0 = none, 1 = tail rows per segment, 2 = tail cols per // segment, 3 = flat tail elems per segment. - int transform = 0; + size_t transform = 0; size_t seg = 0, tail = 0; if (base == "c_embedding") { @@ -653,20 +665,76 @@ class Converter { } else if (base == "mtp_hnorm") { src = "mtp.0.hnorm.weight"; } else if (base == "mtp_norm") { - src = "mtp.0.norm.weight"; + const std::string last_mtp = + "mtp." + std::to_string(config.num_mtp_layers > 0 + ? config.num_mtp_layers - 1 + : 0) + + "."; + src = last_mtp + "norm.weight"; } else if (base == "mtp_hc_fn") { - src = "mtp.0.hc_head_fn"; + const std::string last_mtp = + "mtp." + std::to_string(config.num_mtp_layers > 0 + ? config.num_mtp_layers - 1 + : 0) + + "."; + src = last_mtp + "hc_head_fn"; } else if (base == "mtp_hc_base") { - src = "mtp.0.hc_head_base"; + const std::string last_mtp = + "mtp." + std::to_string(config.num_mtp_layers > 0 + ? config.num_mtp_layers - 1 + : 0) + + "."; + src = last_mtp + "hc_head_base"; } else if (base == "mtp_hc_scale") { - src = "mtp.0.hc_head_scale"; + const std::string last_mtp = + "mtp." + std::to_string(config.num_mtp_layers > 0 + ? config.num_mtp_layers - 1 + : 0) + + "."; + src = last_mtp + "hc_head_scale"; + } else if (base == "mtp_main_proj") { + src = "mtp.0.main_proj.weight"; + } else if (base == "mtp_main_norm") { + src = "mtp.0.main_norm.weight"; + } else if (base == "mtp_markov_w1") { + const std::string last_mtp = + "mtp." + std::to_string(config.num_mtp_layers > 0 + ? config.num_mtp_layers - 1 + : 0) + + "."; + src = checkpoint_.Find(last_mtp + "markov_head.markov_w1.weight") != + nullptr + ? last_mtp + "markov_head.markov_w1.weight" + : "mtp.0.markov_head.markov_w1.weight"; + } else if (base == "mtp_markov_w2") { + const std::string last_mtp = + "mtp." + std::to_string(config.num_mtp_layers > 0 + ? config.num_mtp_layers - 1 + : 0) + + "."; + src = checkpoint_.Find(last_mtp + "markov_head.markov_w2.weight") != + nullptr + ? last_mtp + "markov_head.markov_w2.weight" + : "mtp.0.markov_head.markov_w2.weight"; + } else if (base == "mtp_conf_proj") { + const std::string last_mtp = + "mtp." + std::to_string(config.num_mtp_layers > 0 + ? config.num_mtp_layers - 1 + : 0) + + "."; + src = checkpoint_.Find(last_mtp + "confidence_head.proj.weight") != + nullptr + ? last_mtp + "confidence_head.proj.weight" + : "mtp.0.confidence_head.proj.weight"; } else { HWY_ABORT("No mapping for tensor %s (base %s)", name.c_str(), base.c_str()); } - if (args_.verify_only && checkpoint_.Find(src) == nullptr) { - ++num_skipped_; // shard not downloaded yet + if (checkpoint_.Find(src) == nullptr) { + if (args_.verify_only) { + ++num_skipped_; // shard not downloaded yet + } return; } std::vector data = LoadF32(src, num); @@ -739,6 +807,8 @@ int Main(int argc, char** argv) { args.verify_only = true; } else if (a == "--mtp_only") { args.mtp_only = true; + } else if (a == "--dspark") { + args.dspark = true; } else { fprintf(stderr, "Unknown arg %s\n", a.c_str()); return 1; diff --git a/deepseek/deepseek.cc b/deepseek/deepseek.cc index eee52187..66ef6279 100644 --- a/deepseek/deepseek.cc +++ b/deepseek/deepseek.cc @@ -72,6 +72,7 @@ #include "gemma/attention.h" // includes highway.h #include "gemma/gemma-inl.h" #include "ops/ops-inl.h" +#include "hwy/contrib/algo/minmax-inl.h" HWY_BEFORE_NAMESPACE(); namespace gcpp { @@ -95,15 +96,36 @@ static constexpr uint32_t DSFloatToUint32Sortkey(float val) { } // Copies a small 1D/row tensor of any weight type to f32. -static void ReadRowF32(const MatPtr& w, size_t row, float* HWY_RESTRICT out, - size_t n) { - CallUpcastedActivation(&w, [&](const auto* t) { - for (size_t i = 0; i < n; ++i) { - out[i] = hwy::ConvertScalarTo(t->Row(row)[i]); +void ReadRowF32(const MatPtr& w, size_t row, float* HWY_RESTRICT out, + size_t n) { + if (w.GetType() == Type::kF32) { + MatPtrT wf(w); + const float* src = wf.Row(row); + hwy::CopyBytes(src, out, n * sizeof(float)); + if (wf.Scale() != 1.0f) { + MulByConst(wf.Scale(), out, n); } + return; + } + CallUpcasted(&w, [&](const auto* weights_t) { + const size_t ofs = row * weights_t->Stride(); + HWY_ASSERT(weights_t->Cols() == n); + const auto span = MakeSpan(weights_t->Row(0), ofs + n); + DecompressAndZeroPad(hn::ScalableTag(), span, ofs, out, n); + MulByConst(weights_t->Scale(), out, n); }); } +void ReadRowBF16(const MatPtr& w, size_t row, BF16* HWY_RESTRICT out, + size_t n) { + HWY_ALIGN float tmp[256]; + HWY_ASSERT(n <= 256); + ReadRowF32(w, row, tmp, n); + const hn::ScalableTag df; + CompressPerThread tls; + CompressTraits::Compress(df, tmp, n, tls, MakeSpan(out, n), 0); +} + // Returns the element offset of this layer's segment within a flat KV cache // row. Within the segment: [latent, compressed entry, indexer entry]. static size_t LatentLayerOffset(const KVCachePtr& kv, size_t layer_idx, @@ -119,7 +141,11 @@ static size_t LatentLayerOffset(const KVCachePtr& kv, size_t layer_idx, // default gemma RMSNorm* would apply 1 + w to weights exported as scale-1). static void ScaledRMSNorm(const MatPtr& w, float* HWY_RESTRICT x, size_t n, ThreadingContext& ctx, size_t worker) { - CallUpcastedActivation(&w, [&](const auto* t) { + if (!w.HasPtr()) { + RMSNormNoScaleInplace(x, n, ctx, worker); + return; + } + CallUpcasted(&w, [&](const auto* t) { RMSNormInplace(t->PackedScale1(), /*w_ofs=*/0, x, n, ctx, worker); }); @@ -301,7 +327,7 @@ static HWY_NOINLINE void HCReadDynamic(const MatPtr& fn_w, const MatPtr& base_w, } } } - memcpy(comb, c, hc_mult * hc_mult * sizeof(float)); + hwy::CopyBytes(c, comb, hc_mult * hc_mult * sizeof(float)); }); } @@ -333,7 +359,7 @@ static HWY_NOINLINE void HCWriteDynamic(const MatPtrT& block_out, } MulByConstAndAdd(post[j], out, tmp + j * model_dim, model_dim); } - memcpy(s, tmp, hc_mult * model_dim * sizeof(float)); + hwy::CopyBytes(tmp, s, hc_mult * model_dim * sizeof(float)); }); } @@ -348,7 +374,7 @@ void DeepSeekMaybeInitHCStreams(Activations& activations, MatMulEnv& env) { const float* HWY_RESTRICT x = activations.x.Row(token_idx); float* HWY_RESTRICT s = activations.hc_streams.Row(token_idx); for (size_t i = 0; i < hc_mult; ++i) { - memcpy(s + i * model_dim, x, model_dim * sizeof(float)); + hwy::CopyBytes(x, s + i * model_dim, model_dim * sizeof(float)); } }); } @@ -469,7 +495,7 @@ static bool CompressorStep(const float* HWY_RESTRICT kv_row, const size_t offset = (coff == 2 ? rate : 0) + slot; float* HWY_RESTRICT kv_dst = state.kv_state + offset * width; float* HWY_RESTRICT score_dst = state.score_state + offset * width; - memcpy(kv_dst, kv_row, width * sizeof(float)); + hwy::CopyBytes(kv_row, kv_dst, width * sizeof(float)); { const float* HWY_RESTRICT ape_row = ape + slot * width; size_t i = 0; @@ -524,10 +550,10 @@ static bool CompressorStep(const float* HWY_RESTRICT kv_row, if (coff == 2) { // Shift: the sealed block becomes the next block's overlap window. - memcpy(state.kv_state, state.kv_state + rate * width, - rate * width * sizeof(float)); - memcpy(state.score_state, state.score_state + rate * width, - rate * width * sizeof(float)); + hwy::CopyBytes(state.kv_state + rate * width, state.kv_state, + rate * width * sizeof(float)); + hwy::CopyBytes(state.score_state + rate * width, state.score_state, + rate * width * sizeof(float)); } return true; } @@ -711,15 +737,14 @@ static HWY_NOINLINE void DeepSeekRunCompressors( CompressPerThread tls; Compress(entry, idx_dim, tls, MakeSpan(dst, idx_dim), 0); } - // Speculative decoding: snapshot this layer's state at the - // committed/draft boundary so a rejected draft can be rolled back. - if (HWY_UNLIKELY(static_cast(token_idx) == - activations.ds_snapshot_after) && - cache->ds_state_snapshot.Rows() > 0) { + // Speculative decoding: snapshot this layer's state at each verified + // boundary so a rejected draft can be rolled back to any position. + if (HWY_UNLIKELY(activations.ds_snapshot_after >= 0 && + token_idx < cache->ds_state_snapshot.Rows())) { const size_t ofs = cache->ds_state_offsets[layer_idx]; - memcpy(cache->ds_state_snapshot.Row(0) + ofs, - cache->ds_state.Row(0) + ofs, - lc.DSStateSize() * sizeof(float)); + hwy::CopyBytes(cache->ds_state.Row(0) + ofs, + cache->ds_state_snapshot.Row(token_idx) + ofs, + lc.DSStateSize() * sizeof(float)); } } }); @@ -827,7 +852,7 @@ static HWY_NOINLINE void DeepSeekIndexerSelect( idx_dim); } for (size_t b = nb; b < 4; ++b) { // dummies; results ignored - memcpy(k_f[b], k_f[0], idx_dim * sizeof(float)); + hwy::CopyBytes(k_f[0], k_f[b], idx_dim * sizeof(float)); } HWY_ALIGN float s4[4] = {0.0f, 0.0f, 0.0f, 0.0f}; for (size_t h = 0; h < idx_heads; ++h) { @@ -892,9 +917,6 @@ static HWY_NOINLINE void DeepSeekAttention(size_t num_tokens, size_t layer_idx, const hwy::Divisor div_qbatch(static_cast(qbatch.Size())); const size_t num_interleaved = num_tokens * qbatch.Size(); - att.q.OverrideCols(heads * qkv_dim); - activations.mla_kv_a.OverrideCols(kv_a_dim); - // ---- Projections (tiled MatMul). if (lc.q_lora_rank > 0) { activations.mla_q_a.OverrideCols(lc.q_lora_rank); @@ -909,21 +931,43 @@ static HWY_NOINLINE void DeepSeekAttention(size_t num_tokens, size_t layer_idx, CallMatMul(att.pre_att_rms_out, layer.mla_kv_a, /*add=*/nullptr, env, activations.mla_kv_a); + const bool is_mtp = (layer_idx >= config.num_layers); + const size_t start_pos = qbatch.Pos(0); + + // In DSpark MTP, write main_kv (target layer features) at start_pos. + if (is_mtp && activations.x_bf.Rows() > 0) { + MatPtrT main_x_view("main_x", Extents2D(1, config.model_dim)); + main_x_view.SetPtr(activations.x_bf.Row(0), config.model_dim); + HWY_ALIGN float main_kv[kDSMaxLatentDim]; + MatPtrT main_kv_view("main_kv", Extents2D(1, kv_a_dim)); + main_kv_view.SetPtr(main_kv, kv_a_dim); + CallMatMul(main_x_view, layer.mla_kv_a, /*add=*/nullptr, env, main_kv_view); + ScaledRMSNorm(layer.mla_kv_a_norm, main_kv, kv_a_dim, env.ctx, 0); + Rope(main_kv + kv_a_dim - rope_dim, rope_dim, inv_ts, + static_cast(start_pos), env.ctx, 0); + KV_t* HWY_RESTRICT dst = + qbatch.KV(0).kv_cache.Row(start_pos) + + LatentLayerOffset(qbatch.KV(0), layer_idx, lc); + CompressPerThread tls; + Compress(main_kv, kv_a_dim, tls, MakeSpan(dst, kv_a_dim), 0); + } + // ---- Normalize the full latent, RoPE the decoupled key, write to cache. ParallelFor( Parallelism::kFlat, num_interleaved, env.ctx, /*cluster_idx=*/0, Callers::kAttComputeQKV, [&](size_t task, size_t worker) HWY_ATTR { const size_t qi = div_qbatch.Remainder(static_cast(task)); const size_t token_idx = div_qbatch.Divide(static_cast(task)); - const size_t cache_pos = qbatch.Pos(qi) + token_idx; - HWY_DASSERT(cache_pos < att.SeqLen()); + const size_t kv_pos = is_mtp ? start_pos + 1 + token_idx + : qbatch.Pos(qi) + token_idx; + HWY_DASSERT(kv_pos < att.SeqLen()); float* HWY_RESTRICT kv = activations.mla_kv_a.Row(task); // V4: kv_norm covers the full latent (rope applied after). ScaledRMSNorm(layer.mla_kv_a_norm, kv, kv_a_dim, env.ctx, worker); Rope(kv + kv_a_dim - rope_dim, rope_dim, inv_ts, - static_cast(cache_pos), env.ctx, worker); + static_cast(kv_pos), env.ctx, worker); KV_t* HWY_RESTRICT dst = - qbatch.KV(qi).kv_cache.Row(cache_pos) + + qbatch.KV(qi).kv_cache.Row(kv_pos) + LatentLayerOffset(qbatch.KV(qi), layer_idx, lc); CompressPerThread tls; Compress(kv, kv_a_dim, tls, MakeSpan(dst, kv_a_dim), 0); @@ -959,15 +1003,17 @@ static HWY_NOINLINE void DeepSeekAttention(size_t num_tokens, size_t layer_idx, div_qbatch.Remainder(static_cast(interleaved_idx)); const size_t token_idx = div_qbatch.Divide(static_cast(interleaved_idx)); - const size_t cache_pos = qbatch.Pos(qi) + token_idx; - const size_t end = cache_pos + 1; + const bool is_mtp = (layer_idx >= activations.attention.config.num_layers); + const size_t start_pos = qbatch.Pos(qi); + const size_t q_pos = start_pos + token_idx; + const size_t end = is_mtp ? (start_pos + 1 + num_tokens) : (q_pos + 1); float* HWY_RESTRICT q = att.q.Row(interleaved_idx) + head * qkv_dim; // Per-head RMS (no scale), then RoPE the decoupled part. This task // owns the range. RMSNormNoScaleInplace(q, qkv_dim, env.ctx, worker); Rope(q + qkv_dim - rope_dim, rope_dim, inv_ts, - static_cast(cache_pos), env.ctx, worker); + static_cast(q_pos), env.ctx, worker); const size_t layer_offset = LatentLayerOffset(qbatch.KV(qi), layer_idx, lc); @@ -1021,11 +1067,11 @@ static HWY_NOINLINE void DeepSeekAttention(size_t num_tokens, size_t layer_idx, // Inverse-RoPE the rope dims of the output (values carry rotation). Rope(softmax.acc + kv_a_dim - rope_dim, rope_dim, inv_ts, - -static_cast(cache_pos), env.ctx, worker); + -static_cast(q_pos), env.ctx, worker); float* HWY_RESTRICT out = att.att_out.Row(interleaved_idx) + head * qkv_dim; - memcpy(out, softmax.acc, kv_a_dim * sizeof(float)); + hwy::CopyBytes(softmax.acc, out, kv_a_dim * sizeof(float)); }); // ---- Grouped low-rank output projection. @@ -1267,14 +1313,20 @@ struct DeepSeekMoE { } // Hidden layer -> output layer, via a buffer of the expert's exact width. - MatStorageT C1_narrow("C1_n", - Extents2D(expert_size, expert_ff_hidden_dim), - env.ctx.allocator, MatPadding::kOdd); - for (size_t i = 0; i < expert_size; ++i) { - memcpy(C1_narrow.Row(i), C1.Row(i), expert_ff_hidden_dim * sizeof(BF16)); + if (C1.Cols() == expert_ff_hidden_dim) { + CallMatMul(C1, layer.moe_linear_w[expert_idx], + /*add=*/nullptr, env, expert_out, options); + } else { + MatStorageT C1_narrow("C1_n", + Extents2D(expert_size, expert_ff_hidden_dim), + env.ctx.allocator, MatPadding::kOdd); + for (size_t i = 0; i < expert_size; ++i) { + hwy::CopyBytes(C1.Row(i), C1_narrow.Row(i), + expert_ff_hidden_dim * sizeof(BF16)); + } + CallMatMul(C1_narrow, layer.moe_linear_w[expert_idx], + /*add=*/nullptr, env, expert_out, options); } - CallMatMul(C1_narrow, layer.moe_linear_w[expert_idx], - /*add=*/nullptr, env, expert_out, options); } static HWY_NOINLINE void ComputeAllExpertOutputs( @@ -1471,6 +1523,89 @@ void DeepSeekTransformerLayer(size_t num_tokens, size_t layer_idx, } } +void DeepSeekMaybeSaveDSparkTarget(size_t layer_idx, + Activations& activations) { + if (HWY_LIKELY(activations.dspark_main_hiddens.IsEmpty())) { + return; + } + const size_t num_layers = activations.attention.config.num_layers; + if (layer_idx < num_layers - 3 || layer_idx >= num_layers) { + return; + } + const size_t model_dim = activations.x.Cols(); + const size_t hc_mult = activations.attention.config.hc_mult; + const size_t col_offset = (layer_idx - (num_layers - 3)) * model_dim; + const float inv_mult = 1.0f / static_cast(hc_mult); + const hn::ScalableTag df; + using VF = hn::Vec; + const size_t N = hn::Lanes(df); + const VF vinv_mult = hn::Set(df, inv_mult); + HWY_DASSERT(model_dim % N == 0); + for (size_t r = 0; r < activations.x.Rows(); ++r) { + const float* HWY_RESTRICT s = activations.hc_streams.Row(r); + float* HWY_RESTRICT dst = + activations.dspark_main_hiddens.Row(r) + col_offset; + for (size_t c = 0; c < model_dim; c += N) { + VF vsum = hn::Zero(df); + for (size_t i = 0; i < hc_mult; ++i) { + vsum = hn::Add(vsum, hn::LoadU(df, s + i * model_dim + c)); + } + hn::StoreU(hn::Mul(vsum, vinv_mult), df, dst + c); + } + } +} + +void DeepSeekCommitDSparkKV(size_t num_tokens, size_t pos_base, + const WeightsPtrs& weights, + Activations& activations, QBatch& qbatch, + MatMulEnv& env) { + if (weights.mtp_layers.empty() || activations.dspark_main_hiddens.IsEmpty() || + num_tokens == 0) { + return; + } + const ModelConfig& config = activations.attention.config; + const size_t model_dim = config.model_dim; + const LayerConfig mtp_lc = config.MTPLayerConfig(); + const size_t kv_a_dim = mtp_lc.KVLatentDim(); + const size_t rope_dim = mtp_lc.rope_head_dim; + const float* HWY_RESTRICT inv_ts = + activations.mla_inv_timescale.PackedScale1(); + CompressPerThread tls; + + for (size_t r = 0; r < num_tokens; ++r) { + const size_t pos = pos_base + r; + HWY_ALIGN float ffw_row[kDSMaxHeadDim * 8]; + MatPtrT ffw_view("ffw_row", Extents2D(1, model_dim)); + ffw_view.SetPtr(ffw_row, model_dim); + + MatPtrT r_view("r_view", Extents2D(1, 3 * model_dim)); + r_view.SetPtr(activations.dspark_main_hiddens.Row(r), 3 * model_dim); + CallMatMul(r_view, weights.mtp_main_proj, /*add=*/nullptr, env, ffw_view); + + HWY_ALIGN BF16 main_x[kDSMaxHeadDim * 8]; + MatPtrT main_x_view("main_x", Extents2D(1, model_dim)); + main_x_view.SetPtr(main_x, model_dim); + RMSNormBatched(ffw_view, weights.mtp_main_norm, + main_x_view, env.ctx); + + for (size_t l = 0; l < weights.mtp_layers.size(); ++l) { + const size_t mtp_layer_idx = config.num_layers + l; + const LayerWeightsPtrs& layer = weights.mtp_layers[l]; + HWY_ALIGN float main_kv[kDSMaxLatentDim]; + MatPtrT main_kv_view("main_kv", Extents2D(1, kv_a_dim)); + main_kv_view.SetPtr(main_kv, kv_a_dim); + CallMatMul(main_x_view, layer.mla_kv_a, /*add=*/nullptr, env, main_kv_view); + ScaledRMSNorm(layer.mla_kv_a_norm, main_kv, kv_a_dim, env.ctx, 0); + Rope(main_kv + kv_a_dim - rope_dim, rope_dim, inv_ts, + static_cast(pos), env.ctx, 0); + KV_t* HWY_RESTRICT dst = + qbatch.KV(0).kv_cache.Row(pos) + + LatentLayerOffset(qbatch.KV(0), mtp_layer_idx, mtp_lc); + Compress(main_kv, kv_a_dim, tls, MakeSpan(dst, kv_a_dim), 0); + } + } +} + // Final norm before the output head with plain weights: DeepSeek checkpoints // store the true scale, so gemma's (1 + w) RMSNormBatched must not be used. void DeepSeekFinalNorm(const WeightsPtrs& weights, Activations& activations, @@ -1505,8 +1640,8 @@ void DeepSeekMTPStep(size_t num_tokens, const int* next_tokens, // Stash h = streams in hc_tmp, embed the next tokens into x. activations.token_ids.resize(num_tokens); for (size_t r = 0; r < num_tokens; ++r) { - memcpy(activations.hc_tmp.Row(r), activations.hc_streams.Row(r), - hc_mult * model_dim * sizeof(float)); + hwy::CopyBytes(activations.hc_streams.Row(r), activations.hc_tmp.Row(r), + hc_mult * model_dim * sizeof(float)); activations.token_ids[r] = next_tokens[r]; // Raw embedding row (DeepSeek does not scale embeddings), then enorm. float* HWY_RESTRICT e = activations.x.Row(r); @@ -1554,6 +1689,214 @@ void DeepSeekMTPStep(size_t num_tokens, const int* next_tokens, CallMatMul(activations.x_bf, head, /*add=*/nullptr, env, activations.logits); } +static void ApplyMarkovHeadBias(const float* HWY_RESTRICT markov_embed, + const MatPtr& markov_w2, + float* HWY_RESTRICT logits, + size_t vocab_size, MatMulEnv& env) { + if (!markov_w2.HasPtr() || markov_w2.IsEmpty() || vocab_size == 0) { + return; + } + const size_t cols = markov_w2.Cols(); + if (cols != 256) return; + + namespace hn = hwy::HWY_NAMESPACE; + const hn::ScalableTag df; + using VF = hn::Vec; + const size_t N = hn::Lanes(df); + + // Find max base logit to prune candidate vocabulary items. + const float max_val = hn::MaxValue(df, logits, vocab_size); + const float threshold = max_val - 15.0f; + + for (size_t v = 0; v < vocab_size; ++v) { + if (logits[v] >= threshold) { + HWY_ALIGN float w2_row[256]; + ReadRowF32(markov_w2, v, w2_row, 256); + VF vsum = hn::Zero(df); + for (size_t c = 0; c < 256; c += N) { + vsum = hn::Add(vsum, hn::Mul(hn::Load(df, markov_embed + c), + hn::Load(df, w2_row + c))); + } + logits[v] += hn::ReduceSum(df, vsum); + } + } +} + +static float ComputeDSparkConfidence(const float* HWY_RESTRICT x_row, + const float* HWY_RESTRICT markov_embed, + bool has_markov, const MatPtr& conf_proj, + size_t model_dim) { + if (!conf_proj.HasPtr() || conf_proj.Rows() == 0 || conf_proj.Cols() == 0) { + return 1.0f; + } + const size_t conf_cols = conf_proj.Cols(); + HWY_ALIGN float conf_weights[kDSMaxHeadDim * 8 + 256]; + if (conf_cols > sizeof(conf_weights) / sizeof(conf_weights[0])) return 1.0f; + ReadRowF32(conf_proj, 0, conf_weights, conf_cols); + namespace hn = hwy::HWY_NAMESPACE; + const hn::ScalableTag df; + using VF = hn::Vec; + const size_t N = hn::Lanes(df); + + const size_t x_dim = HWY_MIN(model_dim, conf_cols); + VF vscore = hn::Zero(df); + size_t c = 0; + for (; c + N <= x_dim; c += N) { + vscore = hn::MulAdd(hn::LoadU(df, x_row + c), + hn::Load(df, conf_weights + c), vscore); + } + float score = hn::ReduceSum(df, vscore); + for (; c < x_dim; ++c) { + score += x_row[c] * conf_weights[c]; + } + if (has_markov && conf_cols > model_dim) { + const size_t m_dim = HWY_MIN(size_t{256}, conf_cols - model_dim); + VF vmarkov = hn::Zero(df); + size_t mc = 0; + for (; mc + N <= m_dim; mc += N) { + vmarkov = hn::MulAdd(hn::LoadU(df, markov_embed + mc), + hn::LoadU(df, conf_weights + model_dim + mc), + vmarkov); + } + score += hn::ReduceSum(df, vmarkov); + for (; mc < m_dim; ++mc) { + score += markov_embed[mc] * conf_weights[model_dim + mc]; + } + } + return 1.0f / (1.0f + expf(-score)); +} + +// ------------------------------ DSpark MTP ------------------------------ +// Runs the parallel DSpark multi-layer speculative block. +// In prefill mode (out_drafts == nullptr): runs forward pass for `num_draft_tokens` +// prompt tokens in `next_tokens`. +// In draft mode (out_drafts != nullptr): consumes `pending_token` (next_tokens[0]), +// projects target layer features (dspark_main_hiddens.Row(0)) to main_x, embeds +// [pending, noise, noise, ..., noise] of size block_size = 1 + num_draft_tokens, +// runs all MTP layers in parallel, collapses mHC, applies the Markov head +// transition bias autoregressively, and evaluates confidence scores. +size_t DeepSeekDSparkStep(size_t num_draft_tokens, const int* next_tokens, + int* out_drafts, float* out_confidences, + float confidence_threshold, + const WeightsPtrs& weights, Activations& activations, + QBatch& qbatch, MatMulEnv& env) { + const ModelConfig& config = activations.attention.config; + const size_t model_dim = config.model_dim; + const size_t hc_mult = config.hc_mult; + const size_t vocab_size = config.vocab_size; + + if (out_drafts == nullptr) { + // Prefill mode: next_tokens contains num_draft_tokens tokens. + const size_t num_tokens = num_draft_tokens; + activations.SetBatchSize(num_tokens); + activations.token_ids.resize(num_tokens); + for (size_t r = 0; r < num_tokens; ++r) { + activations.token_ids[r] = next_tokens[r]; + } + CallMatMul(activations.dspark_main_hiddens, weights.mtp_main_proj, + /*add=*/nullptr, env, activations.ffw_out); + RMSNormBatched(activations.ffw_out, + weights.mtp_main_norm, + activations.x_bf, env.ctx); + for (size_t r = 0; r < num_tokens; ++r) { + float* HWY_RESTRICT e = activations.x.Row(r); + const size_t tok = static_cast(next_tokens[r]) % vocab_size; + ReadRowF32(weights.embedder_input_embedding, tok, e, model_dim); + float* HWY_RESTRICT h = activations.hc_streams.Row(r); + for (size_t stream = 0; stream < hc_mult; ++stream) { + hwy::CopyBytes(e, h + stream * model_dim, model_dim * sizeof(float)); + } + } + for (size_t l = 0; l < weights.mtp_layers.size(); ++l) { + const size_t mtp_layer_idx = config.num_layers + l; + DeepSeekTransformerLayer(num_tokens, mtp_layer_idx, weights.mtp_layers[l], + activations, qbatch, env); + } + return num_tokens; + } + + // Speculative Draft Mode: + const size_t num_tokens = num_draft_tokens; + activations.SetBatchSize(num_tokens); + activations.token_ids.resize(num_tokens); + activations.token_ids[0] = next_tokens[0]; // committed token (pending) + for (size_t r = 1; r < num_tokens; ++r) { + activations.token_ids[r] = kDSNoiseTokenId; + } + + // Target features from committed token (row 0). + activations.dspark_main_hiddens.OverrideRows(1); + CallMatMul(activations.dspark_main_hiddens, weights.mtp_main_proj, + /*add=*/nullptr, env, activations.ffw_out); + RMSNormBatched(activations.ffw_out, + weights.mtp_main_norm, + activations.x_bf, env.ctx); + + // Embed draft input tokens and expand to residual streams. + for (size_t r = 0; r < num_tokens; ++r) { + float* HWY_RESTRICT e = activations.x.Row(r); + const size_t tok = + static_cast(activations.token_ids[r]) % config.vocab_size; + ReadRowF32(weights.embedder_input_embedding, tok, e, model_dim); + float* HWY_RESTRICT h = activations.hc_streams.Row(r); + for (size_t stream = 0; stream < hc_mult; ++stream) { + hwy::CopyBytes(e, h + stream * model_dim, model_dim * sizeof(float)); + } + } + + // Run all MTP layers in parallel across the num_tokens rows. + for (size_t l = 0; l < weights.mtp_layers.size(); ++l) { + const size_t mtp_layer_idx = config.num_layers + l; + DeepSeekTransformerLayer(num_tokens, mtp_layer_idx, weights.mtp_layers[l], + activations, qbatch, env); + } + + // Collapse mHC residual streams, apply MTP final norm, and LM head projection. + HCHeadCollapse(weights.mtp_hc_fn, weights.mtp_hc_base, weights.mtp_hc_scale, + activations, env); + RMSNormBatched(activations.x, weights.mtp_norm, + activations.x_bf, env.ctx); + const MatPtr& head = weights.lm_head.HasPtr() + ? weights.lm_head + : weights.embedder_input_embedding; + CallMatMul(activations.x_bf, head, /*add=*/nullptr, env, activations.logits); + + // Autoregressive Markov Head refinement and confidence scoring. + int curr_tok = next_tokens[0]; + size_t actual_drafts = 0; + const bool has_markov = + weights.mtp_markov_w1.HasPtr() && weights.mtp_markov_w2.HasPtr(); + + for (size_t i = 0; i < num_draft_tokens; ++i) { + HWY_ALIGN float markov_embed[256]; + if (has_markov) { + ReadRowF32(weights.mtp_markov_w1, + static_cast(curr_tok) % config.vocab_size, + markov_embed, 256); + ApplyMarkovHeadBias(markov_embed, weights.mtp_markov_w2, + activations.logits.Row(i), config.vocab_size, env); + } + const int draft_tok = Top1OfSoftmax(activations.logits.RowSpan(i)).token; + out_drafts[i] = draft_tok; + curr_tok = draft_tok; + actual_drafts = i + 1; + + if (weights.mtp_conf_proj.HasPtr()) { + const float conf = ComputeDSparkConfidence( + activations.x.Row(i), markov_embed, has_markov, weights.mtp_conf_proj, + model_dim); + if (out_confidences != nullptr) { + out_confidences[i] = conf; + } + if (confidence_threshold > 0.0f && conf < confidence_threshold) { + break; + } + } + } + + return actual_drafts; +} + // NOLINTNEXTLINE(google-readability-namespace-comments) } // namespace HWY_NAMESPACE } // namespace gcpp diff --git a/deepseek/deepseek.h b/deepseek/deepseek.h index 18144fd8..be1a116c 100644 --- a/deepseek/deepseek.h +++ b/deepseek/deepseek.h @@ -49,17 +49,47 @@ namespace gcpp { bool compute_logits, const WeightsPtrs& weights, \ Activations& activations, QBatch& qbatch, \ MatMulEnv& env); \ + /* DSpark multi-layer speculative block: consumes dspark_main_hiddens plus \ + next_tokens, runs all num_mtp_layers speculative layers, and optionally \ + computes draft logits via final norm and output head. */ \ + size_t DeepSeekDSparkStep(size_t num_draft_tokens, const int* next_tokens, \ + int* out_drafts, float* out_confidences, \ + float confidence_threshold, \ + const WeightsPtrs& weights, \ + Activations& activations, QBatch& qbatch, \ + MatMulEnv& env); \ /* Final norm (x -> x_bf) with plain weights; DeepSeek checkpoints store \ the true scale, unlike gemma's (1 + w) convention. */ \ void DeepSeekFinalNorm(const WeightsPtrs& weights, Activations& activations, \ MatMulEnv& env); \ - /* Greedy MTP self-speculative decoding driver (deepseek_spec.cc); \ - called by GenerateT when RuntimeConfig::use_mtp is set. */ \ + void ReadRowF32(const MatPtr& w, size_t row, float* HWY_RESTRICT out, \ + size_t n); \ + void ReadRowBF16(const MatPtr& w, size_t row, BF16* HWY_RESTRICT out, \ + size_t n); \ void GenerateSpecV4(const ModelConfig& config, \ const RuntimeConfig& runtime_config, \ const WeightsPtrs& weights, Activations& activations, \ QBatch& qbatch, MatMulEnv& env, \ TimingInfo& timing_info); \ + /* DSpark EAGLE3 feature fusion hook: saves mean residual stream at target \ + layers 40, 41, 42 into activations.dspark_main_hiddens. No-op unless \ + config.num_mtp_layers > 1 and layer_idx is a target layer. */ \ + void DeepSeekMaybeSaveDSparkTarget(size_t layer_idx, \ + Activations& activations); \ + /* Commits target features for num_tokens starting at pos_base into the \ + draft layers' SWA KV cache. */ \ + void DeepSeekCommitDSparkKV(size_t num_tokens, size_t pos_base, \ + const WeightsPtrs& weights, \ + Activations& activations, QBatch& qbatch, \ + MatMulEnv& env); \ + /* DSpark EAGLE3 multi-layer speculative decoding driver \ + (deepseek_spec.cc); called by GenerateSpecV4 when \ + config.num_mtp_layers > 1. */ \ + void GenerateDSparkV4(const ModelConfig& config, \ + const RuntimeConfig& runtime_config, \ + const WeightsPtrs& weights, Activations& activations, \ + QBatch& qbatch, MatMulEnv& env, \ + TimingInfo& timing_info); \ /* NOLINTNEXTLINE(google-readability-namespace-comments) */ \ } // namespace NAMESPACE diff --git a/deepseek/deepseek_dims.h b/deepseek/deepseek_dims.h index 1db6d1e3..8b285ed4 100644 --- a/deepseek/deepseek_dims.h +++ b/deepseek/deepseek_dims.h @@ -29,6 +29,9 @@ namespace gcpp { +// DSpark speculative noise/mask token ID for parallel draft positions. +static constexpr int kDSNoiseTokenId = 128799; + static inline size_t MaxIndexerHeads(const ModelConfig& config) { size_t max_heads = 0; for (const LayerConfig& lc : config.layer_configs) { diff --git a/deepseek/deepseek_spec.cc b/deepseek/deepseek_spec.cc index e07d881f..b5f66860 100644 --- a/deepseek/deepseek_spec.cc +++ b/deepseek/deepseek_spec.cc @@ -28,6 +28,7 @@ #include "gemma/gemma.h" #include "gemma/kv_cache.h" #include "gemma/weights.h" +#include "util/basics.h" #include "hwy/timer.h" // Compiles this file for multiple architectures via "foreach_target.h", to @@ -57,6 +58,207 @@ static void ComputeLogits(const ModelConfig& config, const WeightsPtrs& weights, FinalLogits(weights, activations, env); } +// Greedy EAGLE3 self-speculative decoding with the DeepSeek V4 DSpark block. +// Each iteration feeds [committed, draft_0..3] (5 tokens) through the main model, +// verifying all 4 speculative tokens in a single forward pass. +void GenerateDSparkV4(const ModelConfig& config, + const RuntimeConfig& runtime_config, + const WeightsPtrs& weights, Activations& activations, + QBatch& qbatch, MatMulEnv& env, TimingInfo& timing_info) { + HWY_ASSERT(qbatch.Size() == 1); + const size_t prefill_max_steps = PrefillTBatchOrQBatch( + config, runtime_config, weights, activations, qbatch, env, timing_info); + const size_t max_gen_steps = + runtime_config.max_generated_tokens > 0 + ? HWY_MIN(prefill_max_steps, runtime_config.max_generated_tokens) + : prefill_max_steps; + const size_t last_prefilled_row = qbatch.Pos(0) - qbatch.InitialPos(0) - 1; + if (last_prefilled_row > 0) { + activations.dspark_main_hiddens.OverrideRows(last_prefilled_row + 1); + hwy::CopyBytes(activations.dspark_main_hiddens.Row(last_prefilled_row), + activations.dspark_main_hiddens.Row(0), + 3 * config.model_dim * sizeof(float)); + DeepSeekCommitDSparkKV(last_prefilled_row + 1, qbatch.InitialPos(0), + weights, activations, qbatch, env); + activations.dspark_main_hiddens.OverrideRows(1); + } + env.ctx.profiler.PrintResults(); + + hwy::BitSet4096<> non_eos; + non_eos.Set(0); + StreamAndUpdateEOSAfterPrefill(config, runtime_config, qbatch, non_eos, 0); + if (!non_eos.Any() || max_gen_steps == 0) return; + + timing_info.generate_start = hwy::platform::Now(); + + size_t stream_pos = qbatch.Pos(0) + 1; + size_t gen = 0; + size_t accepted = 0, rejected = 0; + size_t decode_steps = 0; + + const auto emit = [&](int token) HWY_ATTR -> bool { + timing_info.NotifyGenerated(1); + const bool ok = + runtime_config.StreamToken(qbatch.QueryIdx(0), stream_pos, token, 0.0f); + ++stream_pos; + ++gen; + return ok && !config.IsEOS(token) && gen < max_gen_steps; + }; + + // Bootstrap first token from prompt. + Transformer(config, runtime_config, weights, activations, qbatch, env); + ComputeLogits(config, weights, activations, env); + int pending = Top1OfSoftmax(activations.logits.RowSpan(0)).token; + bool more = emit(pending); + + const size_t max_drafts = HWY_MIN( + size_t{16}, HWY_MAX(size_t{1}, runtime_config.mtp_draft_horizon)); + HWY_ALIGN int drafts[16] = {}; + HWY_ALIGN float confidences[16] = {}; + size_t actual_drafts = 0; + HWY_ALIGN size_t slot_tested[16] = {}; + HWY_ALIGN size_t slot_accepted[16] = {}; + + // Parallel DSpark Drafter: generates up to max_drafts speculative tokens in a single parallel pass. + const auto generate_drafts = [&]() HWY_ATTR { + for (size_t j = 0; j < 16; ++j) { + drafts[j] = 0; + confidences[j] = 1.0f; + } + actual_drafts = DeepSeekDSparkStep( + max_drafts, &pending, drafts, confidences, + runtime_config.mtp_confidence_threshold, weights, activations, + qbatch, env); + + for (size_t i = 0; i < actual_drafts; ++i) { + MaybePrint(2, runtime_config.verbosity, + " [draft slot %zu] draft_top=%d, conf=%.3f", i, drafts[i], + confidences[i]); + } + }; + + if (more) { + generate_drafts(); + } + qbatch.MutablePos(0) += 1; + + while (more) { + ++decode_steps; + const size_t pos = qbatch.Pos(0); + const size_t block_size = 1 + actual_drafts; + activations.SetBatchSize(block_size); + activations.token_ids.resize(block_size); + activations.token_ids[0] = pending; + for (size_t i = 0; i < actual_drafts; ++i) { + activations.token_ids[i + 1] = drafts[i]; + } + for (size_t i = 0; i < block_size; ++i) { + EmbedMMToken(activations.token_ids[i], i, pos + i, /*pos_in_prompt=*/0, + config, weights, activations.x, env.ctx, + /*image_tokens=*/nullptr, /*image_token_position=*/0); + } + DeepSeekMaybeInitHCStreams(activations, env); + activations.ds_snapshot_after = 0; + for (size_t layer_idx = 0; layer_idx < weights.c_layers.size(); + ++layer_idx) { + TransformerLayer(block_size, layer_idx, *weights.GetLayer(layer_idx), + activations, qbatch, env); + } + activations.ds_snapshot_after = -1; + DeepSeekMaybeFinalizeHCStreams(weights, activations, env); + ComputeLogits(config, weights, activations, env); + + // Verify drafts sequentially. + size_t num_acc = 0; + int next_committed = Top1OfSoftmax(activations.logits.RowSpan(0)).token; + const bool debug_log = runtime_config.verbosity >= 2; + if (debug_log) { + fprintf(stderr, "[STEP %zu] pending=%d, %zu drafts: [", decode_steps, pending, actual_drafts); + for (size_t d = 0; d < actual_drafts; ++d) fprintf(stderr, "%d ", drafts[d]); + fprintf(stderr, "] -> base top0=%d", next_committed); + } + + more = emit(next_committed); + while (more && num_acc < actual_drafts && + next_committed == drafts[num_acc]) { + ++num_acc; + next_committed = Top1OfSoftmax(activations.logits.RowSpan(num_acc)).token; + if (debug_log) { + fprintf(stderr, ", acc draft[%zu]=%d -> top%zu=%d", num_acc - 1, drafts[num_acc - 1], num_acc, next_committed); + } + more = emit(next_committed); + } + if (debug_log) { + fprintf(stderr, ", total accepted=%zu\n", num_acc); + } + + for (size_t d = 0; d < actual_drafts; ++d) { + if (d <= num_acc) { + slot_tested[d]++; + if (d < num_acc) { + slot_accepted[d]++; + } + } + } + + accepted += num_acc; + rejected += (actual_drafts - num_acc); + pending = next_committed; + + if (num_acc < actual_drafts) { + // Roll back compressor state to after the last accepted token. + KVCache* cache = qbatch.KV(0).cache; + if (cache != nullptr && cache->ds_state.Rows() > 0) { + hwy::CopyBytes(cache->ds_state_snapshot.Row(num_acc), + cache->ds_state.Row(0), + cache->ds_state.Cols() * sizeof(float)); + } + } + + if (more) { + if (num_acc > 0) { + DeepSeekCommitDSparkKV(num_acc, pos, weights, activations, qbatch, env); + hwy::CopyBytes(activations.dspark_main_hiddens.Row(num_acc), + activations.dspark_main_hiddens.Row(0), + 3 * config.model_dim * sizeof(float)); + } + qbatch.MutablePos(0) += 1 + num_acc; + generate_drafts(); + } + } + + timing_info.NotifyGenerateDone(); + if (runtime_config.verbosity >= 1) { + const size_t total_drafts = accepted + rejected; + const double avg_tokens_per_step = + decode_steps ? static_cast(gen) / static_cast(decode_steps) : 1.0; + const double avg_accepted_per_step = + decode_steps ? static_cast(accepted) / static_cast(decode_steps) : 0.0; + fprintf(stderr, + "\n[ DSpark MTP Summary ]\n" + " Total generated tokens : %zu\n" + " Total decode steps : %zu\n" + " Avg tokens / step : %.2f\n" + " Marginal accept rate : %zu / %zu (%.1f%%)\n" + " Avg accepted / step : %.2f\n" + " --- Per-Slot Conditional Acceptance Rates ---\n", + gen, decode_steps, avg_tokens_per_step, + accepted, total_drafts, + total_drafts ? 100.0 * static_cast(accepted) / static_cast(total_drafts) : 0.0, + avg_accepted_per_step); + for (size_t d = 0; d < max_drafts; ++d) { + if (slot_tested[d] > 0) { + const double cond_rate = + 100.0 * static_cast(slot_accepted[d]) / + static_cast(slot_tested[d]); + fprintf(stderr, + " Slot %zu: tested=%zu, accepted=%zu -> Conditional Rate: %.1f%%\n", + d, slot_tested[d], slot_accepted[d], cond_rate); + } + } + } +} + // Greedy self-speculative decoding with the DeepSeek V4 MTP block: each // iteration feeds [committed, draft] through the main model in one 2-token // pass and reads logits at both positions. An accepted draft yields two @@ -70,6 +272,12 @@ void GenerateSpecV4(const ModelConfig& config, const WeightsPtrs& weights, Activations& activations, QBatch& qbatch, MatMulEnv& env, TimingInfo& timing_info) { HWY_ASSERT(qbatch.Size() == 1); + if (weights.mtp_main_proj.HasPtr() || + config.model == Model::DEEPSEEK4_FLASH) { + GenerateDSparkV4(config, runtime_config, weights, activations, qbatch, env, + timing_info); + return; + } const size_t max_gen_steps = PrefillTBatchOrQBatch( config, runtime_config, weights, activations, qbatch, env, timing_info); env.ctx.profiler.PrintResults(); @@ -87,7 +295,7 @@ void GenerateSpecV4(const ModelConfig& config, size_t accepted = 0, rejected = 0; // Streams `token`; returns false if generation should stop. - const auto emit = [&](int token) -> bool { + const auto emit = [&](int token) HWY_ATTR -> bool { timing_info.NotifyGenerated(1); const bool ok = runtime_config.StreamToken(qbatch.QueryIdx(0), stream_pos, token, 0.0f); @@ -138,10 +346,9 @@ void GenerateSpecV4(const ModelConfig& config, if (true1 == draft) { ++accepted; const int true2 = Top1OfSoftmax(activations.logits.RowSpan(1)).token; - if (getenv("GCPP_SPEC_DEBUG")) { - fprintf(stderr, "[spec] pos=%zu ACCEPT pending=%d draft=%d true2=%d\n", - pos, pending, draft, true2); - } + MaybePrint(2, runtime_config.verbosity, + "[spec] pos=%zu ACCEPT pending=%d draft=%d true2=%d", pos, + pending, draft, true2); more = emit(true1); if (more) more = emit(true2); if (more) { @@ -155,10 +362,9 @@ void GenerateSpecV4(const ModelConfig& config, qbatch.MutablePos(0) += 2; } else { ++rejected; - if (getenv("GCPP_SPEC_DEBUG")) { - fprintf(stderr, "[spec] pos=%zu REJECT pending=%d draft=%d true1=%d\n", - pos, pending, draft, true1); - } + MaybePrint(2, runtime_config.verbosity, + "[spec] pos=%zu REJECT pending=%d draft=%d true1=%d", pos, + pending, draft, true1); more = emit(true1); // Roll back the draft row's compressor state to the boundary snapshot. KVCache* cache = qbatch.KV(0).cache; @@ -176,11 +382,14 @@ void GenerateSpecV4(const ModelConfig& config, } } timing_info.NotifyGenerateDone(); - const size_t drafts = accepted + rejected; - fprintf(stderr, "MTP: accepted %zu / %zu drafts (%.1f%%)\n", accepted, drafts, - drafts ? 100.0 * static_cast(accepted) / - static_cast(drafts) - : 0.0); + if (runtime_config.verbosity >= 1) { + const size_t drafts = accepted + rejected; + fprintf(stderr, "MTP: accepted %zu / %zu drafts (%.1f%%)\n", accepted, + drafts, + drafts ? 100.0 * static_cast(accepted) / + static_cast(drafts) + : 0.0); + } } // NOLINTNEXTLINE(google-readability-namespace-comments) diff --git a/deepseek/deepseek_tensors.cc b/deepseek/deepseek_tensors.cc index 22b51236..53a0b773 100644 --- a/deepseek/deepseek_tensors.cc +++ b/deepseek/deepseek_tensors.cc @@ -65,8 +65,44 @@ void TensorInfoRegistry::AddDeepSeekModelTensors(const ModelConfig& config) { add_hc_collapse("hc_head", ""); } if (config.num_mtp_layers > 0) { - // Multi-token-prediction block extras (DeepSeek V4 `mtp.0.*`). The block - // itself is registered as an extra layer, see the ctor. + Add(no_suffix, { + .base_name = "mtp_main_proj", + .source_names = {"mtp.0.main_proj.weight"}, + .axes = {0, 1}, + .shape = {config.model_dim, 3 * config.model_dim}, + }); + Add(no_suffix, { + .base_name = "mtp_main_norm", + .source_names = {"mtp.0.main_norm.weight"}, + .axes = {0}, + .shape = {config.model_dim}, + .min_size = Type::kBF16, + }); + const std::string last_mtp = + "mtp." + std::to_string(config.num_mtp_layers > 0 + ? config.num_mtp_layers - 1 + : 0) + + "."; + Add(no_suffix, { + .base_name = "mtp_markov_w1", + .source_names = {last_mtp + "markov_head.markov_w1.weight", + "mtp.0.markov_head.markov_w1.weight"}, + .axes = {0, 1}, + .shape = {config.vocab_size, 256}, + }); + Add(no_suffix, { + .base_name = "mtp_markov_w2", + .source_names = {last_mtp + "markov_head.markov_w2.weight", + "mtp.0.markov_head.markov_w2.weight"}, + .axes = {0, 1}, + .shape = {config.vocab_size, 256}, + }); + Add(no_suffix, { + .base_name = "mtp_conf_proj", + .source_names = {"mtp.0.confidence_head.proj.weight"}, + .axes = {0, 1}, + .shape = {1, config.model_dim + 256}, + }); Add(no_suffix, { .base_name = "mtp_e_proj", .source_names = {"mtp.0.e_proj.weight"}, @@ -95,12 +131,13 @@ void TensorInfoRegistry::AddDeepSeekModelTensors(const ModelConfig& config) { }); Add(no_suffix, { .base_name = "mtp_norm", - .source_names = {"mtp.0.norm.weight"}, + .source_names = {last_mtp + "norm.weight", + "mtp.0.norm.weight"}, .axes = {0}, .shape = {config.model_dim}, .min_size = Type::kBF16, }); - add_hc_collapse("mtp_hc", "mtp.0."); + add_hc_collapse("mtp_hc", last_mtp); } } diff --git a/deepseek/deepseek_test.cc b/deepseek/deepseek_test.cc index 1d688da3..46b9e461 100644 --- a/deepseek/deepseek_test.cc +++ b/deepseek/deepseek_test.cc @@ -193,7 +193,7 @@ void TestDeepSeekTiny() { // "gating_ein" alias (mutually exclusive with the split w1/w2 tensors) and // the optional skip_scale. uint64_t seed = 1; - weights.ForEachTensor(nullptr, nullptr, [&](const TensorArgs& t) { + weights.ForEachTensor(nullptr, nullptr, [&](const TensorArgs& t) HWY_ATTR { const char* name = t.mat.Name(); if (strncmp(name, "gating_ein", 10) == 0) return; if (strncmp(name, "skip_scale", 10) == 0) return; @@ -204,7 +204,7 @@ void TestDeepSeekTiny() { { LayerWeightsPtrs* layer = weights.GetLayer(3); const uint32_t num_experts = layer->layer_config.NumExperts(); - SetF32(layer->hash_tid2eid, [&](size_t r, size_t c) { + SetF32(layer->hash_tid2eid, [&](size_t r, size_t c) HWY_ATTR { return static_cast((r + c) % num_experts); }); } @@ -212,10 +212,10 @@ void TestDeepSeekTiny() { // weights. for (size_t i = 0; i < config.num_layers; ++i) { LayerWeightsPtrs* layer = weights.GetLayer(i); - SetF32(layer->hc_att_scale, [](size_t, size_t) { return 0.1f; }); - SetF32(layer->hc_ffw_scale, [](size_t, size_t) { return 0.1f; }); + SetF32(layer->hc_att_scale, [](size_t, size_t) HWY_ATTR { return 0.1f; }); + SetF32(layer->hc_ffw_scale, [](size_t, size_t) HWY_ATTR { return 0.1f; }); } - SetF32(weights.hc_head_scale, [](size_t, size_t) { return 0.1f; }); + SetF32(weights.hc_head_scale, [](size_t, size_t) HWY_ATTR { return 0.1f; }); InferenceArgs inference_args; RuntimeConfig runtime_config; @@ -297,7 +297,8 @@ static std::vector GreedyGenerate(const ModelConfig& config, runtime_config.top_k = 1; // greedy runtime_config.verbosity = 0; runtime_config.use_mtp = use_mtp; - runtime_config.batch_stream_token = [&](size_t, size_t, int token, float) { + runtime_config.batch_stream_token = + [&](size_t, size_t, int token, float) HWY_ATTR { if (++seen <= prompt_size) return true; generated.push_back(token); return true; @@ -336,7 +337,7 @@ void TestDeepSeekMTPEquivalence() { WeightsPtrs weights(config); ASSERT_EQ(weights.mtp_layers.size(), size_t{1}); uint64_t seed = 100; - weights.ForEachTensor(nullptr, nullptr, [&](const TensorArgs& t) { + weights.ForEachTensor(nullptr, nullptr, [&](const TensorArgs& t) HWY_ATTR { const char* name = t.mat.Name(); if (strncmp(name, "gating_ein", 10) == 0) return; if (strncmp(name, "skip_scale", 10) == 0) return; @@ -345,25 +346,25 @@ void TestDeepSeekMTPEquivalence() { { LayerWeightsPtrs* layer = weights.GetLayer(3); const uint32_t num_experts = layer->layer_config.NumExperts(); - SetF32(layer->hash_tid2eid, [&](size_t r, size_t c) { + SetF32(layer->hash_tid2eid, [&](size_t r, size_t c) HWY_ATTR { return static_cast((r + c) % num_experts); }); } for (size_t i = 0; i < config.num_layers; ++i) { LayerWeightsPtrs* layer = weights.GetLayer(i); - SetF32(layer->hc_att_scale, [](size_t, size_t) { return 0.1f; }); - SetF32(layer->hc_ffw_scale, [](size_t, size_t) { return 0.1f; }); + SetF32(layer->hc_att_scale, [](size_t, size_t) HWY_ATTR { return 0.1f; }); + SetF32(layer->hc_ffw_scale, [](size_t, size_t) HWY_ATTR { return 0.1f; }); } - SetF32(weights.hc_head_scale, [](size_t, size_t) { return 0.1f; }); + SetF32(weights.hc_head_scale, [](size_t, size_t) HWY_ATTR { return 0.1f; }); SetF32(weights.mtp_layers[0].hc_att_scale, - [](size_t, size_t) { return 0.1f; }); + [](size_t, size_t) HWY_ATTR { return 0.1f; }); SetF32(weights.mtp_layers[0].hc_ffw_scale, - [](size_t, size_t) { return 0.1f; }); - SetF32(weights.mtp_hc_scale, [](size_t, size_t) { return 0.1f; }); + [](size_t, size_t) HWY_ATTR { return 0.1f; }); + SetF32(weights.mtp_hc_scale, [](size_t, size_t) HWY_ATTR { return 0.1f; }); // Sparse output head: only tokens 3/7/11 can win, with O(1) margins; the // remaining rows are exactly zero, so their logits tie at 0.0 bitwise. - SetF32(weights.lm_head, [](size_t r, size_t c) { + SetF32(weights.lm_head, [](size_t r, size_t c) HWY_ATTR { if (r != 3 && r != 7 && r != 11) return 0.0f; return 2.0f * sinf(131.3f * static_cast(r) + 0.71f * static_cast(c)); @@ -377,15 +378,15 @@ void TestDeepSeekMTPEquivalence() { config, weights, prompt_vec, kMaxTokens, /*use_mtp=*/false, ctx, env); EXPECT_EQ(ref.size(), kMaxTokens); - for (int regime = 0; regime < 2; ++regime) { + for (size_t regime = 0; regime < 2; ++regime) { if (regime == 1) { // Corrupt the MTP input projections: drafts become unrelated to the // main model's output, so most verify steps reject and roll back. - SetF32(weights.mtp_e_proj, [](size_t r, size_t c) { + SetF32(weights.mtp_e_proj, [](size_t r, size_t c) HWY_ATTR { return 0.5f * cosf(17.7f * static_cast(r) - 1.3f * static_cast(c)); }); - SetF32(weights.mtp_h_proj, [](size_t r, size_t c) { + SetF32(weights.mtp_h_proj, [](size_t r, size_t c) HWY_ATTR { return 0.5f * sinf(3.9f * static_cast(r) + 11.1f * static_cast(c)); }); @@ -415,7 +416,7 @@ void TestDeepSeekVerifyStepState() { const ModelConfig config = TinyDeepSeekConfig(); WeightsPtrs weights(config); uint64_t seed = 500; - weights.ForEachTensor(nullptr, nullptr, [&](const TensorArgs& t) { + weights.ForEachTensor(nullptr, nullptr, [&](const TensorArgs& t) HWY_ATTR { const char* name = t.mat.Name(); if (strncmp(name, "gating_ein", 10) == 0) return; if (strncmp(name, "skip_scale", 10) == 0) return; @@ -424,16 +425,16 @@ void TestDeepSeekVerifyStepState() { { LayerWeightsPtrs* layer = weights.GetLayer(3); const uint32_t num_experts = layer->layer_config.NumExperts(); - SetF32(layer->hash_tid2eid, [&](size_t r, size_t c) { + SetF32(layer->hash_tid2eid, [&](size_t r, size_t c) HWY_ATTR { return static_cast((r + c) % num_experts); }); } for (size_t i = 0; i < config.num_layers; ++i) { LayerWeightsPtrs* layer = weights.GetLayer(i); - SetF32(layer->hc_att_scale, [](size_t, size_t) { return 0.1f; }); - SetF32(layer->hc_ffw_scale, [](size_t, size_t) { return 0.1f; }); + SetF32(layer->hc_att_scale, [](size_t, size_t) HWY_ATTR { return 0.1f; }); + SetF32(layer->hc_ffw_scale, [](size_t, size_t) HWY_ATTR { return 0.1f; }); } - SetF32(weights.hc_head_scale, [](size_t, size_t) { return 0.1f; }); + SetF32(weights.hc_head_scale, [](size_t, size_t) HWY_ATTR { return 0.1f; }); InferenceArgs inference_args; RuntimeConfig runtime_config; @@ -453,7 +454,7 @@ void TestDeepSeekVerifyStepState() { std::vector tokens(kTotal); std::iota(tokens.begin(), tokens.end(), 2); - const auto embed = [&](int token, size_t row, Activations& acts) { + const auto embed = [&](int token, size_t row, Activations& acts) HWY_ATTR { MatPtrT emb(weights.embedder_input_embedding); memcpy(acts.x.Row(row), emb.Row(static_cast(token)), config.model_dim * sizeof(float)); @@ -475,12 +476,12 @@ void TestDeepSeekVerifyStepState() { Activations aB(runtime_config, config, kPrompt, spec_kv.SeqLen(), ctx, env.row_ptrs); - const auto layers = [&](Activations& a, QBatch& q, size_t num_tokens) { + const auto layers = [&](Activations& a, QBatch& q, size_t num_tokens) HWY_ATTR { for (size_t l = 0; l < config.num_layers; ++l) { DeepSeekTransformerLayer(num_tokens, l, *weights.GetLayer(l), a, q, env); } }; - const auto step1 = [&](Activations& a, QBatch& q, int tok) { + const auto step1 = [&](Activations& a, QBatch& q, int tok) HWY_ATTR { a.SetBatchSize(1); a.token_ids.assign(1, tok); embed(tok, 0, a); @@ -488,7 +489,7 @@ void TestDeepSeekVerifyStepState() { layers(a, q, 1); q.MutablePos(0) += 1; }; - const auto step2 = [&](Activations& a, QBatch& q, int tok0, int tok1) { + const auto step2 = [&](Activations& a, QBatch& q, int tok0, int tok1) HWY_ATTR { a.SetBatchSize(2); a.token_ids.assign(2, tok0); a.token_ids[1] = tok1; @@ -500,7 +501,7 @@ void TestDeepSeekVerifyStepState() { a.ds_snapshot_after = -1; }; - const auto prefill = [&](Activations& a, QBatch& q) { + const auto prefill = [&](Activations& a, QBatch& q) HWY_ATTR { a.SetBatchSize(kPrompt); a.token_ids.resize(kPrompt); for (size_t i = 0; i < kPrompt; ++i) { @@ -568,6 +569,111 @@ void TestDeepSeekVerifyStepState() { EXPECT_EQ(kv_mismatches, size_t{0}); } +void TestDeepSeekDSpark() { + ThreadingContext ctx({}); + MatMulEnv env(ctx); + std::vector mat_owners; + + ModelConfig config = TinyDeepSeekConfig(); + config.num_mtp_layers = 3; + config.eos_id = static_cast(config.vocab_size); + config.secondary_eos_id = static_cast(config.vocab_size); + + WeightsPtrs weights(config); + ASSERT_EQ(weights.mtp_layers.size(), size_t{3}); + uint64_t seed = 200; + weights.ForEachTensor(nullptr, nullptr, [&](const TensorArgs& t) HWY_ATTR { + const char* name = t.mat.Name(); + if (strncmp(name, "gating_ein", 10) == 0) return; + if (strncmp(name, "skip_scale", 10) == 0) return; + AllocateAndFillRandom(t.mat, ctx.allocator, mat_owners, ++seed); + }); + { + LayerWeightsPtrs* layer = weights.GetLayer(3); + const uint32_t num_experts = layer->layer_config.NumExperts(); + SetF32(layer->hash_tid2eid, [&](size_t r, size_t c) HWY_ATTR { + return static_cast((r + c) % num_experts); + }); + } + for (size_t i = 0; i < config.num_layers; ++i) { + LayerWeightsPtrs* layer = weights.GetLayer(i); + SetF32(layer->hc_att_scale, [](size_t, size_t) HWY_ATTR { return 0.1f; }); + SetF32(layer->hc_ffw_scale, [](size_t, size_t) HWY_ATTR { return 0.1f; }); + } + SetF32(weights.hc_head_scale, [](size_t, size_t) HWY_ATTR { return 0.1f; }); + for (size_t l = 0; l < weights.mtp_layers.size(); ++l) { + SetF32(weights.mtp_layers[l].hc_att_scale, + [](size_t, size_t) HWY_ATTR { return 0.0f; }); + SetF32(weights.mtp_layers[l].hc_ffw_scale, + [](size_t, size_t) HWY_ATTR { return 0.0f; }); + } + SetF32(weights.mtp_hc_scale, [](size_t, size_t) HWY_ATTR { return 0.1f; }); + + SetF32(weights.lm_head, [](size_t r, size_t c) HWY_ATTR { + if (r != 3 && r != 7 && r != 11) return 0.0f; + return 2.0f * + sinf(131.3f * static_cast(r) + 0.71f * static_cast(c)); + }); + + // Align DSpark MTP weights with base model output in regime 0. + SetF32(weights.mtp_main_proj, [](size_t r, size_t c) HWY_ATTR { + return (r % 3 == 0 && c < 64) ? 0.2f : 0.0f; + }); + SetF32(weights.mtp_main_norm, [](size_t, size_t) HWY_ATTR { return 1.0f; }); + SetF32(weights.mtp_norm, [](size_t, size_t) HWY_ATTR { return 1.0f; }); + + // Rank-256 Markov transition bias: predicts dominant token 3 for high acceptance. + SetF32(weights.mtp_markov_w1, [](size_t r, size_t c) HWY_ATTR { + if (r == 3 && c == 0) return 3.0f; + if (r == 7 && c == 1) return 3.0f; + if (r == 11 && c == 2) return 3.0f; + return 0.0f; + }); + SetF32(weights.mtp_markov_w2, [](size_t r, size_t c) HWY_ATTR { + if (r == 3 && c == 0) return 5.0f; + if (r == 7 && c == 1) return 5.0f; + if (r == 11 && c == 2) return 5.0f; + return 0.0f; + }); + + // Confidence head: positive bias for high speculative confidence (>0.9). + SetF32(weights.mtp_conf_proj, [](size_t, size_t c) HWY_ATTR { + return (c == 0) ? 2.5f : 0.01f; + }); + + std::vector prompt_vec(12); + std::iota(prompt_vec.begin(), prompt_vec.end(), 2); + const size_t kMaxTokens = 40; + + const std::vector ref = GreedyGenerate( + config, weights, prompt_vec, kMaxTokens, /*use_mtp=*/false, ctx, env); + EXPECT_EQ(ref.size(), kMaxTokens); + + for (size_t regime = 0; regime < 2; ++regime) { + if (regime == 1) { + SetF32(weights.mtp_main_proj, [](size_t r, size_t c) HWY_ATTR { + return 0.5f * cosf(17.7f * static_cast(r) - + 1.3f * static_cast(c)); + }); + SetF32(weights.mtp_markov_w1, [](size_t r, size_t c) HWY_ATTR { + return 0.2f * sinf(static_cast(r * 256 + c)); + }); + } + const std::vector spec = GreedyGenerate( + config, weights, prompt_vec, kMaxTokens, /*use_mtp=*/true, ctx, env); + ASSERT_EQ(spec.size(), kMaxTokens) << "regime " << regime; + EXPECT_EQ(ref[0], spec[0]) << "regime " << regime; + for (size_t i = 0; i < spec.size(); ++i) { + const int t = spec[i]; + ASSERT_TRUE(t == 0 || t == 3 || t == 7 || t == 11) + << "regime " << regime << ": unreachable token " << t << " at " << i; + } + if (regime == 0) { + EXPECT_EQ(spec, ref) << "Exact token equivalence with base greedy decode"; + } + } +} + } // namespace HWY_NAMESPACE } // namespace gcpp HWY_AFTER_NAMESPACE(); @@ -579,6 +685,7 @@ HWY_BEFORE_TEST(DeepSeekTest); HWY_EXPORT_AND_TEST_P(DeepSeekTest, TestDeepSeekTiny); HWY_EXPORT_AND_TEST_P(DeepSeekTest, TestDeepSeekVerifyStepState); HWY_EXPORT_AND_TEST_P(DeepSeekTest, TestDeepSeekMTPEquivalence); +HWY_EXPORT_AND_TEST_P(DeepSeekTest, TestDeepSeekDSpark); HWY_AFTER_TEST(); } // namespace gcpp diff --git a/deepseek/run_dsv4.cc b/deepseek/run_dsv4.cc index 3bed4a05..c1f06b9d 100644 --- a/deepseek/run_dsv4.cc +++ b/deepseek/run_dsv4.cc @@ -82,6 +82,9 @@ int Main(int argc, char** argv) { std::string prompt_text; std::string tokenizer_json; bool use_mtp = false; + size_t mtp_draft_horizon = 7; + float mtp_confidence_threshold = 0.0f; + size_t max_generated_tokens = 0; bool thinking = false; bool raw = false; bool tokenize_only = false; @@ -97,6 +100,13 @@ int Main(int argc, char** argv) { tokenizer_json = argv[++i]; } else if (arg == "--mtp") { use_mtp = true; + } else if (arg == "--mtp_draft_horizon" && i + 1 < argc) { + mtp_draft_horizon = static_cast(std::stoull(argv[++i])); + } else if (arg == "--mtp_confidence_threshold" && i + 1 < argc) { + mtp_confidence_threshold = std::stof(argv[++i]); + } else if ((arg == "--max_generated_tokens" || arg == "--max_tokens") && + i + 1 < argc) { + max_generated_tokens = static_cast(std::stoull(argv[++i])); } else if (arg == "--thinking") { thinking = true; } else if (arg == "--raw") { @@ -186,8 +196,13 @@ int Main(int argc, char** argv) { .use_spinning = args.threading.spin, }; args.inference.CopyTo(runtime_config); + if (max_generated_tokens > 0) { + runtime_config.max_generated_tokens = max_generated_tokens; + } runtime_config.use_mtp = use_mtp; - if (use_mtp) fprintf(stderr, "MTP speculative decoding enabled.\n"); + runtime_config.mtp_draft_horizon = mtp_draft_horizon; + runtime_config.mtp_confidence_threshold = mtp_confidence_threshold; + if (use_mtp) fprintf(stderr, "DSpark MTP speculative decoding enabled.\n"); TimingInfo timing_info = {.verbosity = args.inference.verbosity}; const PromptTokens prompt(prompt_vec); diff --git a/gemma/activations.h b/gemma/activations.h index bfcbf56b..724931b9 100644 --- a/gemma/activations.h +++ b/gemma/activations.h @@ -49,7 +49,10 @@ static inline size_t MaxQkvDim(const ModelConfig& config) { return max_dim; } static inline size_t MaxFFHiddenDim(const ModelConfig& config) { - size_t max_dim = 0; + size_t max_dim = config.model_dim; + if (config.num_mtp_layers > 1) { + max_dim = HWY_MAX(max_dim, size_t{256}); + } for (const auto& lc : config.layer_configs) { max_dim = HWY_MAX(max_dim, static_cast(lc.ff_hidden_dim)); } @@ -494,6 +497,9 @@ struct Activations { hc_streams(MatFactory("hc_streams", config.hc_mult > 1 ? batch_size : 0, config.hc_mult * config.model_dim, ctx.allocator)), + dspark_main_hiddens(MatFactory( + "dspark_hiddens", config.num_mtp_layers > 0 ? batch_size : 0, + 3 * config.model_dim, ctx.allocator)), hc_tmp(MatFactory("hc_tmp", config.hc_mult > 1 ? batch_size : 0, config.hc_mult * config.model_dim, ctx.allocator)), hc_mixes(MatFactory("hc_mixes", config.hc_mult > 1 ? batch_size : 0, @@ -666,6 +672,9 @@ struct Activations { hc_post_w.OverrideRows(batch_size); hc_comb.OverrideRows(batch_size); } + if (dspark_main_hiddens.Cols() > 0 && dspark_main_hiddens.Rows() > 0) { + dspark_main_hiddens.OverrideRows(batch_size); + } attention_storage.SetBatchSize(batch_size); // `AttentionActivationsPtrs` holds `MatPtrT` which also require updating; @@ -744,6 +753,7 @@ struct Activations { // Manifold-constrained hyper-connections: parallel residual streams, // [batch_size, hc_mult * model_dim] (zero-sized unless hc_mult > 1). MatStorageT hc_streams; + MatStorageT dspark_main_hiddens; MatStorageT hc_tmp; // Per-token dynamic mHC weights: raw mixes and the split post/comb parts // persisted between the block's read and write phases. diff --git a/gemma/configs.h b/gemma/configs.h index 37346098..4051259b 100644 --- a/gemma/configs.h +++ b/gemma/configs.h @@ -774,8 +774,10 @@ struct ModelConfig : public IFields { for (const auto& lc : layer_configs) { cols += lc.CacheLayerSize(); } - // The MTP block caches its latents in an extra trailing segment. - if (num_mtp_layers > 0) cols += MTPLayerConfig().CacheLayerSize(); + // The MTP block caches its latents in an extra trailing segment per layer. + if (num_mtp_layers > 0) { + cols += num_mtp_layers * MTPLayerConfig().CacheLayerSize(); + } return cols; } diff --git a/gemma/gemma.cc b/gemma/gemma.cc index 5dd8535e..b0999334 100644 --- a/gemma/gemma.cc +++ b/gemma/gemma.cc @@ -114,6 +114,7 @@ HWY_NOINLINE void TransformerLayer(const size_t num_tokens, if (layer_config.type == LayerAttentionType::kDeepSeekMLA) { DeepSeekTransformerLayer(num_tokens, layer_idx, layer, activations, qbatch, env); + DeepSeekMaybeSaveDSparkTarget(layer_idx, activations); return; } if (layer_config.IsMoE() && @@ -1233,8 +1234,16 @@ static HWY_NOINLINE void PrefillTBatch(const ModelConfig& config, for (size_t ti = 0; ti < tbatch_size; ++ti) { next[ti] = qbatch_1.Prompt(0)[tbatch_start + ti + 1]; } - DeepSeekMTPStep(tbatch_size, next.data(), /*compute_logits=*/false, - weights, activations, qbatch_1, env); + if (weights.mtp_main_proj.HasPtr() || + config.model == Model::DEEPSEEK4_FLASH) { + DeepSeekDSparkStep(tbatch_size, next.data(), /*out_drafts=*/nullptr, + /*out_confidences=*/nullptr, + /*confidence_threshold=*/0.0f, weights, + activations, qbatch_1, env); + } else { + DeepSeekMTPStep(tbatch_size, next.data(), /*compute_logits=*/false, + weights, activations, qbatch_1, env); + } } qbatch_1.MutablePos(0) += tbatch_size; @@ -1605,14 +1614,13 @@ static void GenerateT(const ModelConfig& config, if (HWY_UNLIKELY(runtime_config.use_mtp)) { if (!weights.mtp_layers.empty() && qbatch.Size() == 1 && - runtime_config.top_k == 1 && !runtime_config.sample_func && !runtime_config.accept_token) { GenerateSpecV4(config, runtime_config, weights, activations, qbatch, env, timing_info); return; } HWY_WARN( - "use_mtp requires MTP weights, a single query and top_k == 1 (greedy);" + "use_mtp requires MTP weights and a single query;" " falling back to normal decoding."); } diff --git a/gemma/gemma_args.h b/gemma/gemma_args.h index 785f0b78..8ba8f0d7 100644 --- a/gemma/gemma_args.h +++ b/gemma/gemma_args.h @@ -185,6 +185,8 @@ struct RuntimeConfig { // prediction block. Requires num_mtp_layers > 0 weights, a single query and // top_k == 1; otherwise falls back to normal decoding with a warning. bool use_mtp = false; + size_t mtp_draft_horizon = 7; + float mtp_confidence_threshold = 0.0f; }; struct InferenceArgs : public ArgsBase { @@ -217,6 +219,9 @@ struct InferenceArgs : public ArgsBase { std::string eot_line; std::string attention_impl; std::string kv_cache_type; + bool use_mtp; + size_t mtp_draft_horizon; + float mtp_confidence_threshold; template void ForEach(const Visitor& visitor) { @@ -280,6 +285,13 @@ struct InferenceArgs : public ArgsBase { "KV cache data type (f32, bf16, int8). If empty, deduced from " "attention_impl.", 2); + visitor(use_mtp, "use_mtp", false, + "Enable DeepSeek MTP self-speculative decoding (default: false)", 2); + visitor(mtp_draft_horizon, "mtp_draft_horizon", size_t{7}, + "Number of MTP draft tokens to speculate per step (default: 7)", 2); + visitor(mtp_confidence_threshold, "mtp_confidence_threshold", 0.0f, + "Minimum top-1 probability threshold to accept MTP drafts (default: 0.0)", + 2); } void CopyTo(RuntimeConfig& runtime_config) const { @@ -313,6 +325,9 @@ struct InferenceArgs : public ArgsBase { HWY_ABORT("Unknown kv_cache_type: %s\n", kv_cache_type.c_str()); } } + runtime_config.use_mtp = use_mtp; + runtime_config.mtp_draft_horizon = mtp_draft_horizon; + runtime_config.mtp_confidence_threshold = mtp_confidence_threshold; } }; diff --git a/gemma/kv_cache.cc b/gemma/kv_cache.cc index baa3df31..90087a9c 100644 --- a/gemma/kv_cache.cc +++ b/gemma/kv_cache.cc @@ -112,15 +112,17 @@ static void InitDSState(const ModelConfig& config, const Allocator& allocator, // The MTP block is dense (no compressor state), but give it an offset entry // so `ds_state_offsets[num_layers]` is valid. if (config.num_mtp_layers > 0) { - ds_state_offsets.push_back(static_cast(accum)); + for (size_t i = 0; i < config.num_mtp_layers; ++i) { + ds_state_offsets.push_back(static_cast(accum)); + } } if (accum == 0) return; ds_state = MatStorageT("ds_state", Extents2D(1, accum), allocator, MatPadding::kPacked); ZeroInit(ds_state); - // Boundary snapshot for speculative decoding: state after the committed - // token of a verify step, restored if the draft is rejected. - ds_state_snapshot = MatStorageT("ds_snap", Extents2D(1, accum), + // Boundary snapshot for speculative decoding: state after each verified + // token of a verify step, restored if a draft is rejected. + ds_state_snapshot = MatStorageT("ds_snap", Extents2D(32, accum), allocator, MatPadding::kPacked); ZeroInit(ds_state_snapshot); } @@ -186,7 +188,11 @@ KVCache::KVCache(const ModelConfig& config, const InferenceArgs& inference_args, tiled_seq_len = num_tiles * kTileSize; // Trailing segment for the MTP block (indexed as layer `num_layers`). if (config.num_mtp_layers > 0) { - layer_flat_offsets.push_back(static_cast(flat_accum)); + const size_t mtp_size = config.MTPLayerConfig().CacheLayerSize(); + for (size_t i = 0; i < config.num_mtp_layers; ++i) { + layer_flat_offsets.push_back(static_cast(flat_accum)); + flat_accum += mtp_size; + } } InitDSState(config, allocator, ds_state, ds_state_snapshot, ds_state_offsets); } @@ -430,7 +436,11 @@ KVCache::KVCache(const ModelConfig& config, const InferenceArgs& inference_args, allocator, MatPadding::kOdd); } if (config.num_mtp_layers > 0) { - layer_flat_offsets.push_back(static_cast(flat_accum)); + const size_t mtp_size = config.MTPLayerConfig().CacheLayerSize(); + for (size_t i = 0; i < config.num_mtp_layers; ++i) { + layer_flat_offsets.push_back(static_cast(flat_accum)); + flat_accum += mtp_size; + } } InitDSState(config, allocator, ds_state, ds_state_snapshot, ds_state_offsets); } diff --git a/gemma/model_store.cc b/gemma/model_store.cc index 7d541b4d..7d49b451 100644 --- a/gemma/model_store.cc +++ b/gemma/model_store.cc @@ -333,6 +333,13 @@ static ModelConfig ReadOrDeduceConfig(BlobReader& reader, (std::string(ModelPrefix(*deduced_model)) + suffix).c_str(), (std::string(ModelPrefix(config.model)) + suffix).c_str()); } + for (const std::string& key : reader.Keys()) { + if (key.find("mtp_main_proj") != std::string::npos || + key.find("mtp.0.main_proj") != std::string::npos) { + config.num_mtp_layers = 3; + break; + } + } return config; } diff --git a/gemma/tensor_info.cc b/gemma/tensor_info.cc index 6d2a98cb..bcd97f81 100644 --- a/gemma/tensor_info.cc +++ b/gemma/tensor_info.cc @@ -1014,10 +1014,13 @@ TensorInfoRegistry::TensorInfoRegistry(const ModelConfig& config) { } } if (config.num_mtp_layers > 0) { - // The MTP block is a full extra layer registered at index `num_layers` - // (suffix `_43` for DeepSeek-V4-Flash), outside the main stack. - AddLayerTensors(config, config.MTPLayerConfig(), - config.layer_configs.size()); + // The MTP block is one or more extra layers registered starting at index + // `num_layers` (e.g. suffix `_43`, `_44`, `_45` for DSpark), outside the + // main stack. + for (size_t i = 0; i < config.num_mtp_layers; ++i) { + AddLayerTensors(config, config.MTPLayerConfig(), + config.layer_configs.size() + i); + } } for (size_t i = 0; i < config.vit_config.layer_configs.size(); ++i) { AddImageLayerTensors(config, config.vit_config.layer_configs[i], i); diff --git a/gemma/weights.cc b/gemma/weights.cc index 0ad8f473..2cc9143c 100644 --- a/gemma/weights.cc +++ b/gemma/weights.cc @@ -646,7 +646,7 @@ static void HWY_MAYBE_UNUSED SplitW1NUQ(const LayerConfig& layer_config) { // Zero-initializes only the allocated tensors in `*this`. void WeightsPtrs::ZeroInit() { ForEachTensor(nullptr, nullptr, [](const TensorArgs& t) { - if (!t.mat.HasPtr()) return; + if (!t.mat.HasPtr() || t.mat.GetType() == Type::kUnknown) return; gcpp::ZeroInit(t.mat); }); } @@ -1010,15 +1010,16 @@ WeightsPtrs::Mode WeightsPtrs::ReadFromBlobs(const ModelStore& model, std::vector tensors; // Enumerate all weights (negligible cost). - ForEachTensor(nullptr, nullptr, [&](const TensorArgs& t) { - const bool is_compressed = t.mat.GetType() == Type::kNUQ || - t.mat.GetType() == Type::kI8 || - t.mat.GetType() == Type::kQ4_0; - const MatPadding padding = (is_compressed || (t.flags & TensorArgs::kPacked)) - ? MatPadding::kPacked - : MatPadding::kOdd; + ForEachTensor(nullptr, nullptr, [&](const TensorArgs& t) HWY_ATTR { size_t key_idx; if (model.FindAndUpdateMatPtr(t.mat, key_idx)) { + const bool is_compressed = t.mat.GetType() == Type::kNUQ || + t.mat.GetType() == Type::kI8 || + t.mat.GetType() == Type::kQ4_0; + const MatPadding padding = + (is_compressed || (t.flags & TensorArgs::kPacked)) + ? MatPadding::kPacked + : MatPadding::kOdd; tensors.push_back( {.mat = &t.mat, .range = reader.Range(key_idx), .padding = padding}); return; diff --git a/gemma/weights.h b/gemma/weights.h index 4b70b0a1..b1f26f07 100644 --- a/gemma/weights.h +++ b/gemma/weights.h @@ -450,7 +450,7 @@ struct LayerWeightsPtrs { // Zero-initializes all allocated tensors in the layer. void ZeroInit() { ForEachTensor(nullptr, nullptr, [](const TensorArgs& t) { - if (!t.mat.HasPtr()) return; + if (!t.mat.HasPtr() || t.mat.GetType() == Type::kUnknown) return; gcpp::ZeroInit(t.mat); }); } @@ -635,6 +635,11 @@ struct WeightsPtrs { mtp_hc_fn(finder_("mtp_hc_fn")), mtp_hc_base(finder_("mtp_hc_base")), mtp_hc_scale(finder_("mtp_hc_scale")), + mtp_main_proj(finder_("mtp_main_proj")), + mtp_main_norm(finder_("mtp_main_norm")), + mtp_markov_w1(finder_("mtp_markov_w1")), + mtp_markov_w2(finder_("mtp_markov_w2")), + mtp_conf_proj(finder_("mtp_conf_proj")), t5gemma_encoder_embedding(finder_("enc_embedding")), t5gemma_decoder_embedding(finder_("dec_embedding")), t5gemma_encoder_final_norm_scale(finder_("enc_final_norm")), @@ -674,11 +679,13 @@ struct WeightsPtrs { vit_layers.emplace_back(idx, layer_config, tensors_); } if (config_.num_mtp_layers > 0) { - // The multi-token-prediction block: an extra layer at index + // The multi-token-prediction block: extra layer(s) starting at index // `num_layers`, outside the main stack (see `MTPLayerConfig`). mtp_layer_config = config_.MTPLayerConfig(); - mtp_layers.emplace_back(config_.layer_configs.size(), mtp_layer_config, - tensors_); + for (size_t idx = 0; idx < config_.num_mtp_layers; ++idx) { + mtp_layers.emplace_back(config_.layer_configs.size() + idx, + mtp_layer_config, tensors_); + } } } @@ -709,6 +716,13 @@ struct WeightsPtrs { MatPtr mtp_hc_base; // [hc_mult] f32 MatPtr mtp_hc_scale; // [1] f32 + // DSpark multi-layer target feature fusion and Markov refinement head. + MatPtr mtp_main_proj; // [model_dim, 3 * model_dim] + MatPtr mtp_main_norm; // [model_dim] + MatPtr mtp_markov_w1; // [vocab_size, markov_rank] + MatPtr mtp_markov_w2; // [vocab_size, markov_rank] + MatPtr mtp_conf_proj; // [1, model_dim + markov_rank] + // T5Gemma text encoder-decoder parts. MatPtr t5gemma_encoder_embedding; // at least BF16. MatPtr t5gemma_decoder_embedding; // at least BF16. @@ -796,15 +810,45 @@ struct WeightsPtrs { func(TENSOR_ARGS(hc_head_scale, kMustRead)); } if (config_.num_mtp_layers > 0) { - func(TENSOR_ARGS(mtp_e_proj, kMustRead)); - func(TENSOR_ARGS(mtp_h_proj, kMustRead)); - func(TENSOR_ARGS(mtp_enorm, kMustRead)); - func(TENSOR_ARGS(mtp_hnorm, kMustRead)); - func(TENSOR_ARGS(mtp_norm, kMustRead)); - if (config_.hc_mult > 1) { - func(TENSOR_ARGS(mtp_hc_fn, kMustRead)); - func(TENSOR_ARGS(mtp_hc_base, kMustRead)); - func(TENSOR_ARGS(mtp_hc_scale, kMustRead)); + if (mtp_main_proj.HasPtr() && mtp_main_proj.GetType() != Type::kUnknown) { + func(TENSOR_ARGS(mtp_main_proj, kMustRead)); + func(TENSOR_ARGS(mtp_main_norm, kMustRead)); + func(TENSOR_ARGS(mtp_norm, kMustRead)); + func(TENSOR_ARGS(mtp_markov_w1, kMustRead)); + func(TENSOR_ARGS(mtp_markov_w2, kMustRead)); + func(TENSOR_ARGS(mtp_conf_proj, kMustRead)); + if (config_.hc_mult > 1) { + func(TENSOR_ARGS(mtp_hc_fn, kMustRead)); + func(TENSOR_ARGS(mtp_hc_base, kMustRead)); + func(TENSOR_ARGS(mtp_hc_scale, kMustRead)); + } + } else if (mtp_e_proj.HasPtr() && mtp_e_proj.GetType() != Type::kUnknown) { + func(TENSOR_ARGS(mtp_e_proj, kMustRead)); + func(TENSOR_ARGS(mtp_h_proj, kMustRead)); + func(TENSOR_ARGS(mtp_enorm, kMustRead)); + func(TENSOR_ARGS(mtp_hnorm, kMustRead)); + func(TENSOR_ARGS(mtp_norm, kMustRead)); + if (config_.hc_mult > 1) { + func(TENSOR_ARGS(mtp_hc_fn, kMustRead)); + func(TENSOR_ARGS(mtp_hc_base, kMustRead)); + func(TENSOR_ARGS(mtp_hc_scale, kMustRead)); + } + } else { + func(TENSOR_ARGS(mtp_main_proj, kMaybeRead)); + func(TENSOR_ARGS(mtp_main_norm, kMaybeRead)); + func(TENSOR_ARGS(mtp_markov_w1, kMaybeRead)); + func(TENSOR_ARGS(mtp_markov_w2, kMaybeRead)); + func(TENSOR_ARGS(mtp_conf_proj, kMaybeRead)); + func(TENSOR_ARGS(mtp_e_proj, kMaybeRead)); + func(TENSOR_ARGS(mtp_h_proj, kMaybeRead)); + func(TENSOR_ARGS(mtp_enorm, kMaybeRead)); + func(TENSOR_ARGS(mtp_hnorm, kMaybeRead)); + func(TENSOR_ARGS(mtp_norm, kMustRead)); + if (config_.hc_mult > 1) { + func(TENSOR_ARGS(mtp_hc_fn, kMustRead)); + func(TENSOR_ARGS(mtp_hc_base, kMustRead)); + func(TENSOR_ARGS(mtp_hc_scale, kMustRead)); + } } } diff --git a/util/mat.h b/util/mat.h index b4d8870a..481d1d24 100644 --- a/util/mat.h +++ b/util/mat.h @@ -444,7 +444,8 @@ decltype(auto) CallUpcasted(const MatPtr* base, const Func& func, const MatPtrT mat(*base); return func(&mat, std::forward(args)...); } else { - HWY_ABORT("Unhandled type %s.", TypeName(base->GetType())); + HWY_ABORT("Unhandled type %s for tensor %s.", TypeName(base->GetType()), + base->Name()); } } @@ -483,7 +484,8 @@ decltype(auto) CallUpcastedSame(const MatPtr* base1, const MatPtr* base2, const MatPtrT mat2(*base2); return func(&mat1, &mat2, std::forward(args)...); } else { - HWY_ABORT("Unhandled type %s.", TypeName(base1->GetType())); + HWY_ABORT("Unhandled type %s for tensors %s and %s.", + TypeName(base1->GetType()), base1->Name(), base2->Name()); } }