Skip to content

Commit 35ae645

Browse files
committed
refactor(server): generalize target layer-split adapter path
1 parent 6fd8d9e commit 35ae645

18 files changed

Lines changed: 255 additions & 190 deletions

dflash/src/common/dflash_layer_split_runtime.h

Lines changed: 6 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -1,11 +1,11 @@
11
// dflash_layer_split_runtime.h — target-agnostic runtime types for the
2-
// DFlash layer-split pipeline.
2+
// target layer-split pipeline.
33
//
4-
// Hosts the small pieces that are reused by every architecture's layer-split
5-
// driver: a runtime-configuration struct (replaces former globals) and an
6-
// activation double-buffer used to ferry hidden states between shards.
7-
// Architecture-specific shard layouts (e.g. qwen35's TargetLayerSplitShard
8-
// that embeds TargetWeights/TargetCache) live in their own headers.
4+
// Hosts the small runtime pieces reused by layer-split drivers: a
5+
// runtime-configuration struct and the activation double-buffer used to ferry
6+
// hidden states between shards. Shared placement/load-plan/shard metadata lives
7+
// in common/layer_split_utils.h; architecture-specific shard payloads keep their
8+
// own weights/cache/graph types.
99

1010
#pragma once
1111

dflash/src/common/layer_split_utils.cpp

Lines changed: 73 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,17 +1,22 @@
11
#include "layer_split_utils.h"
22

3+
#include "common/peer_access.h"
4+
#include "common/snapshot_backend.h"
5+
#include "ggml-cuda.h"
6+
37
#include <algorithm>
48
#include <cmath>
9+
#include <cstdio>
510
#include <set>
611

712
namespace dflash::common {
813

9-
std::vector<std::pair<int,int>> compute_layer_ranges(
14+
std::vector<LayerSplitRange> compute_layer_ranges(
1015
int n_layer,
1116
int n_gpus,
1217
const std::vector<double> & weights)
1318
{
14-
std::vector<std::pair<int,int>> ranges;
19+
std::vector<LayerSplitRange> ranges;
1520
if (n_layer <= 0 || n_gpus <= 0 || n_gpus > n_layer) return ranges;
1621

1722
std::vector<double> w = weights;
@@ -39,6 +44,72 @@ std::vector<std::pair<int,int>> compute_layer_ranges(
3944
return ranges;
4045
}
4146

47+
bool init_layer_split_shard_metas(
48+
std::vector<LayerSplitShardMeta *> shards,
49+
const std::vector<int> & gpus,
50+
const std::vector<LayerSplitRange> & ranges,
51+
const char * log_prefix) {
52+
if (shards.size() != gpus.size() || shards.size() != ranges.size()) return false;
53+
const char * prefix = log_prefix ? log_prefix : "target-split";
54+
for (size_t i = 0; i < shards.size(); ++i) {
55+
auto * shard = shards[i];
56+
if (!shard) return false;
57+
shard->gpu = gpus[i];
58+
shard->layer_begin = ranges[i].begin;
59+
shard->layer_end = ranges[i].end;
60+
shard->backend = ggml_backend_cuda_init(shard->gpu);
61+
if (!shard->backend) {
62+
std::fprintf(stderr, "[%s] backend init failed gpu=%d\n",
63+
prefix, shard->gpu);
64+
return false;
65+
}
66+
}
67+
return true;
68+
}
69+
70+
bool enable_layer_split_peer_access(
71+
const std::vector<int> & gpus,
72+
bool peer_access) {
73+
if (!peer_access) return true;
74+
for (size_t i = 0; i < gpus.size(); ++i) {
75+
for (size_t j = i + 1; j < gpus.size(); ++j) {
76+
(void)enable_peer_access_pair(gpus[i], gpus[j]);
77+
}
78+
}
79+
return true;
80+
}
81+
82+
bool init_layer_split_snapshot_backends(
83+
const std::vector<LayerSplitShardMeta *> & shards,
84+
std::vector<ggml_backend_t> & snapshot_backends,
85+
const char * log_prefix) {
86+
const char * prefix = log_prefix ? log_prefix : "target-split";
87+
snapshot_backends.assign(shards.size(), nullptr);
88+
for (size_t i = 0; i < shards.size(); ++i) {
89+
const auto * shard = shards[i];
90+
if (!shard || !shard->backend) return false;
91+
snapshot_backends[i] = create_snapshot_backend(shard->backend);
92+
if (!snapshot_backends[i]) {
93+
std::fprintf(stderr,
94+
"[%s] snapshot backend init failed gpu=%d\n",
95+
prefix, shard->gpu);
96+
return false;
97+
}
98+
}
99+
return true;
100+
}
101+
102+
void free_layer_split_snapshot_backends(
103+
const std::vector<LayerSplitShardMeta *> & shards,
104+
std::vector<ggml_backend_t> & snapshot_backends) {
105+
const size_t n = std::min(shards.size(), snapshot_backends.size());
106+
for (size_t i = 0; i < n; ++i) {
107+
if (!shards[i]) continue;
108+
free_snapshot_backend(snapshot_backends[i], shards[i]->backend);
109+
}
110+
snapshot_backends.clear();
111+
}
112+
42113
std::string validate_device_placement(
43114
const DevicePlacement & dp,
44115
int device_count)

dflash/src/common/layer_split_utils.h

Lines changed: 71 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -8,20 +8,89 @@
88
#include "placement/placement_config.h"
99

1010
#include <string>
11-
#include <utility>
1211
#include <vector>
1312

13+
#include "ggml-backend.h"
14+
1415
namespace dflash::common {
1516

17+
struct LayerSplitRange {
18+
int begin = 0;
19+
int end = 0;
20+
};
21+
22+
struct LayerSplitLoadPlan {
23+
int layer_begin = 0; // inclusive
24+
int layer_end = -1; // exclusive; <0 means all layers
25+
bool load_output = true; // final output norm / lm-head tensors
26+
};
27+
28+
struct LayerSplitShardMeta {
29+
int gpu = 0;
30+
int layer_begin = 0;
31+
int layer_end = 0;
32+
ggml_backend_t backend = nullptr;
33+
};
34+
35+
inline LayerSplitLoadPlan make_layer_split_load_plan(
36+
const LayerSplitShardMeta & shard,
37+
bool is_last_shard) {
38+
LayerSplitLoadPlan plan;
39+
plan.layer_begin = shard.layer_begin;
40+
plan.layer_end = shard.layer_end;
41+
plan.load_output = is_last_shard;
42+
return plan;
43+
}
44+
45+
template <typename Shard>
46+
std::vector<LayerSplitShardMeta *> layer_split_shard_metas(
47+
std::vector<Shard> & shards) {
48+
std::vector<LayerSplitShardMeta *> metas;
49+
metas.reserve(shards.size());
50+
for (auto & shard : shards) {
51+
metas.push_back(&shard);
52+
}
53+
return metas;
54+
}
55+
56+
template <typename Shard>
57+
Shard * find_layer_split_shard(std::vector<Shard> & shards, int layer_idx) {
58+
for (auto & shard : shards) {
59+
if (layer_idx >= shard.layer_begin && layer_idx < shard.layer_end) {
60+
return &shard;
61+
}
62+
}
63+
return nullptr;
64+
}
65+
1666
// Compute [begin, end) layer ranges for each GPU shard.
1767
// If weights is empty, splits layers equally.
1868
// If weights has entries, distributes proportionally (at least 1 layer per GPU).
1969
// Returns empty vector on error (n_layer <= 0 or n_gpus <= 0).
20-
std::vector<std::pair<int,int>> compute_layer_ranges(
70+
std::vector<LayerSplitRange> compute_layer_ranges(
2171
int n_layer,
2272
int n_gpus,
2373
const std::vector<double> & weights);
2474

75+
bool init_layer_split_shard_metas(
76+
std::vector<LayerSplitShardMeta *> shards,
77+
const std::vector<int> & gpus,
78+
const std::vector<LayerSplitRange> & ranges,
79+
const char * log_prefix);
80+
81+
bool enable_layer_split_peer_access(
82+
const std::vector<int> & gpus,
83+
bool peer_access);
84+
85+
bool init_layer_split_snapshot_backends(
86+
const std::vector<LayerSplitShardMeta *> & shards,
87+
std::vector<ggml_backend_t> & snapshot_backends,
88+
const char * log_prefix);
89+
90+
void free_layer_split_snapshot_backends(
91+
const std::vector<LayerSplitShardMeta *> & shards,
92+
std::vector<ggml_backend_t> & snapshot_backends);
93+
2594
// Validate a DevicePlacement against system constraints.
2695
// If device_count is negative, only validates structural constraints that do
2796
// not require querying the runtime-visible GPU count.

dflash/src/internal.h

Lines changed: 2 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -22,6 +22,7 @@
2222
#include "ggml-backend.h"
2323
#include "gguf.h"
2424

25+
#include "common/layer_split_utils.h"
2526
#include "dflash27b.h"
2627

2728
namespace dflash::common {
@@ -176,12 +177,6 @@ inline bool is_eos_tok(int tok, const TargetWeights & w) {
176177
|| (w.eos_id >= 0 && tok == w.eos_id);
177178
}
178179

179-
struct TargetLoadPlan {
180-
int layer_begin = 0; // inclusive
181-
int layer_end = -1; // exclusive; <0 means all layers
182-
bool load_output = true; // output_norm + lm_head
183-
};
184-
185180
// Load a Q4_K_M target model from a GGUF file on disk.
186181
// Returns false and sets last_error on failure.
187182
bool load_target_gguf(const std::string & path,
@@ -190,7 +185,7 @@ bool load_target_gguf(const std::string & path,
190185

191186
bool load_target_gguf_partial(const std::string & path,
192187
ggml_backend_t backend,
193-
const TargetLoadPlan & plan,
188+
const LayerSplitLoadPlan & plan,
194189
TargetWeights & out);
195190

196191
void free_target_weights(TargetWeights & w);

dflash/src/qwen35/gguf_target_loader.cpp

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -44,6 +44,7 @@
4444
// tensor's bytes from the mmap'd file.
4545

4646
#include "internal.h"
47+
#include "common/layer_split_utils.h"
4748

4849
#include <cinttypes>
4950
#include <cstdint>
@@ -239,13 +240,13 @@ float get_f32_or(const gguf_context * g, const char * key, float fallback) {
239240
bool load_target_gguf(const std::string & path,
240241
ggml_backend_t backend,
241242
TargetWeights & out) {
242-
TargetLoadPlan plan;
243+
LayerSplitLoadPlan plan;
243244
return load_target_gguf_partial(path, backend, plan, out);
244245
}
245246

246247
bool load_target_gguf_partial(const std::string & path,
247248
ggml_backend_t backend,
248-
const TargetLoadPlan & plan_in,
249+
const LayerSplitLoadPlan & plan_in,
249250
TargetWeights & out) {
250251

251252
// ── 1. Parse metadata + create a ggml_context holding tensor descriptors ─
@@ -361,7 +362,7 @@ bool load_target_gguf_partial(const std::string & path,
361362
}
362363
}
363364

364-
TargetLoadPlan plan = plan_in;
365+
LayerSplitLoadPlan plan = plan_in;
365366
if (plan.layer_begin < 0) plan.layer_begin = 0;
366367
if (plan.layer_end < 0) plan.layer_end = (int)n_layer;
367368
if (plan.layer_begin > plan.layer_end ||

dflash/src/qwen35/layer_split_daemon.cpp

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -14,8 +14,8 @@
1414

1515
namespace dflash::common {
1616

17-
bool run_target_layer_split_request(
18-
std::vector<TargetLayerSplitShard> & shards,
17+
bool run_qwen35_layer_split_request(
18+
std::vector<Qwen35LayerSplitShard> & shards,
1919
DraftWeights * draft_weights,
2020
ggml_backend_t draft_backend,
2121
int draft_gpu,
@@ -41,7 +41,7 @@ bool run_target_layer_split_request(
4141
ubatch = std::max(1, std::atoi(s));
4242
}
4343
int last_tok = -1;
44-
if (!run_target_layer_split_forward(shards, shards.front().weights,
44+
if (!run_qwen35_layer_split_forward(shards, shards.front().weights,
4545
prompt, 0, ubatch, last_tok,
4646
kq_stride_pad, fa_window,
4747
feature_ring)) {
@@ -66,7 +66,7 @@ bool run_target_layer_split_request(
6666
for (; generated < n_gen; generated++) {
6767
std::vector<int32_t> one(1, last_tok);
6868
int next_tok = -1;
69-
if (!run_target_layer_split_forward(shards, shards.front().weights,
69+
if (!run_qwen35_layer_split_forward(shards, shards.front().weights,
7070
one, (int)out_all.size(), 1, next_tok,
7171
kq_stride_pad, fa_window,
7272
feature_ring)) {

dflash/src/qwen35/layer_split_daemon.h

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,6 @@
11
// layer_split_daemon.h — Layer-split request handler for qwen35 daemon mode.
22
//
3-
// run_target_layer_split_request() handles a single inference request:
3+
// run_qwen35_layer_split_request() handles a single inference request:
44
// prefill → (optional spec-decode or AR decode) → output.
55

66
#pragma once
@@ -22,8 +22,8 @@ namespace dflash::common {
2222
// Runs prefill, then either spec-decode (if run_dflash && draft available)
2323
// or plain AR decode. Emits tokens to stream_fd and optionally writes
2424
// the full sequence to out_path.
25-
bool run_target_layer_split_request(
26-
std::vector<TargetLayerSplitShard> & shards,
25+
bool run_qwen35_layer_split_request(
26+
std::vector<Qwen35LayerSplitShard> & shards,
2727
DraftWeights * draft_weights,
2828
ggml_backend_t draft_backend,
2929
int draft_gpu,

0 commit comments

Comments
 (0)