Skip to content
Merged
Show file tree
Hide file tree
Changes from all 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
252 changes: 219 additions & 33 deletions server/src/glm5next/glm5next_backend.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -85,16 +85,83 @@ bool Glm5NextBackend::init_hybrid_model() {
cache_.cur_pos = 0;
cache_.n_past = 0;

// Build MoE hybrid storage for expert evaluation (all-hot, GPU-only)
// Allocate KV cache tensors for MLA layers and KDA state
const int n_mla_layers = (w_.n_layer + w_.full_attn_interval - 1) / w_.full_attn_interval;
const int n_kda_layers = w_.n_layer - n_mla_layers;
const int kv_dim = w_.head_dim; // MLA uses single KV head (absorbed form)

ggml_init_params cache_params = {
/*.mem_size =*/ 512 * 1024 * 1024, // 512MB for cache metadata
/*.mem_buffer =*/ nullptr,
/*.no_alloc =*/ false,
};
ggml_context * cache_ctx = ggml_init(cache_params);
if (!cache_ctx) {
std::fprintf(stderr, "[glm5next] failed to init cache context\n");
return false;
}

// MLA KV cache: [head_dim, n_ctx, n_mla_layers]
cache_.k = ggml_new_tensor_3d(cache_ctx, GGML_TYPE_F16, kv_dim, cache_.n_ctx, n_mla_layers);
cache_.v = ggml_new_tensor_3d(cache_ctx, GGML_TYPE_F16, kv_dim, cache_.n_ctx, n_mla_layers);
ggml_set_name(cache_.k, "cache_k");
ggml_set_name(cache_.v, "cache_v");

// KDA recurrent state: [head_dim, n_head, n_kda_layers] (hidden state per head)
cache_.kda_state = ggml_new_tensor_3d(cache_ctx, GGML_TYPE_F32,
w_.head_dim, w_.n_head, n_kda_layers);
ggml_set_name(cache_.kda_state, "kda_state");

// Allocate cache on backend
ggml_backend_buffer_t cache_buf = ggml_backend_alloc_ctx_tensors(cache_ctx, backend_);
if (!cache_buf) {
std::fprintf(stderr, "[glm5next] failed to allocate cache buffer\n");
ggml_free(cache_ctx);
return false;
}

// Zero-initialize caches
ggml_backend_tensor_memset(cache_.k, 0, 0, ggml_nbytes(cache_.k));
ggml_backend_tensor_memset(cache_.v, 0, 0, ggml_nbytes(cache_.v));
ggml_backend_tensor_memset(cache_.kda_state, 0, 0, ggml_nbytes(cache_.kda_state));

std::fprintf(stderr, "[glm5next] cache allocated: %d MLA layers, %d KDA layers, ctx=%d\n",
n_mla_layers, n_kda_layers, cache_.n_ctx);

// Build MoE hybrid storage for expert evaluation
// Dual-GPU: experts on device 1 (gfx1151), hot path on device 0 (gfx1100)
const int n_moe_layers = w_.n_layer - w_.first_moe_layer;
if (n_moe_layers > 0) {
// Create all-hot placement: all 288 experts on GPU
// Check for dual-GPU setup: device 1 for experts
const int expert_gpu = 1; // gfx1151 UMA
ggml_backend_t expert_backend = nullptr;
bool dual_gpu = false;

// Try to initialize expert backend on device 1
expert_backend = ggml_backend_cuda_init(expert_gpu);
if (expert_backend) {
dual_gpu = true;
std::fprintf(stderr, "[glm5next] dual-GPU: device 0 (hot path), device 1 (experts)\n");
} else {
std::fprintf(stderr, "[glm5next] device 1 unavailable, using single-GPU all-hot\n");
}

// Placement: all experts on "cold" backend (device 1) if dual-GPU, else all-hot
moe_placement_.n_expert = w_.n_expert;
moe_placement_.hot_expert_ids.resize(n_moe_layers);
for (int il = 0; il < n_moe_layers; ++il) {
moe_placement_.hot_expert_ids[il].resize(w_.n_expert);
for (int e = 0; e < w_.n_expert; ++e) {
moe_placement_.hot_expert_ids[il][e] = e; // All hot

if (dual_gpu) {
// Dual-GPU: all experts on device 1 (cold backend), none on device 0 (hot)
for (int il = 0; il < n_moe_layers; ++il) {
moe_placement_.hot_expert_ids[il].clear(); // No experts on device 0
}
} else {
// Single-GPU fallback: all experts on device 0 (hot)
for (int il = 0; il < n_moe_layers; ++il) {
moe_placement_.hot_expert_ids[il].resize(w_.n_expert);
for (int e = 0; e < w_.n_expert; ++e) {
moe_placement_.hot_expert_ids[il][e] = e;
}
}
}

Expand Down Expand Up @@ -123,20 +190,38 @@ bool Glm5NextBackend::init_hybrid_model() {
moe_cfg.n_layer = n_moe_layers;
moe_cfg.first_moe_layer = w_.first_moe_layer;
moe_cfg.swiglu_clamp = w_.swiglu_clamp;
moe_cfg.cold_expert_backend = MoeHybridColdBackend::Cpu;
moe_cfg.materialize_hot_experts = true;
moe_cfg.materialize_cold_experts = false; // All on GPU

if (dual_gpu) {
// Dual-GPU: experts on device 1
moe_cfg.cold_expert_backend = MoeHybridColdBackend::Gpu;
moe_cfg.materialize_hot_experts = false; // No hot experts on device 0
moe_cfg.materialize_cold_experts = true; // All experts on device 1
} else {
// Single-GPU: all experts hot on device 0
moe_cfg.cold_expert_backend = MoeHybridColdBackend::Cpu;
moe_cfg.materialize_hot_experts = true;
moe_cfg.materialize_cold_experts = false;
}

moe_hybrid_ = std::make_shared<MoeHybridStorage>();
std::string err;
if (!build_moe_hybrid_storage(moe_cfg, backend_, moe_placement_,
layer_descs, *moe_hybrid_, &err)) {
layer_descs, *moe_hybrid_, &err,
dual_gpu ? expert_backend : nullptr)) {
std::fprintf(stderr, "[glm5next] failed to build MoE storage: %s\n", err.c_str());
if (dual_gpu && expert_backend) {
ggml_backend_free(expert_backend);
}
return false;
}

std::fprintf(stderr, "[glm5next] MoE storage initialized: %d layers, %d experts (all GPU)\n",
n_moe_layers, w_.n_expert);
if (dual_gpu) {
std::fprintf(stderr, "[glm5next] MoE storage: %d layers, %d experts on device 1 (gfx1151)\n",
n_moe_layers, w_.n_expert);
} else {
std::fprintf(stderr, "[glm5next] MoE storage: %d layers, %d experts on device 0 (all-hot)\n",
n_moe_layers, w_.n_expert);
}
}

std::fprintf(stderr, "[glm5next] backend initialized: ctx=%d\n", cache_.n_ctx);
Expand Down Expand Up @@ -183,21 +268,40 @@ GenerateResult Glm5NextBackend::generate_impl(
GenerateResult result;
result.status = GenerateStatus::OK;

// Simplified generation: just build graph for first token and return
// Full implementation would do proper prefill, decode loop, and sampling

if (req.prompt.empty()) {
result.status = GenerateStatus::Error;
result.error_message = "empty prompt";
return result;
}

// Setup sampler
sampler_ = req.sampler;
if (req.do_sample && sampler_.seed != 0) {
sampler_rng_.seed(sampler_.seed);
}

const bool process_logits = sampler_.needs_logit_processing();
std::vector<int32_t> history;
if (process_logits) {
history = req.prompt;
if (req.n_gen > 0) {
history.reserve(history.size() + (size_t)req.n_gen);
}
}

// Prefill: process prompt tokens
const int prompt_len = (int)req.prompt.size();
std::vector<int32_t> out_tokens;
out_tokens.reserve((size_t)req.n_gen);

std::fprintf(stderr, "[glm5next] prefill: %d tokens\n", prompt_len);

// Allocate graph context
const size_t graph_ctx_size = 128 * 1024 * 1024; // 128MB
ggml_init_params params = {
/*.mem_size =*/ graph_ctx_size,
/*.mem_buffer =*/ nullptr,
/*.no_alloc =*/ true, // Use backend allocator
/*.no_alloc =*/ true,
};

ggml_context * ctx = ggml_init(params);
Expand All @@ -207,47 +311,129 @@ GenerateResult Glm5NextBackend::generate_impl(
return result;
}

// Build forward graph for first token
// Build forward graph for prompt
ggml_cgraph * gf = ggml_new_graph(ctx);

ggml_tensor * logits = glm5next_build_graph(
ctx, w_, cache_,
req.prompt.data(), 1, // Just first token for now
req.prompt.data(), prompt_len,
cache_.cur_pos,
moe_hybrid_.get() // Pass MoE storage for expert evaluation
moe_hybrid_.get()
);

if (!logits) {
ggml_free(ctx);
result.status = GenerateStatus::Error;
result.error_message = "graph construction failed";
result.error_message = "prefill graph construction failed";
return result;
}

ggml_build_forward_expand(gf, logits);

std::fprintf(stderr, "[glm5next] graph built: %d nodes, %d leaves\n",
gf->n_nodes, gf->n_leafs);

// Compute graph
// Compute prefill
if (ggml_backend_graph_compute(backend_, gf) != GGML_STATUS_SUCCESS) {
ggml_free(ctx);
result.status = GenerateStatus::Error;
result.error_message = "graph compute failed";
result.error_message = "prefill compute failed";
return result;
}

std::fprintf(stderr, "[glm5next] graph computed successfully\n");
// Read logits from last token position
std::vector<float> logits_vec((size_t)w_.n_vocab);
const size_t last_token_offset = (size_t)(prompt_len - 1) * (size_t)w_.n_vocab * sizeof(float);
ggml_backend_tensor_get(logits, logits_vec.data(), last_token_offset,
sizeof(float) * (size_t)w_.n_vocab);

// Sample (simplified: just return EOS)
const int32_t eos_token = 2; // Typical EOS
io.emit(eos_token);
io.emit(-1); // Sentinel

result.n_gen = 1;
result.n_past = 1;
// Update cache position after prefill
cache_.cur_pos += prompt_len;
cache_.n_past += prompt_len;

ggml_free(ctx);
ctx = nullptr;

std::fprintf(stderr, "[glm5next] prefill complete, cur_pos=%d\n", cache_.cur_pos);

// Decode loop: generate n_gen tokens
for (int generated = 0; generated < req.n_gen; ++generated) {
if (io.is_cancelled()) break;

// Sample next token
int32_t next_token = 0;
if (process_logits) {
next_token = sample_logits(logits_vec.data(), w_.n_vocab, sampler_,
history, sampler_rng_);
history.push_back(next_token);
} else {
// Greedy: argmax
float max_val = logits_vec[0];
for (int i = 1; i < w_.n_vocab; ++i) {
if (logits_vec[i] > max_val) {
max_val = logits_vec[i];
next_token = i;
}
}
}

// Emit token
io.emit(next_token);
out_tokens.push_back(next_token);

// Check for EOS
const int32_t eos_token = 2; // Typical EOS
if (next_token == eos_token) {
std::fprintf(stderr, "[glm5next] EOS at position %zu\n", out_tokens.size());
break;
}

// Compute next token logits
ctx = ggml_init(params);
if (!ctx) {
result.status = GenerateStatus::Error;
result.error_message = "failed to create decode graph context";
break;
}

gf = ggml_new_graph(ctx);
logits = glm5next_build_graph(
ctx, w_, cache_,
&next_token, 1,
cache_.cur_pos,
moe_hybrid_.get()
);

if (!logits) {
ggml_free(ctx);
result.status = GenerateStatus::Error;
result.error_message = "decode graph construction failed";
break;
}

ggml_build_forward_expand(gf, logits);

if (ggml_backend_graph_compute(backend_, gf) != GGML_STATUS_SUCCESS) {
ggml_free(ctx);
result.status = GenerateStatus::Error;
result.error_message = "decode compute failed";
break;
}

// Read logits (single token, no offset)
ggml_backend_tensor_get(logits, logits_vec.data(), 0,
sizeof(float) * (size_t)w_.n_vocab);

// Update cache position
cache_.cur_pos += 1;
cache_.n_past += 1;

ggml_free(ctx);
ctx = nullptr;
}

// Emit sentinel
io.emit(-1);

result.n_gen = (int)out_tokens.size();
result.n_past = cache_.n_past;
return result;
}

Expand Down
Loading