Skip to content
Merged
Show file tree
Hide file tree
Changes from 3 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 4 additions & 4 deletions server/src/deepseek4/deepseek4_backend.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -831,14 +831,14 @@ int DeepSeek4Backend::do_prefill(const std::vector<int32_t> & tokens,
}
bool ok = false;
if (moe_hybrid_) {
ok = deepseek4_step(backend_, w_, cache_, embed.data(), n_tok, pos,
ok = deepseek4_step(backend_, cfg_.device.gpu, w_, cache_, embed.data(), n_tok, pos,
logits, moe_hybrid_.get(), tokens.data() + i,
&stream_engine_, timing ? &step_tel : nullptr,
routing_stats_.get(), hp);
} else {
std::vector<float> hc_state;
ok = deepseek4_step_layer_range(
backend_, w_, cache_, hc_state, embed.data(), n_tok, pos,
backend_, cfg_.device.gpu, w_, cache_, hc_state, embed.data(), n_tok, pos,
0, w_.n_layer, &logits, tokens.data() + i,
timing ? &step_tel : nullptr,
cfg_.prefill_mode != PrefillAttentionMode::Sparse, hp);
Expand Down Expand Up @@ -919,7 +919,7 @@ bool DeepSeek4Backend::do_decode(int committed, int n_gen,
if (timing) step_tel.embed_us = elapsed_us(embed_t0, Clock::now());

const int pos = std::max(0, committed + generated - 1);
if (!deepseek4_step(backend_, w_, cache_, embed.data(), 1,
if (!deepseek4_step(backend_, cfg_.device.gpu, w_, cache_, embed.data(), 1,
pos, logits,
moe_hybrid_.get(), &tok_to_eval,
moe_hybrid_ ? &stream_engine_ : nullptr,
Expand Down Expand Up @@ -1023,7 +1023,7 @@ GenerateResult DeepSeek4Backend::generate_impl(const GenerateRequest & req,
std::vector<int32_t> spec_toks;
spec_ran = true;
if (!run_deepseek4_dspark_spec_decode(
backend_, w_, cache_, *spec_drafter_, committed, seed,
backend_, cfg_.device.gpu, w_, cache_, *spec_drafter_, committed, seed,
req.n_gen - 1,
win_len > 0 ? spec_feat_window_.data() : nullptr, win_len,
spec_toks, &accept_rate,
Expand Down
2 changes: 2 additions & 0 deletions server/src/deepseek4/deepseek4_dspark.h
Original file line number Diff line number Diff line change
Expand Up @@ -118,6 +118,7 @@ bool deepseek4_dspark_draft_forward(ggml_backend_t backend,
// main_hidden feed.
// Advances/updates the target cache exactly like a decode of these tokens.
bool deepseek4_dspark_verify_forward(ggml_backend_t backend,
int device,
const DeepSeek4Weights & w,
DeepSeek4Cache & cache,
const std::vector<int> & capture_layer_ids,
Expand Down Expand Up @@ -165,6 +166,7 @@ void deepseek4_spec_rollback_apply(const DeepSeek4SpecRollback & rollback,
struct GenerateRequest; // fwd (from common/…); the loop only needs n_gen + committed
bool run_deepseek4_dspark_spec_decode(
ggml_backend_t backend,
int device,
const DeepSeek4Weights & target_w,
DeepSeek4Cache & target_cache,
const DSparkDrafter & drafter,
Expand Down
15 changes: 9 additions & 6 deletions server/src/deepseek4/deepseek4_dspark_spec.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -43,9 +43,9 @@ namespace dflash::common {
class DeepSeek4DFlashTarget : public DFlashTarget {
public:
DeepSeek4DFlashTarget(const DeepSeek4Weights & w, DeepSeek4Cache & cache,
ggml_backend_t backend, ggml_backend_t snap_backend,
ggml_backend_t backend, int device, ggml_backend_t snap_backend,
std::vector<int> capture_ids, int mask_tok)
: w_(w), cache_(cache), backend_(backend), snap_backend_(snap_backend),
: w_(w), cache_(cache), backend_(backend), device_(device), snap_backend_(snap_backend),
capture_ids_(std::move(capture_ids)), mask_tok_(mask_tok) {}

~DeepSeek4DFlashTarget() override { clear_snapshot(); }
Expand Down Expand Up @@ -80,7 +80,7 @@ class DeepSeek4DFlashTarget : public DFlashTarget {
std::vector<int32_t> am1;
std::vector<float> feat1;
std::vector<float> logits1;
if (!deepseek4_dspark_verify_forward(backend_, w_, cache_, capture_ids_,
if (!deepseek4_dspark_verify_forward(backend_, device_, w_, cache_, capture_ids_,
embed_buf_.data() + (size_t) t * w_.n_embd,
tokens.data() + t, 1, base_pos + t, am1,
keep_logits_ ? &logits1 : nullptr,
Expand All @@ -103,7 +103,7 @@ class DeepSeek4DFlashTarget : public DFlashTarget {
std::vector<int32_t> am;
// n==1 must take the dynamic (non-reuse) path: the reused decode graph
// skips the capture/all-logits hooks (backend HC), which this needs.
if (!deepseek4_dspark_verify_forward(backend_, w_, cache_, capture_ids_,
if (!deepseek4_dspark_verify_forward(backend_, device_, w_, cache_, capture_ids_,
embed_buf_.data(), tokens.data(), n, base_pos, am,
keep_logits_ ? &verify_logits_ : nullptr,
verify_features_, telemetry_,
Expand Down Expand Up @@ -195,6 +195,7 @@ class DeepSeek4DFlashTarget : public DFlashTarget {
const DeepSeek4Weights & w_;
DeepSeek4Cache & cache_;
ggml_backend_t backend_;
int device_;
ggml_backend_t snap_backend_;
std::vector<int> capture_ids_;
int mask_tok_;
Expand Down Expand Up @@ -359,6 +360,7 @@ void deepseek4_spec_rollback_apply(const DeepSeek4SpecRollback & rollback,
// touches the fused single-token 23 tok/s path, with the Ds4VerifyHooks that
// add per-layer mean-over-HC capture and full per-position logits.
bool deepseek4_dspark_verify_forward(ggml_backend_t backend,
int device,
const DeepSeek4Weights & w,
DeepSeek4Cache & cache,
const std::vector<int> & capture_layer_ids,
Expand All @@ -378,7 +380,7 @@ bool deepseek4_dspark_verify_forward(ggml_backend_t backend,
hooks.capture_layer_ids = &capture_layer_ids;
hooks.capture_out = &capture_out;
hooks.all_logits_out = &all_logits;
if (!deepseek4_step_layer_range(backend, w, cache, hc_state, embed, n_tokens, kv_start,
if (!deepseek4_step_layer_range(backend, device, w, cache, hc_state, embed, n_tokens, kv_start,
0, w.n_layer, &last_logits, token_ids,
telemetry, allow_graph_reuse,
&hooks)) {
Expand All @@ -404,6 +406,7 @@ bool deepseek4_dspark_verify_forward(ggml_backend_t backend,

bool run_deepseek4_dspark_spec_decode(
ggml_backend_t backend,
int device,
const DeepSeek4Weights & target_w,
DeepSeek4Cache & target_cache,
const DSparkDrafter & drafter,
Expand Down Expand Up @@ -475,7 +478,7 @@ bool run_deepseek4_dspark_spec_decode(
ggml_backend_t snap_backend = ggml_backend_cpu_init();
if (!snap_backend) { std::fprintf(stderr, "[ds4-spec] no CPU snapshot backend\n"); return false; }

DeepSeek4DFlashTarget target(target_w, target_cache, backend, snap_backend,
DeepSeek4DFlashTarget target(target_w, target_cache, backend, device, snap_backend,
drafter.capture_layer_ids, drafter.mask_token_id);
DraftWeights dw = make_dspark_shim(drafter);
DeepSeek4SpecRollback rollback;
Expand Down
Loading