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
13 changes: 9 additions & 4 deletions server/src/common/layer_split_backend.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -45,7 +45,8 @@ bool LayerSplitBackend::unpark(const std::string & what) {
GenerateResult LayerSplitBackend::run_from_state(const GenerateRequest & req,
const DaemonIO & io,
int base_pos,
bool reset_state) {
bool reset_state,
const std::vector<int32_t> & history_prefix) {
GenerateResult result;
if (!adapter_) {
result.fail(GenerateErrorCode::AdapterUnavailable);
Expand Down Expand Up @@ -115,6 +116,7 @@ GenerateResult LayerSplitBackend::run_from_state(const GenerateRequest & req,
? adapter_->decode_dflash(req.prompt, base_pos, last_tok, req.n_gen,
result.tokens, out_io, dflash_accept_rate)
: adapter_->decode_ar(last_tok, base_pos + (int)req.prompt.size(), req.n_gen,
history_prefix,
result.tokens, out_io);
if (use_dflash) result.accept_rate = dflash_accept_rate;
if (!ok) {
Expand All @@ -131,7 +133,8 @@ GenerateResult LayerSplitBackend::run_from_state(const GenerateRequest & req,

GenerateResult LayerSplitBackend::generate_impl(const GenerateRequest & req,
const DaemonIO & io) {
return run_from_state(req, io, /*base_pos=*/0, /*reset_state=*/true);
return run_from_state(req, io, /*base_pos=*/0, /*reset_state=*/true,
req.prompt);
}

bool LayerSplitBackend::snapshot_save(int slot) {
Expand Down Expand Up @@ -177,12 +180,14 @@ GenerateResult LayerSplitBackend::restore_and_generate_impl(
std::fprintf(stderr,
"[pc] snapshot longer than prompt (snap=%d > prompt=%zu) — "
"fresh prefill fallback\n", snap_pos, req.prompt.size());
return run_from_state(req, io, /*base_pos=*/0, /*reset_state=*/true);
return run_from_state(req, io, /*base_pos=*/0, /*reset_state=*/true,
req.prompt);
}
GenerateRequest delta_req = req;
delta_req.prompt = std::vector<int32_t>(
req.prompt.begin() + snap_pos, req.prompt.end());
return run_from_state(delta_req, io, snap_pos, /*reset_state=*/false);
return run_from_state(delta_req, io, snap_pos, /*reset_state=*/false,
req.prompt);
}

ModelBackend::CompressResult
Expand Down
6 changes: 5 additions & 1 deletion server/src/common/layer_split_backend.h
Original file line number Diff line number Diff line change
Expand Up @@ -28,7 +28,10 @@ class LayerSplitAdapter {
virtual int prefill_chunk_tokens() const { return 0; }
virtual bool prefill(const std::vector<int32_t> & prompt,
int base_pos, int & last_tok) = 0;
// history_prefix is the full original request prompt (not the delta
// prefill after a prefix-cache restore); it seeds sampler penalty history.
virtual bool decode_ar(int last_tok, int committed, int n_gen,
const std::vector<int32_t> & history_prefix,
std::vector<int32_t> & out_tokens,
const DaemonIO & io) = 0;
virtual bool supports_cpu_sampling() const { return false; }
Expand Down Expand Up @@ -125,7 +128,8 @@ class LayerSplitBackend : public ModelBackend {
GenerateResult run_from_state(const GenerateRequest & req,
const DaemonIO & io,
int base_pos,
bool reset_state);
bool reset_state,
const std::vector<int32_t> & history_prefix);

std::unique_ptr<LayerSplitAdapter> adapter_;
bool shutdown_done_ = false;
Expand Down
11 changes: 9 additions & 2 deletions server/src/common/layer_split_runtime.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -10,19 +10,25 @@ bool run_layer_split_ar_decode(
const std::vector<float> & prefill_last_logits,
const SamplerCfg & sampler,
std::mt19937_64 & rng,
const std::vector<int32_t> & history_prefix,
const LayerSplitForwardStep & forward_one,
const std::function<bool(int)> & is_eos,
std::vector<int32_t> & out_tokens,
const DaemonIO & io) {
if (n_gen <= 0) return true;

std::vector<int32_t> history;
if (sampler.needs_logit_processing()) {
history.reserve(history_prefix.size() + out_tokens.size() + (size_t)n_gen);
history.insert(history.end(), history_prefix.begin(), history_prefix.end());
history.insert(history.end(), out_tokens.begin(), out_tokens.end());
if ((int)prefill_last_logits.size() != vocab) return false;
last_tok = sample_logits(prefill_last_logits.data(), vocab, sampler,
out_tokens, rng);
history, rng);
}

out_tokens.push_back(last_tok);
if (sampler.needs_logit_processing()) history.push_back(last_tok);
io.emit(last_tok);
if (io.cancelled) {
io.emit(-1);
Expand All @@ -46,11 +52,12 @@ bool run_layer_split_ar_decode(
if (sampler.needs_logit_processing()) {
if ((int)logits_buf.size() != vocab) return false;
next_tok = sample_logits(logits_buf.data(), vocab, sampler,
out_tokens, rng);
history, rng);
}

last_tok = next_tok;
out_tokens.push_back(last_tok);
if (sampler.needs_logit_processing()) history.push_back(last_tok);
io.emit(last_tok);
++committed;
if (io.cancelled) break;
Expand Down
1 change: 1 addition & 0 deletions server/src/common/layer_split_runtime.h
Original file line number Diff line number Diff line change
Expand Up @@ -87,6 +87,7 @@ bool run_layer_split_ar_decode(
const std::vector<float> & prefill_last_logits,
const SamplerCfg & sampler,
std::mt19937_64 & rng,
const std::vector<int32_t> & history_prefix,
const LayerSplitForwardStep & forward_one,
const std::function<bool(int)> & is_eos,
std::vector<int32_t> & out_tokens,
Expand Down
7 changes: 6 additions & 1 deletion server/src/deepseek4/deepseek4_layer_split_adapter.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -405,7 +405,10 @@ bool DeepSeek4LayerSplitAdapter::init_mixed_target_split() {
}

void DeepSeek4LayerSplitAdapter::begin_request(const GenerateRequest & req) {
(void)req;
sampler_ = req.sampler;
if (req.do_sample && sampler_.seed != 0) {
sampler_rng_.seed(sampler_.seed);
}
}

void DeepSeek4LayerSplitAdapter::reset_request_state() {
Expand Down Expand Up @@ -598,6 +601,7 @@ bool DeepSeek4LayerSplitAdapter::decode_ar(
int last_tok_in,
int committed,
int n_gen,
const std::vector<int32_t> & history_prefix,
std::vector<int32_t> & out_tokens,
const DaemonIO & io) {
if (shards_.empty()) return false;
Expand All @@ -619,6 +623,7 @@ bool DeepSeek4LayerSplitAdapter::decode_ar(
return run_layer_split_ar_decode(
last_tok_in, committed, n_gen, vocab,
prefill_last_logits_, sampler_, sampler_rng_,
history_prefix,
forward_one, is_eos, out_tokens, io);
}

Expand Down
1 change: 1 addition & 0 deletions server/src/deepseek4/deepseek4_layer_split_adapter.h
Original file line number Diff line number Diff line change
Expand Up @@ -53,6 +53,7 @@ class DeepSeek4LayerSplitAdapter : public LayerSplitAdapter {
bool prefill(const std::vector<int32_t> & prompt,
int base_pos, int & last_tok) override;
bool decode_ar(int last_tok, int committed, int n_gen,
const std::vector<int32_t> & history_prefix,
std::vector<int32_t> & out_tokens,
const DaemonIO & io) override;
bool supports_cpu_sampling() const override { return true; }
Expand Down
3 changes: 2 additions & 1 deletion server/src/gemma4/gemma4_layer_split_adapter.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -1040,6 +1040,7 @@ bool Gemma4LayerSplitAdapter::decode_ar(
int last_tok,
int committed,
int n_gen,
const std::vector<int32_t> & history_prefix,
std::vector<int32_t> & out_tokens,
const DaemonIO & io) {
if (n_gen <= 0) return true;
Expand All @@ -1049,7 +1050,7 @@ bool Gemma4LayerSplitAdapter::decode_ar(
const int vocab = w.n_vocab;
const bool ok = run_layer_split_ar_decode(
last_tok, committed, n_gen, vocab, prefill_last_logits_, sampler_,
sampler_rng_,
sampler_rng_, history_prefix,
[&](const std::vector<int32_t> & one, int pos, int & next_tok,
std::vector<float> * logits_out) {
return use_mixed_target_split()
Expand Down
1 change: 1 addition & 0 deletions server/src/gemma4/gemma4_layer_split_adapter.h
Original file line number Diff line number Diff line change
Expand Up @@ -60,6 +60,7 @@ class Gemma4LayerSplitAdapter : public LayerSplitAdapter {
bool prefill(const std::vector<int32_t> & prompt,
int base_pos, int & last_tok) override;
bool decode_ar(int last_tok, int committed, int n_gen,
const std::vector<int32_t> & history_prefix,
std::vector<int32_t> & out_tokens,
const DaemonIO & io) override;
bool supports_cpu_sampling() const override { return true; }
Expand Down
3 changes: 2 additions & 1 deletion server/src/laguna/laguna_layer_split_adapter.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -855,6 +855,7 @@ bool LagunaLayerSplitAdapter::decode_ar(
int last_tok,
int committed,
int n_gen,
const std::vector<int32_t> & history_prefix,
std::vector<int32_t> & out_tokens,
const DaemonIO & io) {
if (n_gen <= 0) return true;
Expand All @@ -864,7 +865,7 @@ bool LagunaLayerSplitAdapter::decode_ar(
const int vocab = (int)w.embedder.n_vocab;
const bool ok = run_layer_split_ar_decode(
last_tok, committed, n_gen, vocab, prefill_last_logits_, sampler_,
sampler_rng_,
sampler_rng_, history_prefix,
[&](const std::vector<int32_t> & one, int pos, int & next_tok,
std::vector<float> * logits_out) {
return use_mixed_target_split()
Expand Down
1 change: 1 addition & 0 deletions server/src/laguna/laguna_layer_split_adapter.h
Original file line number Diff line number Diff line change
Expand Up @@ -59,6 +59,7 @@ class LagunaLayerSplitAdapter : public LayerSplitAdapter {
bool prefill(const std::vector<int32_t> & prompt,
int base_pos, int & last_tok) override;
bool decode_ar(int last_tok, int committed, int n_gen,
const std::vector<int32_t> & history_prefix,
std::vector<int32_t> & out_tokens,
const DaemonIO & io) override;
bool supports_cpu_sampling() const override { return true; }
Expand Down
3 changes: 2 additions & 1 deletion server/src/qwen35/qwen35_layer_split_adapter.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -1270,14 +1270,15 @@ int Qwen35LayerSplitAdapter::current_last_token() const {

bool Qwen35LayerSplitAdapter::decode_ar(
int last_tok, int committed, int n_gen,
const std::vector<int32_t> & history_prefix,
std::vector<int32_t> & out_tokens,
const DaemonIO & io) {
if (n_gen <= 0) return true;
const auto & w = shards_.front().weights;
const int vocab = w.n_vocab;
const bool ok = run_layer_split_ar_decode(
last_tok, committed, n_gen, vocab, prefill_last_logits_, sampler_,
sampler_rng_,
sampler_rng_, history_prefix,
[&](const std::vector<int32_t> & one, int pos, int & next_tok,
std::vector<float> * logits_out) {
if (use_mixed_target_split()) {
Expand Down
1 change: 1 addition & 0 deletions server/src/qwen35/qwen35_layer_split_adapter.h
Original file line number Diff line number Diff line change
Expand Up @@ -61,6 +61,7 @@ class Qwen35LayerSplitAdapter : public LayerSplitAdapter {
bool prefill(const std::vector<int32_t> & prompt,
int base_pos, int & last_tok) override;
bool decode_ar(int last_tok, int committed, int n_gen,
const std::vector<int32_t> & history_prefix,
std::vector<int32_t> & out_tokens,
const DaemonIO & io) override;
bool supports_cpu_sampling() const override { return true; }
Expand Down
141 changes: 140 additions & 1 deletion server/tests/test_deepseek4_unit.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,8 @@

#include "common/backend_ipc.h"
#include "common/dspark_head.h"
#include "common/layer_split_backend.h"
#include "common/layer_split_runtime.h"
#include "common/layer_split_utils.h"
#include "deepseek4/deepseek4_dspark.h"

Expand Down Expand Up @@ -899,6 +901,139 @@ static void test_hc_state_dimensions() {
std::fprintf(stderr, g_failures ? " done\n" : " ok\n");
}

static void test_layer_split_request_propagates_sampler() {
std::fprintf(stderr, " test_layer_split_request_propagates_sampler ...");

DeepSeek4LayerSplitAdapter adapter({});
GenerateRequest req;
req.do_sample = true;
req.sampler.temp = 0.25f;
req.sampler.top_p = 0.9f;
req.sampler.seed = 42;

adapter.begin_request(req);

TEST_ASSERT(adapter.sampler_.temp == req.sampler.temp);
TEST_ASSERT(adapter.sampler_.top_p == req.sampler.top_p);
TEST_ASSERT(adapter.sampler_.seed == req.sampler.seed);
std::mt19937_64 expected(req.sampler.seed);
TEST_ASSERT(adapter.sampler_rng_() == expected());

std::fprintf(stderr, g_failures ? " done\n" : " ok\n");
}

static void test_layer_split_sampler_uses_prompt_history() {
std::fprintf(stderr, " test_layer_split_sampler_uses_prompt_history ...");

SamplerCfg sampler;
sampler.rep_pen = 2.0f;
const std::vector<float> logits = {0.0f, 4.0f, 3.0f};
const std::vector<int32_t> prompt_history = {1};
std::vector<int32_t> out_tokens;
std::mt19937_64 rng(42);
const bool ok = run_layer_split_ar_decode(
/*last_tok=*/1, /*committed=*/1, /*n_gen=*/1, /*vocab=*/3,
logits, sampler, rng, prompt_history,
[](const std::vector<int32_t> &, int, int &,
std::vector<float> *) { return false; },
[](int) { return false; }, out_tokens, DaemonIO{});

TEST_ASSERT(ok);
TEST_ASSERT(out_tokens.size() == 1);
if (out_tokens.size() == 1) {
TEST_ASSERT(out_tokens[0] == 2);
}

std::fprintf(stderr, g_failures ? " done\n" : " ok\n");
}

static void test_layer_split_sampler_appends_generated_tokens() {
std::fprintf(stderr,
" test_layer_split_sampler_appends_generated_tokens ...");

SamplerCfg sampler;
sampler.rep_pen = 2.0f;
const std::vector<float> logits = {0.0f, 4.0f, 3.0f};
std::vector<int32_t> out_tokens;
std::mt19937_64 rng(42);
const bool ok = run_layer_split_ar_decode(
/*last_tok=*/0, /*committed=*/0, /*n_gen=*/2, /*vocab=*/3,
logits, sampler, rng, /*history_prefix=*/{},
[&logits](const std::vector<int32_t> &, int, int &,
std::vector<float> * logits_out) {
if (logits_out) *logits_out = logits;
return true;
},
[](int) { return false; }, out_tokens, DaemonIO{});

TEST_ASSERT(ok);
// With no prompt history the first sample is the plain argmax (token 1);
// the second sees token 1 in history, so rep_pen drops its logit to 2 and
// token 2 (logit 3) wins.
TEST_ASSERT(out_tokens == std::vector<int32_t>({1, 2}));

std::fprintf(stderr, g_failures ? " done\n" : " ok\n");
}

class TestLayerSplitHistoryAdapter final : public LayerSplitAdapter {
public:
const char * name() const override { return "test-history"; }
bool init() override { return true; }
int max_context() const override { return 64; }
void reset_request_state() override { reset_called = true; }
bool prefill(const std::vector<int32_t> & prompt,
int, int & last_tok) override {
prefilled.insert(prefilled.end(), prompt.begin(), prompt.end());
last_tok = 2;
return true;
}
bool decode_ar(int, int, int,
const std::vector<int32_t> & history_prefix,
std::vector<int32_t> & out_tokens,
const DaemonIO &) override {
decoded_history = history_prefix;
out_tokens.push_back(2);
return true;
}
bool supports_cpu_sampling() const override { return true; }
void free_drafter() override {}
int snapshot_cur_pos(int) const override { return 1; }
bool snapshot_restore(int) override {
restore_called = true;
return true;
}
int current_last_token() const override { return 1; }
void shutdown() override {}

bool reset_called = false;
bool restore_called = false;
std::vector<int32_t> prefilled;
std::vector<int32_t> decoded_history;
};

static void test_layer_split_restore_preserves_full_sampling_history() {
std::fprintf(stderr,
" test_layer_split_restore_preserves_full_sampling_history ...");

auto adapter = std::make_unique<TestLayerSplitHistoryAdapter>();
auto * observed = adapter.get();
LayerSplitBackend backend(std::move(adapter));
GenerateRequest req;
req.prompt = {1, 2, 3};
req.n_gen = 1;

const GenerateResult result =
backend.restore_and_generate_impl(0, req, DaemonIO{});

TEST_ASSERT(result.ok());
TEST_ASSERT(observed->restore_called);
TEST_ASSERT(!observed->reset_called);
TEST_ASSERT(observed->prefilled == std::vector<int32_t>({2, 3}));
TEST_ASSERT(observed->decoded_history == req.prompt);

std::fprintf(stderr, g_failures ? " done\n" : " ok\n");
}

static void test_loader_rejects_missing_required_metadata(ggml_backend_t backend) {
std::fprintf(stderr, " test_loader_rejects_missing_required_metadata ...");

Expand Down Expand Up @@ -1364,7 +1499,7 @@ static void test_adapter_guard_paths() {
adapter.remote_target_shard_.active_ = false;

std::vector<int32_t> out_tokens;
TEST_ASSERT(!adapter.decode_ar(1, 0, 1, out_tokens, DaemonIO{}));
TEST_ASSERT(!adapter.decode_ar(1, 0, 1, {}, out_tokens, DaemonIO{}));

std::fprintf(stderr, g_failures ? " done\n" : " ok\n");
}
Expand Down Expand Up @@ -2049,6 +2184,10 @@ int main() {
test_auto_split_computation();
test_layer_range_validation();
test_hc_state_dimensions();
test_layer_split_request_propagates_sampler();
test_layer_split_sampler_uses_prompt_history();
test_layer_split_sampler_appends_generated_tokens();
test_layer_split_restore_preserves_full_sampling_history();
test_loader_rejects_missing_required_metadata(backend);
test_loader_rejects_invalid_compress_ratio_type(backend);
test_loader_rejects_zero_vocab_size(backend);
Expand Down
Loading