Skip to content

Commit 6fd8d9e

Browse files
committed
feat(server): add target-layer-split backend adapter path
1 parent ad8662b commit 6fd8d9e

14 files changed

Lines changed: 1211 additions & 25 deletions

dflash/CMakeLists.txt

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -242,10 +242,12 @@ add_library(dflash_common STATIC
242242
src/common/dflash_draft_ipc_daemon.cpp
243243
src/common/dflash_draft_graph.cpp
244244
src/common/dflash_spec_decode.cpp
245+
src/common/layer_split_backend.cpp
245246
src/qwen35/graph_builders.cpp
246247
src/qwen35/layer_split_forward.cpp
247248
src/qwen35/layer_split_daemon.cpp
248249
src/qwen35/qwen35_backend.cpp
250+
src/qwen35/qwen35_layer_split_adapter.cpp
249251
src/qwen35/qwen35_dflash_target.cpp
250252
src/qwen35/qwen35_layer_split_dflash_target.cpp
251253
src/qwen35/layer_split_daemon_loop.cpp

dflash/src/common/backend_factory.cpp

Lines changed: 27 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -7,8 +7,11 @@
77
#include "laguna_backend.h"
88
#include "qwen3_backend.h"
99
#include "gemma4_backend.h"
10+
#include "layer_split_backend.h"
11+
#include "qwen35_layer_split_adapter.h"
1012

1113
#include <cstdio>
14+
#include <algorithm>
1215

1316
namespace dflash::common {
1417

@@ -37,6 +40,30 @@ std::unique_ptr<ModelBackend> create_backend(const BackendArgs & args) {
3740
std::fprintf(stderr, "[backend_factory] detected arch=%s\n", arch.c_str());
3841

3942
if (arch == "qwen35") {
43+
if (args.device.is_layer_split()) {
44+
Qwen35LayerSplitAdapterConfig cfg;
45+
cfg.target_path = args.model_path;
46+
cfg.draft_path = args.draft_path;
47+
cfg.device = args.device;
48+
cfg.draft_gpu = args.draft_device.gpu;
49+
cfg.remote_draft = args.remote_draft;
50+
cfg.fa_window = args.fa_window;
51+
cfg.kq_stride_pad = args.kq_stride_pad;
52+
cfg.draft_ctx_max = args.draft_ctx_max;
53+
cfg.max_verify_tokens = args.ddtree_mode
54+
? std::max<int>(DFLASH27B_DRAFT_BLOCK_SIZE, args.ddtree_budget + 1)
55+
: DFLASH27B_DRAFT_BLOCK_SIZE;
56+
cfg.run_dflash = args.draft_path != nullptr;
57+
58+
auto adapter = std::make_unique<Qwen35LayerSplitAdapter>(cfg);
59+
auto backend = std::make_unique<LayerSplitBackend>(std::move(adapter));
60+
if (!backend->init()) {
61+
std::fprintf(stderr, "[backend_factory] LayerSplitBackend(qwen35) init failed\n");
62+
return nullptr;
63+
}
64+
return backend;
65+
}
66+
4067
Qwen35Config cfg;
4168
cfg.target_path = args.model_path;
4269
cfg.draft_path = args.draft_path;

dflash/src/common/dflash_spec_decode.cpp

Lines changed: 33 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -34,7 +34,30 @@ bool run_dflash_spec_decode(
3434
int draft_ctx_max,
3535
int stream_fd,
3636
DFlashDraftIpcClient * remote_draft,
37-
const std::vector<int32_t> * hint_tokens) {
37+
const std::vector<int32_t> * hint_tokens,
38+
int base_pos) {
39+
DaemonIO io;
40+
io.stream_fd = stream_fd;
41+
return run_dflash_spec_decode(target, draft_weights, draft_backend,
42+
feature_ring, prompt, n_gen, last_tok,
43+
out_path, draft_ctx_max, io,
44+
remote_draft, hint_tokens, base_pos);
45+
}
46+
47+
bool run_dflash_spec_decode(
48+
DFlashTarget & target,
49+
DraftWeights & draft_weights,
50+
ggml_backend_t draft_backend,
51+
DraftFeatureMirror & feature_ring,
52+
const std::vector<int32_t> & prompt,
53+
int n_gen,
54+
int last_tok,
55+
const char * out_path,
56+
int draft_ctx_max,
57+
const DaemonIO & io,
58+
DFlashDraftIpcClient * remote_draft,
59+
const std::vector<int32_t> * hint_tokens,
60+
int base_pos) {
3861
const bool use_remote_draft = remote_draft && remote_draft->active();
3962
if (!use_remote_draft && !feature_ring.target_feat) return false;
4063

@@ -54,7 +77,7 @@ bool run_dflash_spec_decode(
5477
std::vector<float> remote_hidden; // host buffer for remote-draft hidden states
5578

5679
std::vector<int32_t> out_all = prompt;
57-
int committed = (int)prompt.size();
80+
int committed = base_pos + (int)prompt.size();
5881
int n_generated = 0;
5982
int n_draft_steps = 0;
6083
int n_accept_sum = 0;
@@ -199,15 +222,19 @@ bool run_dflash_spec_decode(
199222
last_tok = replay_last_tok;
200223

201224
bool hit_eos = false;
225+
int emitted = 0;
202226
for (int i = 0; i < commit_n; i++) {
203227
out_all.push_back(replay_tok[i]);
204-
stream_emit_fd(stream_fd, replay_tok[i]);
228+
io.emit(replay_tok[i]);
229+
if (io.cancelled) break;
230+
++emitted;
205231
if (target.is_eos(replay_tok[i])) hit_eos = true;
206232
}
207-
committed += commit_n;
208-
n_generated += commit_n;
209-
n_accept_sum += std::min(accept_n, commit_n);
233+
committed += emitted;
234+
n_generated += emitted;
235+
n_accept_sum += std::min(accept_n, emitted);
210236
n_draft_steps++;
237+
if (io.cancelled) break;
211238
if (hit_eos) break;
212239
}
213240
if (!use_remote_draft && draft_backend) ggml_backend_synchronize(draft_backend);
@@ -231,4 +258,3 @@ bool run_dflash_spec_decode(
231258
}
232259

233260
} // namespace dflash::common
234-

dflash/src/common/dflash_spec_decode.h

Lines changed: 18 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,7 @@
1414
#include "dflash_target.h"
1515
#include "dflash_feature_ring.h"
1616
#include "dflash_draft_ipc.h"
17+
#include "model_backend.h"
1718

1819
#include "ggml.h"
1920
#include "ggml-backend.h"
@@ -53,6 +54,22 @@ bool run_dflash_spec_decode(
5354
int draft_ctx_max,
5455
int stream_fd = -1,
5556
DFlashDraftIpcClient * remote_draft = nullptr,
56-
const std::vector<int32_t> * hint_tokens = nullptr);
57+
const std::vector<int32_t> * hint_tokens = nullptr,
58+
int base_pos = 0);
59+
60+
bool run_dflash_spec_decode(
61+
DFlashTarget & target,
62+
DraftWeights & draft_weights,
63+
ggml_backend_t draft_backend,
64+
DraftFeatureMirror & feature_ring,
65+
const std::vector<int32_t> & prompt,
66+
int n_gen,
67+
int last_tok,
68+
const char * out_path,
69+
int draft_ctx_max,
70+
const DaemonIO & io,
71+
DFlashDraftIpcClient * remote_draft = nullptr,
72+
const std::vector<int32_t> * hint_tokens = nullptr,
73+
int base_pos = 0);
5774

5875
} // namespace dflash::common
Lines changed: 225 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,225 @@
1+
// Generic server-facing backend for target layer split.
2+
3+
#include "layer_split_backend.h"
4+
5+
#include "io_utils.h"
6+
7+
#include <chrono>
8+
#include <cstdio>
9+
#include <utility>
10+
11+
namespace dflash::common {
12+
13+
LayerSplitBackend::LayerSplitBackend(std::unique_ptr<LayerSplitAdapter> adapter)
14+
: adapter_(std::move(adapter)) {}
15+
16+
LayerSplitBackend::~LayerSplitBackend() { shutdown(); }
17+
18+
bool LayerSplitBackend::init() {
19+
if (!adapter_) {
20+
std::fprintf(stderr, "[target-split] missing model adapter\n");
21+
return false;
22+
}
23+
return adapter_->init();
24+
}
25+
26+
void LayerSplitBackend::print_ready_banner() const {
27+
std::printf("[daemon] ready\n");
28+
std::fflush(stdout);
29+
}
30+
31+
bool LayerSplitBackend::park(const std::string & what) {
32+
std::fprintf(stderr, "[target-split] park is not supported yet (%s)\n",
33+
what.c_str());
34+
return false;
35+
}
36+
37+
bool LayerSplitBackend::unpark(const std::string & what) {
38+
std::fprintf(stderr, "[target-split] unpark is not supported yet (%s)\n",
39+
what.c_str());
40+
return false;
41+
}
42+
43+
GenerateResult LayerSplitBackend::run_from_state(const GenerateRequest & req,
44+
const DaemonIO & io,
45+
int base_pos,
46+
bool reset_state) {
47+
GenerateResult result;
48+
if (!adapter_) {
49+
result.error = "adapter";
50+
return result;
51+
}
52+
53+
DaemonIO out_io = io.with_token_callback(req.on_token);
54+
if (base_pos + (int)req.prompt.size() + req.n_gen + 1 > adapter_->max_context()) {
55+
result.error = "context";
56+
return result;
57+
}
58+
if (req.do_sample && req.sampler.temp > 0.0f) {
59+
result.error = "sampling_unsupported";
60+
return result;
61+
}
62+
63+
adapter_->begin_request(req);
64+
if (reset_state) adapter_->reset_request_state();
65+
66+
const int prompt_len = (int)req.prompt.size();
67+
int last_tok = (base_pos > 0 && prompt_len == 0)
68+
? adapter_->current_last_token()
69+
: -1;
70+
int consumed = 0;
71+
auto t_prefill_start = std::chrono::steady_clock::now();
72+
while (consumed < prompt_len) {
73+
int n_tokens = prompt_len - consumed;
74+
if (req.snap_pos >= 0 && req.snap_slot >= 0 &&
75+
req.snap_pos > base_pos + consumed &&
76+
req.snap_pos < base_pos + consumed + n_tokens) {
77+
n_tokens = req.snap_pos - (base_pos + consumed);
78+
}
79+
std::vector<int32_t> chunk(req.prompt.begin() + consumed,
80+
req.prompt.begin() + consumed + n_tokens);
81+
if (!adapter_->prefill(chunk, base_pos + consumed, last_tok)) {
82+
result.error = "prefill";
83+
return result;
84+
}
85+
consumed += n_tokens;
86+
if (req.snap_pos >= 0 && req.snap_slot >= 0 &&
87+
base_pos + consumed == req.snap_pos) {
88+
if (adapter_->snapshot_save(req.snap_slot)) {
89+
std::printf("[snap] inline slot=%d cur_pos=%d\n",
90+
req.snap_slot, req.snap_pos);
91+
std::fflush(stdout);
92+
}
93+
}
94+
}
95+
result.prefill_s = std::chrono::duration<double>(
96+
std::chrono::steady_clock::now() - t_prefill_start).count();
97+
98+
if (req.n_gen > 0) {
99+
if (last_tok < 0) {
100+
result.error = "decode_seed";
101+
return result;
102+
}
103+
auto t_decode_start = std::chrono::steady_clock::now();
104+
const bool ok = (base_pos == 0 && adapter_->can_dflash_decode())
105+
? adapter_->decode_dflash(req.prompt, base_pos, last_tok, req.n_gen,
106+
result.tokens, out_io)
107+
: adapter_->decode_ar(last_tok, base_pos + (int)req.prompt.size(), req.n_gen,
108+
result.tokens, out_io);
109+
if (!ok) {
110+
result.error = "decode";
111+
return result;
112+
}
113+
result.decode_s = std::chrono::duration<double>(
114+
std::chrono::steady_clock::now() - t_decode_start).count();
115+
}
116+
117+
result.ok = true;
118+
return result;
119+
}
120+
121+
GenerateResult LayerSplitBackend::generate(const GenerateRequest & req,
122+
const DaemonIO & io) {
123+
return run_from_state(req, io, /*base_pos=*/0, /*reset_state=*/true);
124+
}
125+
126+
bool LayerSplitBackend::snapshot_save(int slot) {
127+
return adapter_ && adapter_->snapshot_save(slot);
128+
}
129+
130+
void LayerSplitBackend::snapshot_free(int slot) {
131+
if (adapter_) adapter_->snapshot_free(slot);
132+
}
133+
134+
bool LayerSplitBackend::snapshot_used(int slot) const {
135+
return adapter_ && adapter_->snapshot_used(slot);
136+
}
137+
138+
int LayerSplitBackend::snapshot_cur_pos(int slot) const {
139+
return adapter_ ? adapter_->snapshot_cur_pos(slot) : 0;
140+
}
141+
142+
GenerateResult LayerSplitBackend::restore_and_generate(
143+
int slot, const GenerateRequest & req, const DaemonIO & io) {
144+
GenerateResult result;
145+
if (!adapter_ || !adapter_->snapshot_restore(slot)) {
146+
result.error = "bad slot";
147+
io.emit(-1);
148+
return result;
149+
}
150+
const int snap_pos = adapter_->snapshot_cur_pos(slot);
151+
if ((int)req.prompt.size() < snap_pos) {
152+
result.error = "snapshot_longer_than_prompt";
153+
io.emit(-1);
154+
return result;
155+
}
156+
GenerateRequest delta_req = req;
157+
delta_req.prompt = std::vector<int32_t>(
158+
req.prompt.begin() + snap_pos, req.prompt.end());
159+
return run_from_state(delta_req, io, snap_pos, /*reset_state=*/false);
160+
}
161+
162+
ModelBackend::CompressResult
163+
LayerSplitBackend::compress(const CompressRequest & req) {
164+
return adapter_ ? adapter_->compress(req) : CompressResult{};
165+
}
166+
167+
bool LayerSplitBackend::handle_compress(const std::string & line,
168+
const DaemonIO & io) {
169+
std::string args = line.size() > 9 ? line.substr(9) : std::string{};
170+
bool skip_park = false;
171+
const std::string suffix = " nopark";
172+
if (args.size() >= suffix.size() &&
173+
args.compare(args.size() - suffix.size(), suffix.size(), suffix) == 0) {
174+
skip_park = true;
175+
args.resize(args.size() - suffix.size());
176+
}
177+
178+
char ppath[1024];
179+
int keep_x1000 = 0;
180+
char drafter_path[1024] = {0};
181+
const int n = std::sscanf(args.c_str(), "%1023s %d %1023s",
182+
ppath, &keep_x1000, drafter_path);
183+
if (n < 2) {
184+
std::fprintf(stderr, "[target-split][compress] bad args\n");
185+
io.emit(-1);
186+
return false;
187+
}
188+
189+
CompressRequest req;
190+
req.input_ids = read_int32_file(ppath);
191+
req.keep_ratio = (float)keep_x1000 / 1000.0f;
192+
if (n >= 3 && drafter_path[0]) {
193+
req.drafter_path = drafter_path;
194+
} else if (adapter_) {
195+
req.drafter_path = adapter_->default_compress_drafter_path();
196+
}
197+
req.skip_park = skip_park;
198+
199+
CompressResult result = compress(req);
200+
for (int32_t t : result.compressed_ids) io.emit(t);
201+
io.emit(-1);
202+
return result.ok;
203+
}
204+
205+
void LayerSplitBackend::free_drafter() {
206+
if (adapter_) adapter_->free_drafter();
207+
}
208+
209+
bool LayerSplitBackend::supports_dflash_spec_decode() const {
210+
return adapter_ && adapter_->supports_dflash_spec_decode();
211+
}
212+
213+
DFlashTarget * LayerSplitBackend::dflash_target() {
214+
return adapter_ ? adapter_->dflash_target() : nullptr;
215+
}
216+
217+
bool LayerSplitBackend::supports_remote_draft() const {
218+
return adapter_ && adapter_->supports_remote_draft();
219+
}
220+
221+
void LayerSplitBackend::shutdown() {
222+
if (adapter_) adapter_->shutdown();
223+
}
224+
225+
} // namespace dflash::common

0 commit comments

Comments
 (0)