engine: true parallel batching of the talker and predictor, per slot codec streams

This commit is contained in:
Pascal
2026-07-20 16:53:36 +02:00
parent fbff3317a4
commit 37ea692be6
14 changed files with 1663 additions and 672 deletions
+178 -126
View File
@@ -1,26 +1,29 @@
#pragma once #pragma once
// code-predictor-forward.h: run the 5-layer Qwen3 code predictor over a // code-predictor-forward.h: run the 5-layer Qwen3 code predictor over a
// growing context to produce the 15 acoustic codes of one audio frame, // growing context to produce the 15 acoustic codes of one audio frame
// KV cached. // for every active slot in one batched pass, KV cached per slot.
// //
// Input: // Input:
// hidden_bridge [hidden] f32 -- persistent backend tensor holding // hidden_bridge [hidden, max_batch] f32 -- persistent backend tensor
// the talker last position hidden // holding the talker last
// (post final norm), written on // position hidden per slot (post
// device by the talker graph and // final norm), written on device
// read here as a graph leaf // by the talker graph and read
// c0 -- semantic code sampled from the // here as a graph leaf
// Talker codec_head (codebook 0) // c0[N] -- semantic codes sampled from the
// Talker codec_head (codebook 0),
// one per slot
// Output: // Output:
// codes[16] = [c0, c1, ..., c15] -- the full set of codes for one // codes[N * 16] -- per slot [c0, c1, ..., c15],
// frame, ready for decode through // slot major, ready for decode
// the codec // through the codec
// //
// The predictor cache is local to a single frame: we reset it at every // The predictor cache is local to a single frame and holds one set per
// frame, prefill the first two positions (talker_hidden + embed(c0)), // slot: the prefill writes the first two positions (talker_hidden +
// then decode 14 single-token steps. Total work drops from // embed(c0)) of every slot, then 14 batched single-token steps follow.
// O(sum_{g=0..14} (g+2)^2) = O(1496 token-steps) to O(16) per frame, // Every slot runs the same sub-step sequence every frame, so the batch
// roughly 90x for the inner loop. // stays in perfect lockstep and the graphs bake positions, kv rows and
// masks at build time.
// //
// Architecture mirrors the Talker block, only differences are: // Architecture mirrors the Talker block, only differences are:
// - 5 layers instead of 28 // - 5 layers instead of 28
@@ -28,10 +31,10 @@
// - one private embedding table and one private linear head per // - one private embedding table and one private linear head per
// acoustic codebook (1..15) // acoustic codebook (1..15)
// //
// Graph metadata lives in caller owned static graphs, one for the T=2 // Graph metadata lives in caller owned static graphs, one per flavor
// prefill and one per T=1 step: each graph is built and allocated once // and batch width N (prefill plus one per acoustic step), built lazily
// at load, then replayed directly on the backend with a 4 byte code id // on the first frame at a given N, then replayed directly on the
// upload per call, the positions, kv rows, and causal mask baked in. // backend with an N * 4 byte code id upload per call.
#include "code-predictor-graph.h" #include "code-predictor-graph.h"
#include "code-predictor-weights.h" #include "code-predictor-weights.h"
@@ -51,13 +54,15 @@
#include <vector> #include <vector>
struct CodePredictorOutput { struct CodePredictorOutput {
// Sixteen codes: c0 from the talker plus c1..c15 from the predictor. // Per slot sixteen codes, slot major: codes[slot * 16 + g] holds c0
// from the talker plus c1..c15 from the predictor.
std::vector<int32_t> codes; std::vector<int32_t> codes;
}; };
// Manual F32 attention chain for the code predictor block. Same shape // Manual F32 attention chain for the code predictor block. Same shape
// contract as talker_attn_f32: q [hd, T, n_q_heads], k/v [hd, T_full, // contract as talker_attn_f32: q [hd, T, n_q_heads, N], k/v
// n_kv], output [hd, n_q_heads, T]. Used when use_flash_attn is false. // [hd, T_full, n_kv, N], output [hd, n_q_heads, T, N]; the mul_mat
// broadcasts over dims 2 and 3. Used when use_flash_attn is false.
static struct ggml_tensor * code_predictor_attn_f32(struct ggml_context * ctx, static struct ggml_tensor * code_predictor_attn_f32(struct ggml_context * ctx,
struct ggml_tensor * q, struct ggml_tensor * q,
struct ggml_tensor * k, struct ggml_tensor * k,
@@ -76,10 +81,12 @@ static int code_predictor_graph_max_nodes(int n_layers) {
return 48 * n_layers + 64; return 48 * n_layers + 64;
} }
// One Qwen3 decoder block, KV cached. K and V for the T fresh positions // One batched Qwen3 decoder block, KV cached over sets [0, N). x holds
// are written into the cache at [n_past, n_past+T) on dim 1; the // the N slots' token columns flattened slot major: column j = n * T + t
// attention reads the fixed [0, n_kv_pad) window with the mask carrying // is position t of slot n. K and V for the fresh positions are written
// neg inf beyond n_past+T. Returns the layer output [hidden, T]. // into each slot's set at the rows carried by kv_rows; the attention
// reads the fixed [0, n_kv_pad) window per set with the mask carrying
// neg inf beyond n_past + T. Returns the layer output [hidden, T * N].
// use_flash_attn and clamp_fp16 follow the same contract as in // use_flash_attn and clamp_fp16 follow the same contract as in
// talker-forward.h. // talker-forward.h.
static struct ggml_tensor * code_predictor_layer_forward(struct ggml_context * ctx, static struct ggml_tensor * code_predictor_layer_forward(struct ggml_context * ctx,
@@ -89,9 +96,10 @@ static struct ggml_tensor * code_predictor_layer_forward(struct ggml_context *
struct ggml_tensor * positions, struct ggml_tensor * positions,
struct ggml_tensor * mask, struct ggml_tensor * mask,
struct ggml_tensor * kv_rows, struct ggml_tensor * kv_rows,
struct ggml_tensor * k_cache, struct ggml_tensor * k4,
struct ggml_tensor * v_cache, struct ggml_tensor * v4,
int T, int T,
int N,
int n_kv_pad, int n_kv_pad,
bool use_flash_attn, bool use_flash_attn,
bool clamp_fp16, bool clamp_fp16,
@@ -100,6 +108,7 @@ static struct ggml_tensor * code_predictor_layer_forward(struct ggml_context *
const int n_kv = cw->num_key_value_heads; const int n_kv = cw->num_key_value_heads;
const int hd = cw->head_dim; const int hd = cw->head_dim;
const float eps = cw->rms_norm_eps; const float eps = cw->rms_norm_eps;
const int TN = T * N;
struct ggml_tensor * h = ggml_rms_norm(ctx, x, eps); struct ggml_tensor * h = ggml_rms_norm(ctx, x, eps);
h = ggml_mul(ctx, h, layer.input_norm_w); h = ggml_mul(ctx, h, layer.input_norm_w);
@@ -108,34 +117,43 @@ static struct ggml_tensor * code_predictor_layer_forward(struct ggml_context *
struct ggml_tensor * k = ggml_mul_mat(ctx, layer.attn.k_proj_w, h); struct ggml_tensor * k = ggml_mul_mat(ctx, layer.attn.k_proj_w, h);
struct ggml_tensor * v = ggml_mul_mat(ctx, layer.attn.v_proj_w, h); struct ggml_tensor * v = ggml_mul_mat(ctx, layer.attn.v_proj_w, h);
q = ggml_reshape_3d(ctx, q, hd, n_q_heads, T); q = ggml_reshape_3d(ctx, q, hd, n_q_heads, TN);
k = ggml_reshape_3d(ctx, k, hd, n_kv, T); k = ggml_reshape_3d(ctx, k, hd, n_kv, TN);
v = ggml_reshape_3d(ctx, v, hd, n_kv, T); v = ggml_reshape_3d(ctx, v, hd, n_kv, TN);
q = ggml_rms_norm(ctx, q, eps); q = ggml_rms_norm(ctx, q, eps);
q = ggml_mul(ctx, q, layer.attn.q_norm_w); q = ggml_mul(ctx, q, layer.attn.q_norm_w);
k = ggml_rms_norm(ctx, k, eps); k = ggml_rms_norm(ctx, k, eps);
k = ggml_mul(ctx, k, layer.attn.k_norm_w); k = ggml_mul(ctx, k, layer.attn.k_norm_w);
// RoPE over the flattened token axis: positions [T * N] repeat the
// same in-frame offsets for every slot.
q = ggml_rope_ext(ctx, q, positions, NULL, hd, GGML_ROPE_TYPE_NEOX, 0, cw->rope_theta, 1.0f, 0.0f, 1.0f, 0.0f, q = ggml_rope_ext(ctx, q, positions, NULL, hd, GGML_ROPE_TYPE_NEOX, 0, cw->rope_theta, 1.0f, 0.0f, 1.0f, 0.0f,
0.0f); 0.0f);
k = ggml_rope_ext(ctx, k, positions, NULL, hd, GGML_ROPE_TYPE_NEOX, 0, cw->rope_theta, 1.0f, 0.0f, 1.0f, 0.0f, k = ggml_rope_ext(ctx, k, positions, NULL, hd, GGML_ROPE_TYPE_NEOX, 0, cw->rope_theta, 1.0f, 0.0f, 1.0f, 0.0f,
0.0f); 0.0f);
// Write the fresh positions into the cache via set_rows: positions // Write the fresh positions into each slot's set via set_rows:
// travel as data so every step keeps an identical topology and the // [hd, heads, T*N] unflattens to [hd, heads, T, N], permutes to
// captured CUDA graph replays without an update. // [hd, T, heads, N] and lands at the kv_rows [T, 1, N] destinations,
struct ggml_tensor * k_perm = ggml_cont(ctx, ggml_permute(ctx, k, 0, 2, 1, 3)); // [hd, T, n_kv] // broadcast across the n_kv head dim.
struct ggml_tensor * v_perm = ggml_cont(ctx, ggml_permute(ctx, v, 0, 2, 1, 3)); struct ggml_tensor * k_perm =
ggml_cont(ctx, ggml_permute(ctx, ggml_reshape_4d(ctx, k, hd, n_kv, T, N), 0, 2, 1, 3));
struct ggml_tensor * v_perm =
ggml_cont(ctx, ggml_permute(ctx, ggml_reshape_4d(ctx, v, hd, n_kv, T, N), 0, 2, 1, 3));
ggml_build_forward_expand(gf, ggml_set_rows(ctx, k_cache, k_perm, kv_rows)); struct ggml_tensor * k_sets = ggml_view_4d(ctx, k4, hd, k4->ne[1], n_kv, N, k4->nb[1], k4->nb[2], k4->nb[3], 0);
ggml_build_forward_expand(gf, ggml_set_rows(ctx, v_cache, v_perm, kv_rows)); struct ggml_tensor * v_sets = ggml_view_4d(ctx, v4, hd, v4->ne[1], n_kv, N, v4->nb[1], v4->nb[2], v4->nb[3], 0);
struct ggml_tensor * k_full = ggml_view_3d(ctx, k_cache, hd, n_kv_pad, n_kv, k_cache->nb[1], k_cache->nb[2], 0); ggml_build_forward_expand(gf, ggml_set_rows(ctx, k_sets, k_perm, kv_rows));
struct ggml_tensor * v_full = ggml_view_3d(ctx, v_cache, hd, n_kv_pad, n_kv, v_cache->nb[1], v_cache->nb[2], 0); ggml_build_forward_expand(gf, ggml_set_rows(ctx, v_sets, v_perm, kv_rows));
// Q permute [hd, n_q_heads, T] -> [hd, T, n_q_heads] for flash_attn_ext. struct ggml_tensor * k_full = ggml_view_4d(ctx, k4, hd, n_kv_pad, n_kv, N, k4->nb[1], k4->nb[2], k4->nb[3], 0);
struct ggml_tensor * q_p = ggml_permute(ctx, q, 0, 2, 1, 3); struct ggml_tensor * v_full = ggml_view_4d(ctx, v4, hd, n_kv_pad, n_kv, N, v4->nb[1], v4->nb[2], v4->nb[3], 0);
// Q [hd, n_q_heads, T, N] -> [hd, T, n_q_heads, N] for
// flash_attn_ext, taken as a view like the talker path does.
struct ggml_tensor * q_p = ggml_permute(ctx, ggml_reshape_4d(ctx, q, hd, n_q_heads, T, N), 0, 2, 1, 3);
// Clamp V before attention when clamp_fp16 is set, same rationale // Clamp V before attention when clamp_fp16 is set, same rationale
// as the talker block: sub Ampere CUDA tensor cores accumulate in // as the talker block: sub Ampere CUDA tensor cores accumulate in
@@ -144,9 +162,7 @@ static struct ggml_tensor * code_predictor_layer_forward(struct ggml_context *
v_full = ggml_clamp(ctx, v_full, -65504.0f, 65504.0f); v_full = ggml_clamp(ctx, v_full, -65504.0f, 65504.0f);
} }
// Attention: fused flash kernel or manual F32 chain. Matches the // Attention: fused flash kernel or manual F32 chain.
// working acestep qw3lm_build_attn pattern on the fused branch and
// the omnivoice qwen3_attn_f32 helper on the manual one.
float scale = 1.0f / sqrtf((float) hd); float scale = 1.0f / sqrtf((float) hd);
struct ggml_tensor * attn; struct ggml_tensor * attn;
if (use_flash_attn) { if (use_flash_attn) {
@@ -156,7 +172,9 @@ static struct ggml_tensor * code_predictor_layer_forward(struct ggml_context *
attn = code_predictor_attn_f32(ctx, q_p, k_full, v_full, mask, scale); attn = code_predictor_attn_f32(ctx, q_p, k_full, v_full, mask, scale);
} }
attn = ggml_reshape_2d(ctx, attn, n_q_heads * hd, T); // [hd, n_q_heads, T, N] -> [n_q_heads*hd, T*N], flatten heads for
// o_proj, token order matching x.
attn = ggml_reshape_2d(ctx, attn, n_q_heads * hd, TN);
struct ggml_tensor * o = ggml_mul_mat(ctx, layer.attn.o_proj_w, attn); struct ggml_tensor * o = ggml_mul_mat(ctx, layer.attn.o_proj_w, attn);
x = ggml_add(ctx, x, o); x = ggml_add(ctx, x, o);
@@ -180,18 +198,20 @@ static struct ggml_tensor * code_predictor_layer_forward(struct ggml_context *
return x; return x;
} }
// Build one static predictor graph. A non NULL hidden_bridge selects // Build one static batched predictor graph over sets [0, N). A non
// the T=2 prefill flavor reading [talker_hidden, embed(c0)] through // NULL hidden_bridge selects the T=2 prefill flavor reading
// lm_head[0]; otherwise the graph is the single token step for g_head, // [talker_hidden, embed(c0)] per slot through lm_head[0]; otherwise
// appending at the fixed cache row g_head + 1. The logits node holds // the graph is the single token step for g_head, appending at the
// the last position only, so every flavor reads back one row at // fixed cache row g_head + 1. The logits node holds the last position
// offset zero. use_flash_attn / clamp_fp16 apply to every layer. // of every slot as [Vg, N]. use_flash_attn / clamp_fp16 apply to every
// layer.
static bool code_predictor_graph_build(const CodePredictorWeights * cw, static bool code_predictor_graph_build(const CodePredictorWeights * cw,
KVCache * kv, KVCache * kv,
ggml_backend_t backend, ggml_backend_t backend,
struct ggml_tensor * embd_table, struct ggml_tensor * embd_table,
struct ggml_tensor * hidden_bridge, struct ggml_tensor * hidden_bridge,
int g_head, int g_head,
int N,
bool use_flash_attn, bool use_flash_attn,
bool clamp_fp16, bool clamp_fp16,
CodePredGraph * cp) { CodePredGraph * cp) {
@@ -214,17 +234,18 @@ static bool code_predictor_graph_build(const CodePredictorWeights * cw,
} }
struct ggml_context * gctx = cp->ctx; struct ggml_context * gctx = cp->ctx;
// Inputs: one code id gathered in graph from embd_table, positions, // Inputs: one code id per slot gathered in graph from embd_table,
// attention mask. The prefill path (T == 2, hidden_bridge non NULL) // positions, kv rows and the attention mask. The prefill path
// concats the resident talker hidden ahead of embed(c0), both on // (T == 2) concats each slot's resident talker hidden ahead of
// device: the sequence is [talker_hidden, embed(c0)] with zero row // embed(c0), both on device: the per slot sequence is
// upload. Steps (T == 1) are pure gathers: the only per step upload // [talker_hidden, embed(c0)] with zero row upload. Steps (T == 1)
// is 4 bytes of code id. // are pure gathers: the only per step upload is N * 4 bytes of
struct ggml_tensor * ids_in = ggml_new_tensor_1d(gctx, GGML_TYPE_I32, 1); // code ids.
struct ggml_tensor * pos_in = ggml_new_tensor_1d(gctx, GGML_TYPE_I32, T); struct ggml_tensor * ids_in = ggml_new_tensor_1d(gctx, GGML_TYPE_I32, N);
struct ggml_tensor * mask_in = ggml_new_tensor_2d(gctx, GGML_TYPE_F16, n_kv_pad, T); struct ggml_tensor * pos_in = ggml_new_tensor_1d(gctx, GGML_TYPE_I32, T * N);
struct ggml_tensor * rows_in = ggml_new_tensor_1d(gctx, GGML_TYPE_I64, T); struct ggml_tensor * mask_in = ggml_new_tensor_4d(gctx, GGML_TYPE_F16, n_kv_pad, T, 1, N);
ggml_set_name(ids_in, "sub_code_id"); struct ggml_tensor * rows_in = ggml_new_tensor_3d(gctx, GGML_TYPE_I64, T, 1, N);
ggml_set_name(ids_in, "sub_code_ids");
ggml_set_name(pos_in, "positions"); ggml_set_name(pos_in, "positions");
ggml_set_name(mask_in, "causal_mask"); ggml_set_name(mask_in, "causal_mask");
ggml_set_name(rows_in, "kv_rows"); ggml_set_name(rows_in, "kv_rows");
@@ -240,9 +261,14 @@ static bool code_predictor_graph_build(const CodePredictorWeights * cw,
ggml_set_input(rows_in); ggml_set_input(rows_in);
ggml_set_output(rows_in); ggml_set_output(rows_in);
struct ggml_tensor * x_in = ggml_get_rows(gctx, embd_table, ids_in); struct ggml_tensor * x_in = ggml_get_rows(gctx, embd_table, ids_in); // [hidden, N]
if (T == 2) { if (T == 2) {
x_in = ggml_concat(gctx, hidden_bridge, x_in, 1); struct ggml_tensor * bridge_cols =
ggml_view_2d(gctx, hidden_bridge, hidden_bridge->ne[0], N, hidden_bridge->nb[1], 0);
struct ggml_tensor * b3 = ggml_reshape_3d(gctx, bridge_cols, hidden_bridge->ne[0], 1, N);
struct ggml_tensor * e3 = ggml_reshape_3d(gctx, x_in, x_in->ne[0], 1, N);
x_in = ggml_concat(gctx, b3, e3, 1); // [in_dim, 2, N]
x_in = ggml_reshape_2d(gctx, x_in, x_in->ne[0], T * N);
} }
ggml_set_name(x_in, "sub_input"); ggml_set_name(x_in, "sub_input");
@@ -261,16 +287,17 @@ static bool code_predictor_graph_build(const CodePredictorWeights * cw,
for (int l = 0; l < n_layers; l++) { for (int l = 0; l < n_layers; l++) {
h = code_predictor_layer_forward(gctx, cw, cw->layers[(size_t) l], h, pos_in, mask_in, rows_in, h = code_predictor_layer_forward(gctx, cw, cw->layers[(size_t) l], h, pos_in, mask_in, rows_in,
kv->k[(size_t) l], kv->v[(size_t) l], T, n_kv_pad, use_flash_attn, clamp_fp16, kv->k4[(size_t) l], kv->v4[(size_t) l], T, N, n_kv_pad, use_flash_attn,
gf); clamp_fp16, gf);
} }
struct ggml_tensor * h_final = ggml_rms_norm(gctx, h, cw->rms_norm_eps); struct ggml_tensor * h_final = ggml_rms_norm(gctx, h, cw->rms_norm_eps);
h_final = ggml_mul(gctx, h_final, cw->norm_w); h_final = ggml_mul(gctx, h_final, cw->norm_w);
if (T > 1) { if (T > 1) {
// Last position only: the prefill pays a single lm_head row. // Last position of every slot: columns (T - 1) + n * T, one
h_final = ggml_cont( // strided view so the prefill pays N lm_head rows.
gctx, ggml_view_2d(gctx, h_final, h_final->ne[0], 1, h_final->nb[1], (size_t) (T - 1) * h_final->nb[1])); h_final = ggml_cont(gctx, ggml_view_2d(gctx, h_final, h_final->ne[0], N, (size_t) T * h_final->nb[1],
(size_t) (T - 1) * h_final->nb[1]));
} }
struct ggml_tensor * logits = ggml_mul_mat(gctx, cw->lm_head[(size_t) g_head], h_final); struct ggml_tensor * logits = ggml_mul_mat(gctx, cw->lm_head[(size_t) g_head], h_final);
@@ -286,30 +313,36 @@ static bool code_predictor_graph_build(const CodePredictorWeights * cw,
} }
{ {
std::vector<int32_t> pos((size_t) T); std::vector<int32_t> pos((size_t) T * (size_t) N);
for (int i = 0; i < T; i++) { for (int n = 0; n < N; n++) {
pos[(size_t) i] = n_past + i; for (int t = 0; t < T; t++) {
pos[(size_t) n * (size_t) T + (size_t) t] = n_past + t;
}
} }
ggml_backend_tensor_set(pos_in, pos.data(), 0, (size_t) T * sizeof(int32_t)); ggml_backend_tensor_set(pos_in, pos.data(), 0, pos.size() * sizeof(int32_t));
std::vector<int64_t> rows((size_t) T); std::vector<int64_t> rows((size_t) T * (size_t) N);
for (int i = 0; i < T; i++) { for (int n = 0; n < N; n++) {
rows[(size_t) i] = (int64_t) (n_past + i); for (int t = 0; t < T; t++) {
rows[(size_t) n * (size_t) T + (size_t) t] = (int64_t) (n_past + t);
}
} }
ggml_backend_tensor_set(rows_in, rows.data(), 0, (size_t) T * sizeof(int64_t)); ggml_backend_tensor_set(rows_in, rows.data(), 0, rows.size() * sizeof(int64_t));
} }
{ {
std::vector<ggml_fp16_t> mask((size_t) T * (size_t) n_kv_pad); std::vector<ggml_fp16_t> mask((size_t) n_kv_pad * (size_t) T * (size_t) N);
const ggml_fp16_t zero = ggml_fp32_to_fp16(0.0f); const ggml_fp16_t zero = ggml_fp32_to_fp16(0.0f);
const ggml_fp16_t neg_inf = ggml_fp32_to_fp16(-INFINITY); const ggml_fp16_t neg_inf = ggml_fp32_to_fp16(-INFINITY);
for (size_t i = 0; i < mask.size(); i++) { for (size_t i = 0; i < mask.size(); i++) {
mask[i] = neg_inf; mask[i] = neg_inf;
} }
for (int q = 0; q < T; q++) { for (int n = 0; n < N; n++) {
const int q_pos = n_past + q; for (int q = 0; q < T; q++) {
for (int k = 0; k <= q_pos; k++) { const int q_pos = n_past + q;
mask[(size_t) q * (size_t) n_kv_pad + (size_t) k] = zero; for (int k = 0; k <= q_pos; k++) {
mask[((size_t) n * (size_t) T + (size_t) q) * (size_t) n_kv_pad + (size_t) k] = zero;
}
} }
} }
ggml_backend_tensor_set(mask_in, mask.data(), 0, mask.size() * sizeof(ggml_fp16_t)); ggml_backend_tensor_set(mask_in, mask.data(), 0, mask.size() * sizeof(ggml_fp16_t));
@@ -318,86 +351,105 @@ static bool code_predictor_graph_build(const CodePredictorWeights * cw,
cp->gf = gf; cp->gf = gf;
cp->ids_in = ids_in; cp->ids_in = ids_in;
cp->logits = logits; cp->logits = logits;
cp->N = N;
return true; return true;
} }
// Replay one static predictor graph: upload the code id, run the graph // Replay one static predictor graph: upload the N code ids, run the
// directly on the backend, read the single logits row back. // graph directly on the backend, read the [Vg, N] logits back.
static bool code_predictor_replay(CodePredGraph * cp, static bool code_predictor_replay(CodePredGraph * cp,
ggml_backend_t backend, ggml_backend_t backend,
int32_t code_id, const int32_t * code_ids,
int N,
std::vector<float> * logits_out) { std::vector<float> * logits_out) {
ggml_backend_tensor_set(cp->ids_in, &code_id, 0, sizeof(int32_t)); ggml_backend_tensor_set(cp->ids_in, code_ids, 0, (size_t) N * sizeof(int32_t));
if (ggml_backend_graph_compute(backend, cp->gf) != GGML_STATUS_SUCCESS) { if (ggml_backend_graph_compute(backend, cp->gf) != GGML_STATUS_SUCCESS) {
fprintf(stderr, "[CodePredictor] FATAL: graph compute failed\n"); fprintf(stderr, "[CodePredictor] FATAL: graph compute failed\n");
return false; return false;
} }
logits_out->resize((size_t) cp->logits->ne[0]); logits_out->resize((size_t) cp->logits->ne[0] * (size_t) N);
ggml_backend_tensor_get(cp->logits, logits_out->data(), 0, (size_t) cp->logits->ne[0] * sizeof(float)); ggml_backend_tensor_get(cp->logits, logits_out->data(), 0, logits_out->size() * sizeof(float));
return true; return true;
} }
// Run the predictor for one audio frame over the static graphs. The // Run the predictor for one audio frame over the static graphs, all N
// prefill replay consumes the persistent hidden bridge already written // slots in lockstep. The prefill replay consumes the persistent hidden
// by the talker graph and the sampled c0; the 14 step replays each // bridge columns already written by the talker graph and the sampled
// feed the code sampled just before. Sampling parameters control // c0 of every slot; the 14 step replays each feed the codes sampled
// greedy (temperature <= 0) vs stochastic. subseq_base is the Philox // just before. Per slot sampling parameters control greedy
// subsequence of the c0 sample for this step; the 15 acoustic samples // (temperature <= 0) vs stochastic; seed and subseq_base index each
// consume subseq_base + 1 .. subseq_base + 15. Returns the full vector // slot's own Philox stream, so per slot outputs match a single
// of 16 codes. dump_dir may be NULL. // sequence run bit for bit. subseq_base[i] is the subsequence of slot
// i's c0 sample for this step; its 15 acoustic samples consume
// subseq_base[i] + 1 .. subseq_base[i] + 15. Fills out->codes as
// [N * 16] slot major. dump_dir may be NULL and applies to slot 0.
static bool code_predictor_step(const CodePredictorWeights * cw, static bool code_predictor_step(const CodePredictorWeights * cw,
ggml_backend_t backend, ggml_backend_t backend,
CodePredGraph * prefill_graph, CodePredGraph * prefill_graph,
CodePredGraph * step_graphs, CodePredGraph * step_graphs,
int c0, const int32_t * c0,
float temperature, int N,
int top_k, const float * temperature,
float top_p, const int * top_k,
int64_t seed, const float * top_p,
int64_t subseq_base, const int64_t * seed,
const int64_t * subseq_base,
const char * dump_dir, const char * dump_dir,
CodePredictorOutput * out) { CodePredictorOutput * out) {
const int n_acoustic = cw->num_acoustic_codebooks; const int n_acoustic = cw->num_acoustic_codebooks;
const int n_codes = n_acoustic + 1;
out->codes.assign((size_t) (n_acoustic + 1), 0); out->codes.assign((size_t) N * (size_t) n_codes, 0);
out->codes[0] = c0; for (int i = 0; i < N; i++) {
out->codes[(size_t) i * (size_t) n_codes] = c0[i];
}
std::vector<float> logits; std::vector<float> logits;
if (!code_predictor_replay(prefill_graph, backend, c0, &logits)) { std::vector<int32_t> ids((size_t) N);
if (!code_predictor_replay(prefill_graph, backend, c0, N, &logits)) {
return false; return false;
} }
{ const int V0 = (int) (logits.size() / (size_t) N);
for (int i = 0; i < N; i++) {
float u_g = 0.0f; float u_g = 0.0f;
int cg = sample_top_k_p(logits.data(), (int) logits.size(), temperature, top_k, top_p, 1.0f, nullptr, 0, seed, int cg = sample_top_k_p(logits.data() + (size_t) i * (size_t) V0, V0, temperature[i], top_k[i], top_p[i], 1.0f,
subseq_base + 1, &u_g); nullptr, 0, seed[i], subseq_base[i] + 1, &u_g);
if (cg < 0) { if (cg < 0) {
fprintf(stderr, "[CodePredictor] FATAL: sample returned no candidate at g=0\n"); fprintf(stderr, "[CodePredictor] FATAL: sample returned no candidate at g=0 (slot %d)\n", i);
return false; return false;
} }
out->codes[1] = cg; out->codes[(size_t) i * (size_t) n_codes + 1] = cg;
} }
// Decode loop: 14 single-token replays. At step g (g=1..14) we feed // Decode loop: 14 batched single-token replays. At step g
// the id of the code we just sampled, gathered in graph from the // (g=1..14) every slot feeds the id of the code it just sampled,
// group's private embedding table, and read lm_head[g]. // gathered in graph from the group's private embedding table, and
// reads lm_head[g].
for (int g = 1; g < n_acoustic; g++) { for (int g = 1; g < n_acoustic; g++) {
if (!code_predictor_replay(&step_graphs[(size_t) (g - 1)], backend, out->codes[(size_t) g], &logits)) { for (int i = 0; i < N; i++) {
ids[(size_t) i] = out->codes[(size_t) i * (size_t) n_codes + (size_t) g];
}
if (!code_predictor_replay(&step_graphs[(size_t) (g - 1)], backend, ids.data(), N, &logits)) {
return false; return false;
} }
float u_g = 0.0f; const int Vg = (int) (logits.size() / (size_t) N);
int cg = sample_top_k_p(logits.data(), (int) logits.size(), temperature, top_k, top_p, 1.0f, nullptr, 0, seed, for (int i = 0; i < N; i++) {
subseq_base + 1 + g, &u_g); float u_g = 0.0f;
if (cg < 0) { int cg = sample_top_k_p(logits.data() + (size_t) i * (size_t) Vg, Vg, temperature[i], top_k[i], top_p[i],
fprintf(stderr, "[CodePredictor] FATAL: sample returned no candidate at g=%d\n", g); 1.0f, nullptr, 0, seed[i], subseq_base[i] + 1 + g, &u_g);
return false; if (cg < 0) {
fprintf(stderr, "[CodePredictor] FATAL: sample returned no candidate at g=%d (slot %d)\n", g, i);
return false;
}
out->codes[(size_t) i * (size_t) n_codes + (size_t) (g + 1)] = cg;
} }
out->codes[(size_t) (g + 1)] = cg;
} }
if (dump_dir) { if (dump_dir) {
DebugDumper d; DebugDumper d;
debug_init(&d, dump_dir); debug_init(&d, dump_dir);
std::vector<int32_t> codes32(out->codes.begin(), out->codes.end()); std::vector<int32_t> codes32(out->codes.begin(), out->codes.begin() + n_codes);
int n = (int) codes32.size(); int n = (int) codes32.size();
debug_dump_i32_as_f32(&d, "codes-step0", codes32.data(), &n, 1); debug_dump_i32_as_f32(&d, "codes-step0", codes32.data(), &n, 1);
} }
+10 -7
View File
@@ -1,9 +1,10 @@
#pragma once #pragma once
// code-predictor-graph.h: one static predictor graph, built once at // code-predictor-graph.h: one static predictor graph per flavor and
// load, allocated once, and replayed with a 4 byte code id upload per // batch width, built lazily, allocated once, and replayed with an N
// call. Positions, kv rows, and the causal mask bake in at build time // code id upload per call. Positions, kv rows, and the causal mask
// because the frame cache layout repeats identically: the prefill // bake in at build time because the frame cache layout repeats
// always writes rows 0..1 and step g always appends at row g + 1. // identically for every slot: the prefill always writes rows 0..1 and
// step g always appends at row g + 1.
#include "ggml-alloc.h" #include "ggml-alloc.h"
#include "ggml.h" #include "ggml.h"
@@ -12,8 +13,9 @@ struct CodePredGraph {
struct ggml_context * ctx = nullptr; struct ggml_context * ctx = nullptr;
struct ggml_cgraph * gf = nullptr; struct ggml_cgraph * gf = nullptr;
ggml_gallocr_t galloc = nullptr; ggml_gallocr_t galloc = nullptr;
struct ggml_tensor * ids_in = nullptr; struct ggml_tensor * ids_in = nullptr; // [N] i32, one code id per slot
struct ggml_tensor * logits = nullptr; struct ggml_tensor * logits = nullptr; // [Vg, N] f32
int N = 0; // batch width this build covers
}; };
static void code_predictor_graph_free(CodePredGraph * cp) { static void code_predictor_graph_free(CodePredGraph * cp) {
@@ -28,4 +30,5 @@ static void code_predictor_graph_free(CodePredGraph * cp) {
cp->gf = nullptr; cp->gf = nullptr;
cp->ids_in = nullptr; cp->ids_in = nullptr;
cp->logits = nullptr; cp->logits = nullptr;
cp->N = 0;
} }
+80 -28
View File
@@ -1,16 +1,19 @@
#pragma once #pragma once
// kv-cache.h: persistent per-layer KV cache for the Talker LM and the // kv-cache.h: persistent per-layer KV cache for the Talker LM and the
// Code Predictor. Mirrors the standard llama.cpp ring approach but kept // Code Predictor, batched over n_sets independent sequences. Kept
// minimal: the cache is sized at init for a fixed max sequence length // minimal: the cache is sized at init for a fixed max sequence length
// and never reallocates. Reset just rewinds cur_len to 0. // and set count and never reallocates. Reset just rewinds cur_len.
// //
// Layout per layer, both K and V : // Layout per layer, both K and V :
// ggml_tensor [hd, max_seq_len, n_kv] f32, contiguous on hd // 4D ggml_tensor [hd, max_seq_len, n_kv, n_sets] f32, contiguous on hd
// This matches the layout the attention block already uses for K, so // plus one 3D view [hd, max_seq_len, n_kv] per set for the per-sequence
// the write path is a ggml_cpy of the freshly RoPE'd K into a view of // prefill path and device side set copies. The batched decode graph
// the cache, and the read path is just a view spanning [0, cur_len+T) // views the 4D tensor directly with ne3 = N over sets [0, N).
// on dim 1. V uses the same layout and gets permuted to [T, n_kv, hd] //
// at read time for the value matmul, identical to the prefill path. // The 3D layout matches what the attention block uses for K, so the
// write path is a set_rows of the freshly RoPE'd K into the set view,
// and the read path is a view spanning the padded causal window on
// dim 1. V uses the same layout.
#include "ggml-alloc.h" #include "ggml-alloc.h"
#include "ggml-backend.h" #include "ggml-backend.h"
@@ -26,9 +29,16 @@ struct KVCache {
int n_kv_heads; int n_kv_heads;
int head_dim; int head_dim;
int max_seq_len; int max_seq_len;
int cur_len; int n_sets;
// One pair per layer, both tensors live in `buffer` allocated below. // Write head per set.
std::vector<int> cur_len;
// Per layer 4D tensors, both in `buffer` allocated below.
std::vector<struct ggml_tensor *> k4;
std::vector<struct ggml_tensor *> v4;
// Per set 3D views into the 4D tensors, indexed [set * n_layers + layer].
std::vector<struct ggml_tensor *> k; std::vector<struct ggml_tensor *> k;
std::vector<struct ggml_tensor *> v; std::vector<struct ggml_tensor *> v;
@@ -36,25 +46,38 @@ struct KVCache {
ggml_backend_buffer_t buffer; ggml_backend_buffer_t buffer;
}; };
static struct ggml_tensor * kv_cache_k(const KVCache * kv, int set, int layer) {
return kv->k[(size_t) set * (size_t) kv->n_layers + (size_t) layer];
}
static struct ggml_tensor * kv_cache_v(const KVCache * kv, int set, int layer) {
return kv->v[(size_t) set * (size_t) kv->n_layers + (size_t) layer];
}
// Allocate a fresh KV cache backed by a dedicated buffer on `backend`. // Allocate a fresh KV cache backed by a dedicated buffer on `backend`.
// Tensors are zero initialised (the attention path never reads past // Tensors are zero initialised: the padded attention window reads past
// cur_len so this is mostly cosmetic, but it keeps debug dumps clean). // cur_len and the masked tail must see finite values.
static bool kv_cache_init(KVCache * kv, static bool kv_cache_init(KVCache * kv,
int n_layers, int n_layers,
int n_kv_heads, int n_kv_heads,
int head_dim, int head_dim,
int max_seq_len, int max_seq_len,
int n_sets,
ggml_backend_t backend) { ggml_backend_t backend) {
kv->n_layers = n_layers; kv->n_layers = n_layers;
kv->n_kv_heads = n_kv_heads; kv->n_kv_heads = n_kv_heads;
kv->head_dim = head_dim; kv->head_dim = head_dim;
kv->max_seq_len = max_seq_len; kv->max_seq_len = max_seq_len;
kv->cur_len = 0; kv->n_sets = n_sets;
kv->k.assign((size_t) n_layers, NULL); kv->cur_len.assign((size_t) n_sets, 0);
kv->v.assign((size_t) n_layers, NULL); kv->k4.assign((size_t) n_layers, NULL);
kv->v4.assign((size_t) n_layers, NULL);
kv->k.assign((size_t) n_sets * (size_t) n_layers, NULL);
kv->v.assign((size_t) n_sets * (size_t) n_layers, NULL);
struct ggml_init_params gp = { const size_t n_tensors = (size_t) (2 * n_layers) * (size_t) (1 + n_sets) + 4;
ggml_tensor_overhead() * (size_t) (2 * n_layers + 4), struct ggml_init_params gp = {
ggml_tensor_overhead() * n_tensors,
NULL, NULL,
true, true,
}; };
@@ -65,13 +88,26 @@ static bool kv_cache_init(KVCache * kv,
} }
for (int l = 0; l < n_layers; l++) { for (int l = 0; l < n_layers; l++) {
kv->k[(size_t) l] = ggml_new_tensor_3d(kv->ctx, GGML_TYPE_F32, head_dim, max_seq_len, n_kv_heads); kv->k4[(size_t) l] = ggml_new_tensor_4d(kv->ctx, GGML_TYPE_F32, head_dim, max_seq_len, n_kv_heads, n_sets);
kv->v[(size_t) l] = ggml_new_tensor_3d(kv->ctx, GGML_TYPE_F32, head_dim, max_seq_len, n_kv_heads); kv->v4[(size_t) l] = ggml_new_tensor_4d(kv->ctx, GGML_TYPE_F32, head_dim, max_seq_len, n_kv_heads, n_sets);
char name[64]; char name[64];
snprintf(name, sizeof(name), "kv_k_l%d", l); snprintf(name, sizeof(name), "kv_k_l%d", l);
ggml_set_name(kv->k[(size_t) l], name); ggml_set_name(kv->k4[(size_t) l], name);
snprintf(name, sizeof(name), "kv_v_l%d", l); snprintf(name, sizeof(name), "kv_v_l%d", l);
ggml_set_name(kv->v[(size_t) l], name); ggml_set_name(kv->v4[(size_t) l], name);
// Per set 3D views, created before the buffer allocation so
// ggml_backend_alloc_ctx_tensors runs its view init on them
// (buffer and data resolve against the owning 4D tensor).
struct ggml_tensor * k4 = kv->k4[(size_t) l];
struct ggml_tensor * v4 = kv->v4[(size_t) l];
for (int s = 0; s < n_sets; s++) {
size_t off = (size_t) s * k4->nb[3];
kv->k[(size_t) s * (size_t) n_layers + (size_t) l] =
ggml_view_3d(kv->ctx, k4, head_dim, max_seq_len, n_kv_heads, k4->nb[1], k4->nb[2], off);
kv->v[(size_t) s * (size_t) n_layers + (size_t) l] =
ggml_view_3d(kv->ctx, v4, head_dim, max_seq_len, n_kv_heads, v4->nb[1], v4->nb[2], off);
}
} }
kv->buffer = ggml_backend_alloc_ctx_tensors(kv->ctx, backend); kv->buffer = ggml_backend_alloc_ctx_tensors(kv->ctx, backend);
@@ -85,16 +121,29 @@ static bool kv_cache_init(KVCache * kv,
// Zero-init the buffer so any out of bounds read returns a known value. // Zero-init the buffer so any out of bounds read returns a known value.
ggml_backend_buffer_clear(kv->buffer, 0); ggml_backend_buffer_clear(kv->buffer, 0);
size_t bytes_per_layer = (size_t) head_dim * (size_t) max_seq_len * (size_t) n_kv_heads * sizeof(float); size_t bytes_per_layer =
size_t total_mb = (size_t) (2 * n_layers) * bytes_per_layer / (1024 * 1024); (size_t) head_dim * (size_t) max_seq_len * (size_t) n_kv_heads * (size_t) n_sets * sizeof(float);
fprintf(stderr, "[KVCache] Allocated: %d layers, %d KV heads, head_dim %d, max_seq_len %d -> %zu MB\n", n_layers, size_t total_mb = (size_t) (2 * n_layers) * bytes_per_layer / (1024 * 1024);
n_kv_heads, head_dim, max_seq_len, total_mb); fprintf(stderr, "[KVCache] Allocated: %d layers, %d KV heads, head_dim %d, max_seq_len %d, %d sets -> %zu MB\n",
n_layers, n_kv_heads, head_dim, max_seq_len, n_sets, total_mb);
return true; return true;
} }
// Rewind the cache so the next forward starts a fresh sequence. // Rewind one set so its next forward starts a fresh sequence.
static void kv_cache_reset(KVCache * kv) { static void kv_cache_reset(KVCache * kv, int set) {
kv->cur_len = 0; kv->cur_len[(size_t) set] = 0;
}
// Device side copy of one whole set into another through the
// persistent 3D views. Used by the batch engine to compact the active
// slot range after a retirement: the tail set moves into the freed one
// so the batched decode keeps viewing a consecutive [0, N) span.
static void kv_cache_copy_set(KVCache * kv, int src, int dst) {
for (int l = 0; l < kv->n_layers; l++) {
ggml_backend_tensor_copy(kv_cache_k(kv, src, l), kv_cache_k(kv, dst, l));
ggml_backend_tensor_copy(kv_cache_v(kv, src, l), kv_cache_v(kv, dst, l));
}
kv->cur_len[(size_t) dst] = kv->cur_len[(size_t) src];
} }
static void kv_cache_free(KVCache * kv) { static void kv_cache_free(KVCache * kv) {
@@ -106,6 +155,9 @@ static void kv_cache_free(KVCache * kv) {
ggml_free(kv->ctx); ggml_free(kv->ctx);
kv->ctx = NULL; kv->ctx = NULL;
} }
kv->k4.clear();
kv->v4.clear();
kv->k.clear(); kv->k.clear();
kv->v.clear(); kv->v.clear();
kv->cur_len.clear();
} }
+51 -8
View File
@@ -259,7 +259,7 @@ static bool pipeline_codec_stream_ensure(PipelineCodec * pc) {
} }
if (!kv_cache_init(&pc->stream_kv, pc->transformer.num_layers, pc->transformer.num_kv_heads, if (!kv_cache_init(&pc->stream_kv, pc->transformer.num_layers, pc->transformer.num_kv_heads,
pc->transformer.head_dim, CODEC_STREAM_RING, pc->backend)) { pc->transformer.head_dim, CODEC_STREAM_RING, 1, pc->backend)) {
ggml_backend_buffer_free(pc->stream_buf); ggml_backend_buffer_free(pc->stream_buf);
pc->stream_buf = NULL; pc->stream_buf = NULL;
ggml_free(pc->stream_ctx); ggml_free(pc->stream_ctx);
@@ -281,7 +281,7 @@ bool pipeline_codec_stream_reset(PipelineCodec * pc) {
// the zeroed KV ring stays hidden behind the sliding window mask. // the zeroed KV ring stays hidden behind the sliding window mask.
ggml_backend_buffer_clear(pc->stream_buf, 0); ggml_backend_buffer_clear(pc->stream_buf, 0);
ggml_backend_buffer_clear(pc->stream_kv.buffer, 0); ggml_backend_buffer_clear(pc->stream_kv.buffer, 0);
kv_cache_reset(&pc->stream_kv); kv_cache_reset(&pc->stream_kv, 0);
pc->stream_pos = 0; pc->stream_pos = 0;
return true; return true;
} }
@@ -305,14 +305,17 @@ static bool codec_snap_ensure(PipelineCodec * pc, CodecStateSnap * s) {
if (s->ctx) { if (s->ctx) {
return true; return true;
} }
// Views alias memory the owning tensors already cover, so the
// walker skips them: the mirror holds one duplicate per real
// tensor and the positional pairing in codec_snap_copy holds.
int n = 0; int n = 0;
for (struct ggml_tensor * t = ggml_get_first_tensor(pc->stream_ctx); t; for (struct ggml_tensor * t = ggml_get_first_tensor(pc->stream_ctx); t;
t = ggml_get_next_tensor(pc->stream_ctx, t)) { t = ggml_get_next_tensor(pc->stream_ctx, t)) {
n++; n += t->view_src ? 0 : 1;
} }
for (struct ggml_tensor * t = ggml_get_first_tensor(pc->stream_kv.ctx); t; for (struct ggml_tensor * t = ggml_get_first_tensor(pc->stream_kv.ctx); t;
t = ggml_get_next_tensor(pc->stream_kv.ctx, t)) { t = ggml_get_next_tensor(pc->stream_kv.ctx, t)) {
n++; n += t->view_src ? 0 : 1;
} }
struct ggml_init_params gp = { ggml_tensor_overhead() * (size_t) n, NULL, true }; struct ggml_init_params gp = { ggml_tensor_overhead() * (size_t) n, NULL, true };
s->ctx = ggml_init(gp); s->ctx = ggml_init(gp);
@@ -322,11 +325,15 @@ static bool codec_snap_ensure(PipelineCodec * pc, CodecStateSnap * s) {
} }
for (struct ggml_tensor * t = ggml_get_first_tensor(pc->stream_ctx); t; for (struct ggml_tensor * t = ggml_get_first_tensor(pc->stream_ctx); t;
t = ggml_get_next_tensor(pc->stream_ctx, t)) { t = ggml_get_next_tensor(pc->stream_ctx, t)) {
ggml_dup_tensor(s->ctx, t); if (!t->view_src) {
ggml_dup_tensor(s->ctx, t);
}
} }
for (struct ggml_tensor * t = ggml_get_first_tensor(pc->stream_kv.ctx); t; for (struct ggml_tensor * t = ggml_get_first_tensor(pc->stream_kv.ctx); t;
t = ggml_get_next_tensor(pc->stream_kv.ctx, t)) { t = ggml_get_next_tensor(pc->stream_kv.ctx, t)) {
ggml_dup_tensor(s->ctx, t); if (!t->view_src) {
ggml_dup_tensor(s->ctx, t);
}
} }
s->buf = ggml_backend_alloc_ctx_tensors(s->ctx, pc->backend); s->buf = ggml_backend_alloc_ctx_tensors(s->ctx, pc->backend);
if (!s->buf) { if (!s->buf) {
@@ -343,12 +350,20 @@ static bool codec_snap_ensure(PipelineCodec * pc, CodecStateSnap * s) {
static void codec_snap_copy(PipelineCodec * pc, CodecStateSnap * s, bool save) { static void codec_snap_copy(PipelineCodec * pc, CodecStateSnap * s, bool save) {
struct ggml_tensor * m = ggml_get_first_tensor(s->ctx); struct ggml_tensor * m = ggml_get_first_tensor(s->ctx);
for (struct ggml_tensor * t = ggml_get_first_tensor(pc->stream_ctx); t; for (struct ggml_tensor * t = ggml_get_first_tensor(pc->stream_ctx); t;
t = ggml_get_next_tensor(pc->stream_ctx, t), m = ggml_get_next_tensor(s->ctx, m)) { t = ggml_get_next_tensor(pc->stream_ctx, t)) {
if (t->view_src) {
continue;
}
ggml_backend_tensor_copy(save ? t : m, save ? m : t); ggml_backend_tensor_copy(save ? t : m, save ? m : t);
m = ggml_get_next_tensor(s->ctx, m);
} }
for (struct ggml_tensor * t = ggml_get_first_tensor(pc->stream_kv.ctx); t; for (struct ggml_tensor * t = ggml_get_first_tensor(pc->stream_kv.ctx); t;
t = ggml_get_next_tensor(pc->stream_kv.ctx, t), m = ggml_get_next_tensor(s->ctx, m)) { t = ggml_get_next_tensor(pc->stream_kv.ctx, t)) {
if (t->view_src) {
continue;
}
ggml_backend_tensor_copy(save ? t : m, save ? m : t); ggml_backend_tensor_copy(save ? t : m, save ? m : t);
m = ggml_get_next_tensor(s->ctx, m);
} }
} }
@@ -389,6 +404,34 @@ bool pipeline_codec_stream_snapshot(PipelineCodec * pc, uint64_t key) {
return true; return true;
} }
bool pipeline_codec_stream_save(PipelineCodec * pc, CodecStateSnap * s) {
if (!pc->stream_ready || !codec_snap_ensure(pc, s)) {
return false;
}
codec_snap_copy(pc, s, true);
s->pos = pc->stream_pos;
return true;
}
bool pipeline_codec_stream_load(PipelineCodec * pc, CodecStateSnap * s) {
if (!pc->stream_ready || !s->ctx) {
return false;
}
codec_snap_copy(pc, s, false);
pc->stream_pos = s->pos;
return true;
}
void pipeline_codec_snap_free(CodecStateSnap * s) {
if (s->buf) {
ggml_backend_buffer_free(s->buf);
}
if (s->ctx) {
ggml_free(s->ctx);
}
*s = {};
}
// Build the static stream graph of chunk width T = 1 << cls: the // Build the static stream graph of chunk width T = 1 << cls: the
// same module chain as the T=1 frame graph, every stream state and // same module chain as the T=1 frame graph, every stream state and
// KV ring tensor shared across classes, inputs and intermediates // KV ring tensor shared across classes, inputs and intermediates
+9
View File
@@ -198,6 +198,15 @@ bool pipeline_codec_stream_restore(PipelineCodec * pc, uint64_t key);
// the least recently used slot when all are taken. // the least recently used slot when all are taken.
bool pipeline_codec_stream_snapshot(PipelineCodec * pc, uint64_t key); bool pipeline_codec_stream_snapshot(PipelineCodec * pc, uint64_t key);
// Caller owned stream state mirror, enabling multiple interleaved
// streamed utterances over the single live state: save parks the live
// state (and position cursor) into s, load restores it. The mirror
// allocates lazily on the first save; key and stamp stay unused. Free
// with pipeline_codec_snap_free.
bool pipeline_codec_stream_save(PipelineCodec * pc, CodecStateSnap * s);
bool pipeline_codec_stream_load(PipelineCodec * pc, CodecStateSnap * s);
void pipeline_codec_snap_free(CodecStateSnap * s);
// Encode a 24 kHz mono waveform into RVQ codes. // Encode a 24 kHz mono waveform into RVQ codes.
// audio : [n_samples] f32 mono 24 kHz. Must be a multiple of // audio : [n_samples] f32 mono 24 kHz. Must be a multiple of
// TOKENIZER_HOP_LENGTH (1920); the caller is expected // TOKENIZER_HOP_LENGTH (1920); the caller is expected
+718 -347
View File
File diff suppressed because it is too large Load Diff
+77 -15
View File
@@ -90,6 +90,13 @@ struct PromptCache {
size_t max_prefix_entries; size_t max_prefix_entries;
}; };
// One set of static predictor graphs for a given batch width: the T=2
// prefill plus one T=1 step per acoustic codebook after the first.
struct CodePredGraphSet {
CodePredGraph prefill;
std::vector<CodePredGraph> steps;
};
struct PipelineTTS { struct PipelineTTS {
GGUFModel gguf_talker; GGUFModel gguf_talker;
TalkerWeights talker; TalkerWeights talker;
@@ -108,6 +115,11 @@ struct PipelineTTS {
std::string model_type; std::string model_type;
int num_code_groups; int num_code_groups;
// Batch capacity: number of KV sets, bridge columns and maximum
// concurrent slots the batch engine drives. 1 keeps the exact
// single sequence layout and behavior.
int max_batch;
CodecSpecials codec_specials; CodecSpecials codec_specials;
TextSpecials text_specials; TextSpecials text_specials;
std::vector<LanguageEntry> languages; std::vector<LanguageEntry> languages;
@@ -128,43 +140,47 @@ struct PipelineTTS {
bool use_flash_attn; bool use_flash_attn;
bool clamp_fp16; bool clamp_fp16;
// Persistent KV caches: the talker holds the LM context, the // Persistent KV caches, one set per slot: the talker holds the LM
// predictor holds one frame's 16 sub-steps and gets reset every // contexts, the predictor holds one frame's 16 sub-steps per slot,
// frame in code_predictor_step. // rewritten every frame at baked rows.
KVCache talker_kv; KVCache talker_kv;
KVCache code_predictor_kv; KVCache code_predictor_kv;
// Hidden bridge: the talker last position hidden stays resident on // Hidden bridge: the talker last position hidden of every slot
// device. The talker graph copies it in, the code predictor prefill // stays resident on device as one column of [talker_hidden,
// graph reads it as a leaf, so the AR hot loop never round trips // max_batch] f32. The talker graphs copy their columns in, the code
// the row through the host. [talker_hidden] f32 on `backend`. // predictor prefill graph reads them as a leaf, so the AR hot loop
// never round trips the rows through the host.
struct ggml_context * bridge_ctx; struct ggml_context * bridge_ctx;
ggml_backend_buffer_t bridge_buf; ggml_backend_buffer_t bridge_buf;
struct ggml_tensor * hidden_bridge; struct ggml_tensor * hidden_bridge;
// Persistent graph arena for the talker prefill (T_ctx varies per // Persistent graph arena for the talker prefill (T_ctx varies per
// request, rebuilt through the sched). The talker decode and the // request, rebuilt through the sched). The talker decode and the
// whole predictor run on static graphs instead: the decode keeps // whole predictor run on static batched graphs instead: the decode
// one graph per attention window class built lazily, the predictor // keeps one graph per attention window class built lazily and
// one prefill (T=2) plus one per acoustic step built at load, all // rebuilt when the batch width changes, the predictor one graph
// replayed directly on the backend. // set (prefill T=2 plus one per acoustic step) per batch width
// built lazily, all replayed directly on the backend.
GraphArena talker_arena; GraphArena talker_arena;
std::vector<TalkerDecodeGraph> talker_decode_graphs; // one per 256 step window class, lazy std::vector<TalkerDecodeGraph> talker_decode_graphs; // one per 256 step window class, lazy
CodePredGraph cp_prefill_graph; std::vector<CodePredGraphSet> cp_graphs; // index N - 1, lazy per batch width
std::vector<CodePredGraph> cp_step_graphs;
}; };
// Open the talker GGUF and the codec GGUF, load every module on the // Open the talker GGUF and the codec GGUF, load every module on the
// shared backend. Aborts with a logged error on any missing tensor or // shared backend. Aborts with a logged error on any missing tensor or
// invalid metadata. use_fa is gated on bp.has_gpu inside the load: // invalid metadata. use_fa is gated on bp.has_gpu inside the load:
// CPU only runs always use the manual F32 attention chain. clamp_fp16 // CPU only runs always use the manual F32 attention chain. clamp_fp16
// is forwarded as is. Caller frees with pipeline_tts_free. // is forwarded as is. max_batch sizes the KV sets, the bridge columns
// and the maximum concurrent slots (minimum 1). Caller frees with
// pipeline_tts_free.
bool pipeline_tts_load(PipelineTTS * pt, bool pipeline_tts_load(PipelineTTS * pt,
const char * talker_gguf_path, const char * talker_gguf_path,
const char * codec_gguf_path, const char * codec_gguf_path,
BackendPair bp, BackendPair bp,
bool use_fa, bool use_fa,
bool clamp_fp16); bool clamp_fp16,
int max_batch);
void pipeline_tts_free(PipelineTTS * pt); void pipeline_tts_free(PipelineTTS * pt);
@@ -201,3 +217,49 @@ qt_status pipeline_tts_synthesize(PipelineTTS * pt,
// rate (24000 / TOKENIZER_HOP_LENGTH). Clamps to a // rate (24000 / TOKENIZER_HOP_LENGTH). Clamps to a
// minimum of one frame. // minimum of one frame.
int pipeline_tts_duration_sec_to_tokens(const PipelineTTS * pt, float duration_sec); int pipeline_tts_duration_sec_to_tokens(const PipelineTTS * pt, float duration_sec);
// One synthesis request driven by the batch engine. The caller owns
// params / out for the whole lifetime of the job; status and error
// fill at retirement. error carries the qt_last_error() text captured
// on the thread that ran the engine, so a scheduler on a worker thread
// can replay it into the caller's thread local slot. done is reserved
// for the owner's completion signaling; the engine never touches it.
struct TtsJob {
const struct qt_tts_params * params;
int64_t resolved_seed;
struct qt_audio * out;
qt_status status;
std::string error;
bool done;
};
// Batch engine: drives up to pt->max_batch concurrent synthesis slots
// in lockstep over the batched talker decode and code predictor
// graphs. Slots always occupy KV sets [0, N); a retirement compacts
// the range with one device side set copy so the batched views stay
// consecutive. Single threaded: every call runs on the thread that
// owns the GPU. pipeline_tts_synthesize drives a transient engine
// synchronously for the one request case; the facade scheduler keeps a
// long lived one on its worker thread.
struct TtsEngine;
TtsEngine * tts_engine_new(PipelineTTS * pt, BPETokenizer * tok);
void tts_engine_free(TtsEngine * e);
// Admit one job: prompt assembly, reference handling and talker
// prefill into the next free slot (KV set N). Returns false when the
// admit fails; job->status and job->error are then final and the job
// never occupies a slot. Logs the prefill stall the join imposes on
// the already active slots.
bool tts_engine_admit(TtsEngine * e, TtsJob * job);
// Run one frame for every active slot: batched talker decode for the
// slots past their prefill, per slot c0 sampling, batched code
// predictor, per slot codec streaming, then retirement of finished
// slots (EOS, max_new_tokens, cancel, error) including their buffered
// codec decode or streaming drain. Retired jobs append to *retired
// with status, error and out final.
void tts_engine_step(TtsEngine * e, std::vector<TtsJob *> * retired);
// Number of currently occupied slots.
int tts_engine_active(const TtsEngine * e);
+139 -2
View File
@@ -25,24 +25,49 @@
#include "version.h" #include "version.h"
#include <atomic> #include <atomic>
#include <condition_variable>
#include <cstdarg> #include <cstdarg>
#include <cstdio> #include <cstdio>
#include <cstdlib> #include <cstdlib>
#include <cstring> #include <cstring>
#include <deque>
#include <mutex>
#include <new> #include <new>
#include <random> #include <random>
#include <stdexcept> #include <stdexcept>
#include <string> #include <string>
#include <thread>
#include <vector> #include <vector>
// Internal definition of the opaque handle. C++ types are fine here // Internal definition of the opaque handle. C++ types are fine here
// because nothing in this struct ever crosses the public ABI boundary : // because nothing in this struct ever crosses the public ABI boundary :
// callers only ever see `struct qt_context *`. PipelineTTS already // callers only ever see `struct qt_context *`. PipelineTTS already
// embeds the PipelineCodec, so no separate codec field is needed. // embeds the PipelineCodec, so no separate codec field is needed.
//
// gpu_mu serializes every GPU touching entry: the batch worker holds
// it per engine phase (admit, frame step) so a voice extraction or a
// max_batch == 1 synthesize slips between frames instead of racing the
// backend from another thread.
//
// The scheduler members drive the max_batch > 1 mode: qt_synthesize
// enqueues a TtsJob and blocks on cv_done until the worker retires it.
// The worker owns the long lived TtsEngine, admits queued jobs into
// free slots between frames and steps the batch until both the queue
// and the slots drain.
struct qt_context { struct qt_context {
BackendPair bp; BackendPair bp;
PipelineTTS pt; PipelineTTS pt;
BPETokenizer tok; BPETokenizer tok;
std::mutex gpu_mu;
int max_batch = 1;
std::thread worker;
std::mutex mu;
std::condition_variable cv_work;
std::condition_variable cv_done;
std::deque<TtsJob *> queue;
bool stop = false;
}; };
// Thread-local backing store for qt_last_error(). std::string sized once // Thread-local backing store for qt_last_error(). std::string sized once
@@ -193,6 +218,7 @@ void qt_init_default_params(struct qt_init_params * p) {
p->codec_path = nullptr; p->codec_path = nullptr;
p->use_fa = true; p->use_fa = true;
p->clamp_fp16 = false; p->clamp_fp16 = false;
p->max_batch = 1;
} }
void qt_tts_default_params(struct qt_tts_params * p) { void qt_tts_default_params(struct qt_tts_params * p) {
@@ -236,6 +262,62 @@ int qt_num_codebooks(const struct qt_context * q) {
return q->pt.num_code_groups; return q->pt.num_code_groups;
} }
// Batch worker: owns the long lived TtsEngine and the GPU hot loop.
// Sleeps until work arrives, then admits queued jobs into free slots
// between frames and steps the batch until both the queue and the
// slots drain. Every GPU touching phase runs under gpu_mu so a voice
// extraction or another entry can slip between frames from its own
// thread. Retired jobs get their done flag under mu and a cv_done
// broadcast so the blocked qt_synthesize callers wake.
static void qt_batch_worker(qt_context * q) {
TtsEngine * e = tts_engine_new(&q->pt, &q->tok);
std::vector<TtsJob *> retired;
std::unique_lock<std::mutex> lk(q->mu);
for (;;) {
q->cv_work.wait(lk, [&] { return q->stop || !q->queue.empty(); });
if (q->stop && q->queue.empty()) {
break;
}
while (!q->queue.empty() || tts_engine_active(e) > 0) {
// Admit up to max_batch: joins happen at frame boundaries,
// each one stalls the active slots for one prefill.
while (!q->queue.empty() && tts_engine_active(e) < q->max_batch) {
TtsJob * j = q->queue.front();
q->queue.pop_front();
lk.unlock();
bool admitted;
{
std::lock_guard<std::mutex> gpu(q->gpu_mu);
admitted = tts_engine_admit(e, j);
}
lk.lock();
if (!admitted) {
j->done = true;
q->cv_done.notify_all();
}
}
if (tts_engine_active(e) == 0) {
break;
}
lk.unlock();
retired.clear();
{
std::lock_guard<std::mutex> gpu(q->gpu_mu);
tts_engine_step(e, &retired);
}
lk.lock();
for (TtsJob * j : retired) {
j->done = true;
}
if (!retired.empty()) {
q->cv_done.notify_all();
}
}
}
lk.unlock();
tts_engine_free(e);
}
struct qt_context * qt_init(const struct qt_init_params * params) { struct qt_context * qt_init(const struct qt_init_params * params) {
if (!params || !params->talker_path || !params->codec_path) { if (!params || !params->talker_path || !params->codec_path) {
qt_set_error("qt_init: params, talker_path or codec_path is NULL"); qt_set_error("qt_init: params, talker_path or codec_path is NULL");
@@ -252,10 +334,14 @@ struct qt_context * qt_init(const struct qt_init_params * params) {
qt_log(QT_LOG_INFO, "[Qwen] qwentts.cpp %s", qt_version()); qt_log(QT_LOG_INFO, "[Qwen] qwentts.cpp %s", qt_version());
// ABI v3 tail field: zero init from older callers means 1.
const int max_batch = (params->abi_version >= 3 && params->max_batch > 1) ? params->max_batch : 1;
// new qt_context() value-initialises every field: POD aggregates // new qt_context() value-initialises every field: POD aggregates
// (BackendPair, PipelineTTS) are zero-init, std containers in // (BackendPair, PipelineTTS) are zero-init, std containers in
// BPETokenizer construct empty. // BPETokenizer construct empty.
qt_context * q = new qt_context(); qt_context * q = new qt_context();
q->max_batch = max_batch;
// The load chain runs inside a try block. Any failure deep in the // The load chain runs inside a try block. Any failure deep in the
// GGUF reader, the codec load or the LM weight load throws via // GGUF reader, the codec load or the LM weight load throws via
@@ -269,7 +355,7 @@ struct qt_context * qt_init(const struct qt_init_params * params) {
} }
if (!pipeline_tts_load(&q->pt, params->talker_path, params->codec_path, q->bp, params->use_fa, if (!pipeline_tts_load(&q->pt, params->talker_path, params->codec_path, q->bp, params->use_fa,
params->clamp_fp16)) { params->clamp_fp16, max_batch)) {
qt_throw("qt_init: pipeline_tts_load failed for '%s' / '%s'", params->talker_path, params->codec_path); qt_throw("qt_init: pipeline_tts_load failed for '%s' / '%s'", params->talker_path, params->codec_path);
} }
@@ -292,6 +378,13 @@ struct qt_context * qt_init(const struct qt_init_params * params) {
return nullptr; return nullptr;
} }
// Batch mode: one worker thread owns the long lived engine and the
// GPU hot loop; qt_synthesize enqueues and blocks on completion.
if (q->max_batch > 1) {
q->worker = std::thread(qt_batch_worker, q);
qt_log(QT_LOG_INFO, "[Qwen] Batch scheduler started (max_batch=%d)", q->max_batch);
}
return q; return q;
} }
@@ -299,6 +392,14 @@ void qt_free(struct qt_context * q) {
if (!q) { if (!q) {
return; return;
} }
if (q->worker.joinable()) {
{
std::lock_guard<std::mutex> lk(q->mu);
q->stop = true;
}
q->cv_work.notify_all();
q->worker.join();
}
pipeline_tts_free(&q->pt); pipeline_tts_free(&q->pt);
backend_release(q->bp.backend, q->bp.cpu_backend); backend_release(q->bp.backend, q->bp.cpu_backend);
delete q; delete q;
@@ -353,6 +454,10 @@ enum qt_status qt_extract_voice_ref(struct qt_context * q,
} }
try { try {
// Serialize against the batch worker and any concurrent
// synthesize: the extraction slips between two engine frames.
std::lock_guard<std::mutex> gpu(q->gpu_mu);
// Lazy residency: the first reference audio request pays the // Lazy residency: the first reference audio request pays the
// weight load once, mirroring the qt_synthesize ref_audio path. // weight load once, mirroring the qt_synthesize ref_audio path.
if (!q->pt.spk_enc_loaded) { if (!q->pt.spk_enc_loaded) {
@@ -534,7 +639,39 @@ enum qt_status qt_synthesize(struct qt_context * q, const struct qt_tts_params *
// crosses the extern "C" boundary. // crosses the extern "C" boundary.
try { try {
const int64_t resolved_seed = qt_resolve_seed(params->seed); const int64_t resolved_seed = qt_resolve_seed(params->seed);
return pipeline_tts_synthesize(&q->pt, &q->tok, params, resolved_seed, out);
if (q->max_batch <= 1) {
// Single sequence mode: run synchronously on the calling
// thread under gpu_mu, so concurrent callers serialize FIFO
// and callbacks fire on their own caller's thread.
std::lock_guard<std::mutex> gpu(q->gpu_mu);
return pipeline_tts_synthesize(&q->pt, &q->tok, params, resolved_seed, out);
}
// Batch mode: enqueue and block until the worker retires the
// job. qt_last_error is thread local, so the engine captured
// the worker side message into job.error at retirement; replay
// it into this caller's slot so the errno style contract holds
// across the thread hop.
TtsJob job;
job.params = params;
job.resolved_seed = resolved_seed;
job.out = out;
job.status = QT_STATUS_OK;
job.done = false;
{
std::lock_guard<std::mutex> lk(q->mu);
q->queue.push_back(&job);
}
q->cv_work.notify_all();
{
std::unique_lock<std::mutex> lk(q->mu);
q->cv_done.wait(lk, [&] { return job.done; });
}
if (job.status != QT_STATUS_OK && !job.error.empty()) {
qt_set_error("%s", job.error.c_str());
}
return job.status;
} catch (const std::exception & e) { } catch (const std::exception & e) {
qt_set_error("%s", e.what()); qt_set_error("%s", e.what());
qt_log(QT_LOG_ERROR, "[Qwen] %s", e.what()); qt_log(QT_LOG_ERROR, "[Qwen] %s", e.what());
+20 -2
View File
@@ -57,7 +57,7 @@ extern "C" {
// git short hash + commit date string returned by qt_version(); for // git short hash + commit date string returned by qt_version(); for
// binding compat checks, QT_ABI_VERSION is the only number that // binding compat checks, QT_ABI_VERSION is the only number that
// matters. // matters.
#define QT_ABI_VERSION 2 #define QT_ABI_VERSION 3
// Returns a static string of the form "<git-hash> (<date>)" identifying // Returns a static string of the form "<git-hash> (<date>)" identifying
// the exact commit this binary was built from. Safe to call from any // the exact commit this binary was built from. Safe to call from any
@@ -121,6 +121,19 @@ struct qt_init_params {
const char * codec_path; const char * codec_path;
bool use_fa; bool use_fa;
bool clamp_fp16; bool clamp_fp16;
// ABI v3. Maximum number of concurrent synthesis requests batched
// on the GPU. 0 and 1 select the single sequence behavior (zero
// init from older callers keeps the previous semantics); values
// above 1 size the KV cache sets accordingly and start an internal
// worker thread that coalesces concurrent qt_synthesize calls into
// batched decode steps, queueing FIFO beyond max_batch. With
// max_batch > 1 the on_chunk and cancel callbacks of every request
// are invoked from that worker thread, not from the calling
// thread; callbacks must be safe to run there and must not call
// back into the qwen_* API. qt_synthesize itself stays blocking
// and thread safe in both modes.
int max_batch;
}; };
// Initialise to the standard defaults: both paths NULL (caller must set // Initialise to the standard defaults: both paths NULL (caller must set
@@ -289,7 +302,12 @@ struct qt_tts_params {
// the streaming pipeline: audio chunks emit through on_chunk and // the streaming pipeline: audio chunks emit through on_chunk and
// `out` stays empty on success. on_chunk NULL keeps the buffered // `out` stays empty on success. on_chunk NULL keeps the buffered
// path. The last chunk on EOS or max_new flushes whatever frames // path. The last chunk on EOS or max_new flushes whatever frames
// remain. // remain. With qt_init_params.max_batch > 1 the callback runs on
// the internal batch worker thread, not the qt_synthesize caller
// thread: it must be safe there, must not call back into the
// qwen_* API, and a blocking body stalls every batched request, so
// hand the samples to the consumer thread through a queue instead
// of blocking.
qt_audio_chunk_cb on_chunk; qt_audio_chunk_cb on_chunk;
void * on_chunk_user_data; void * on_chunk_user_data;
+20 -13
View File
@@ -1,10 +1,11 @@
#pragma once #pragma once
// talker-decode-graph.h: one static talker decode graph, built once // talker-decode-graph.h: one static batched talker decode graph, built
// per attention window class (the kv window rounds up in 256 step // per (attention window class, batch width N) and replayed directly on
// spans) and replayed directly on the backend. Every input re-uploads // the backend. The kv window rounds up in 256 step spans; every input
// before each replay: positions, kv row, and mask carry the moving // re-uploads before each replay: positions, kv rows, and masks carry
// n_past, the frame code ids and the overlay row carry the previous // the moving per-slot n_past, the frame code ids and the overlay rows
// frame, so nothing bakes and the allocator contract holds. // carry each slot's previous frame, so nothing bakes and the allocator
// contract holds. The graph always covers KV sets [0, N).
#include "ggml-alloc.h" #include "ggml-alloc.h"
#include "ggml.h" #include "ggml.h"
@@ -15,14 +16,17 @@ struct TalkerDecodeGraph {
struct ggml_context * ctx = nullptr; struct ggml_context * ctx = nullptr;
struct ggml_cgraph * gf = nullptr; struct ggml_cgraph * gf = nullptr;
ggml_gallocr_t galloc = nullptr; ggml_gallocr_t galloc = nullptr;
struct ggml_tensor * ids_in = nullptr; // [1 + n_acoustic] i32 struct ggml_tensor * ids_in = nullptr; // [(1 + n_acoustic) * N] i32, group major: g * N + slot
struct ggml_tensor * overlay = nullptr; // [hidden, 1] f32 struct ggml_tensor * overlay = nullptr; // [hidden, N] f32
struct ggml_tensor * pos_in = nullptr; // [1] i32 struct ggml_tensor * pos_in = nullptr; // [N] i32
struct ggml_tensor * rows_in = nullptr; // [1] i64 struct ggml_tensor * rows_in = nullptr; // [1, 1, N] i64
struct ggml_tensor * mask_in = nullptr; // [n_kv_pad, 1] f16 struct ggml_tensor * mask_in = nullptr; // [n_kv_pad, 1, 1, N] f16
struct ggml_tensor * logits = nullptr; // [vocab, 1] f32 struct ggml_tensor * logits = nullptr; // [vocab, N] f32
std::vector<ggml_fp16_t> mask; // [n_kv_pad] f16 std::vector<ggml_fp16_t> mask; // [n_kv_pad * N] f16
std::vector<int32_t> pos_data; // [N] host staging
std::vector<int64_t> rows_data; // [N] host staging
int n_kv_pad = 0; // window class width, 0 marks an empty slot int n_kv_pad = 0; // window class width, 0 marks an empty slot
int N = 0; // batch width this build covers
}; };
static void talker_decode_graph_free(TalkerDecodeGraph * tg) { static void talker_decode_graph_free(TalkerDecodeGraph * tg) {
@@ -42,5 +46,8 @@ static void talker_decode_graph_free(TalkerDecodeGraph * tg) {
tg->mask_in = nullptr; tg->mask_in = nullptr;
tg->logits = nullptr; tg->logits = nullptr;
tg->mask.clear(); tg->mask.clear();
tg->pos_data.clear();
tg->rows_data.clear();
tg->n_kv_pad = 0; tg->n_kv_pad = 0;
tg->N = 0;
} }
+242 -80
View File
@@ -253,16 +253,17 @@ static struct ggml_tensor * talker_layer_forward(struct ggml_context * ctx,
// Prefill core: builds the graph in the caller owned arena, allocates // Prefill core: builds the graph in the caller owned arena, allocates
// through the sched, uploads the raw embedding, runs it and pulls out // through the sched, uploads the raw embedding, runs it and pulls out
// the last position logits. T tokens are appended to the cache // the last position logits. T tokens are appended to KV set `kv_set`
// starting at n_past. The last position hidden copies in graph into // starting at n_past. The last position hidden copies in graph into
// the caller owned persistent hidden_bridge tensor the code predictor // column `slot` of the caller owned persistent hidden_bridge tensor
// prefill reads on device; the host copy in out->hidden_last fills // [hidden, max_batch] the code predictor prefill reads on device; the
// only under read_hidden_host. When n_past == 0 and dump_dir is set, // host copy in out->hidden_last fills only under read_hidden_host.
// the bisect taps fire. use_fa / clamp_fp16 are forwarded as is to // When n_past == 0 and dump_dir is set, the bisect taps fire. use_fa /
// every layer. The decode hot path runs on the static graphs below // clamp_fp16 are forwarded as is to every layer. The decode hot path
// instead. // runs on the static batched graphs below instead.
static bool talker_forward_core(const TalkerWeights * tw, static bool talker_forward_core(const TalkerWeights * tw,
KVCache * kv, KVCache * kv,
int kv_set,
ggml_backend_sched_t sched, ggml_backend_sched_t sched,
GraphArena * arena, GraphArena * arena,
struct ggml_tensor * hidden_bridge, struct ggml_tensor * hidden_bridge,
@@ -317,8 +318,9 @@ static bool talker_forward_core(const TalkerWeights * tw,
struct ggml_tensor * h = x_in; struct ggml_tensor * h = x_in;
std::vector<struct ggml_tensor *> taps(TALKER_N_BISECT_LAYERS, NULL); std::vector<struct ggml_tensor *> taps(TALKER_N_BISECT_LAYERS, NULL);
for (int l = 0; l < n_layers; l++) { for (int l = 0; l < n_layers; l++) {
h = talker_layer_forward(gctx, tw, tw->layers[(size_t) l], h, pos_in, mask_in, rows_in, kv->k[(size_t) l], h = talker_layer_forward(gctx, tw, tw->layers[(size_t) l], h, pos_in, mask_in, rows_in,
kv->v[(size_t) l], T, n_kv_pad, use_flash_attn, clamp_fp16, gf); kv_cache_k(kv, kv_set, l), kv_cache_v(kv, kv_set, l), T, n_kv_pad, use_flash_attn,
clamp_fp16, gf);
if (record_taps && talker_is_bisect_layer(l)) { if (record_taps && talker_is_bisect_layer(l)) {
for (int i = 0; i < TALKER_N_BISECT_LAYERS; i++) { for (int i = 0; i < TALKER_N_BISECT_LAYERS; i++) {
if (TALKER_BISECT_LAYERS[i] == l) { if (TALKER_BISECT_LAYERS[i] == l) {
@@ -340,12 +342,14 @@ static bool talker_forward_core(const TalkerWeights * tw,
ggml_set_output(h_final); ggml_set_output(h_final);
} }
// Bridge: the last position hidden copies on device into the // Bridge: the last position hidden copies on device into column
// persistent tensor. Constant destination address across steps, so // `kv_set` of the persistent [hidden, max_batch] tensor. Constant
// the decode graph topology stays replayable. // destination address across steps, so the decode graph topology
// stays replayable.
struct ggml_tensor * h_last = struct ggml_tensor * h_last =
ggml_view_1d(gctx, h_final, hidden, (size_t) (T - 1) * (size_t) hidden * sizeof(float)); ggml_view_1d(gctx, h_final, hidden, (size_t) (T - 1) * (size_t) hidden * sizeof(float));
struct ggml_tensor * bridge_cpy = ggml_cpy(gctx, h_last, hidden_bridge); struct ggml_tensor * bridge_col = ggml_view_1d(gctx, hidden_bridge, hidden, (size_t) kv_set * hidden_bridge->nb[1]);
struct ggml_tensor * bridge_cpy = ggml_cpy(gctx, h_last, bridge_col);
// codec_head: [hidden, vocab]. ggml_mul_mat returns [vocab, T]. // codec_head: [hidden, vocab]. ggml_mul_mat returns [vocab, T].
struct ggml_tensor * logits = ggml_mul_mat(gctx, tw->codec_head_w, h_final); struct ggml_tensor * logits = ggml_mul_mat(gctx, tw->codec_head_w, h_final);
@@ -455,14 +459,15 @@ static bool talker_forward_core(const TalkerWeights * tw,
// Advance the cache write head. The graph already executed the cpy // Advance the cache write head. The graph already executed the cpy
// nodes so positions [n_past, n_past + T) are now populated. The // nodes so positions [n_past, n_past + T) are now populated. The
// arena and the sched allocation persist into the next forward. // arena and the sched allocation persist into the next forward.
kv->cur_len = T_full; kv->cur_len[(size_t) kv_set] = T_full;
return true; return true;
} }
// Prefill: reset the cache and write T_ctx positions in one shot. // Prefill: reset KV set `kv_set` and write T_ctx positions in one
// input_embed is [T, hidden] f32 row-major. dump_dir may be NULL. // shot. input_embed is [T, hidden] f32 row-major. dump_dir may be NULL.
static bool talker_forward_prefill(const TalkerWeights * tw, static bool talker_forward_prefill(const TalkerWeights * tw,
KVCache * kv, KVCache * kv,
int kv_set,
ggml_backend_sched_t sched, ggml_backend_sched_t sched,
GraphArena * arena, GraphArena * arena,
struct ggml_tensor * hidden_bridge, struct ggml_tensor * hidden_bridge,
@@ -472,21 +477,147 @@ static bool talker_forward_prefill(const TalkerWeights * tw,
bool clamp_fp16, bool clamp_fp16,
const char * dump_dir, const char * dump_dir,
TalkerForwardOutput * out) { TalkerForwardOutput * out) {
kv_cache_reset(kv); kv_cache_reset(kv, kv_set);
if (T > kv->max_seq_len) { if (T > kv->max_seq_len) {
fprintf(stderr, "[TalkerForward] FATAL: prefill T=%d exceeds cache max_seq_len=%d\n", T, kv->max_seq_len); fprintf(stderr, "[TalkerForward] FATAL: prefill T=%d exceeds cache max_seq_len=%d\n", T, kv->max_seq_len);
return false; return false;
} }
return talker_forward_core(tw, kv, sched, arena, hidden_bridge, input_embed, T, 0, use_flash_attn, clamp_fp16, return talker_forward_core(tw, kv, kv_set, sched, arena, hidden_bridge, input_embed, T, 0, use_flash_attn,
dump_dir != NULL, dump_dir, out); clamp_fp16, dump_dir != NULL, dump_dir, out);
} }
// Build one static decode graph for the given attention window. The // Build the batched per-layer block over KV sets [0, N), one fresh
// previous frame code ids gather and sum in graph with the overlay // token per set. Projections, norms, RoPE and the MLP run on [*, N]
// row on top; the last position hidden copies into hidden_bridge and // with per column math identical to the single sequence layer. The
// the codec head logits are the single output row. Positions, kv row, // attention is per set: fresh K,V write into the 4D cache at the rows
// and mask stay plain inputs re-uploaded before every replay since // carried by kv_rows via set_rows, the read views the padded causal
// n_past moves each step. // window with ne3 = N, and the mask kills everything past each set's
// own context. Returns the layer output [hidden, N].
static struct ggml_tensor * talker_layer_forward_batch(struct ggml_context * ctx,
const TalkerWeights * tw,
const TalkerLayer & layer,
struct ggml_tensor * x,
struct ggml_tensor * positions,
struct ggml_tensor * mask,
struct ggml_tensor * kv_rows,
struct ggml_tensor * k4,
struct ggml_tensor * v4,
int N,
int n_kv_pad,
bool use_flash_attn,
bool clamp_fp16,
struct ggml_cgraph * gf) {
const int n_q_heads = tw->num_attention_heads;
const int n_kv = tw->num_key_value_heads;
const int hd = tw->head_dim;
const float eps = tw->rms_norm_eps;
// Pre-norm
struct ggml_tensor * h = ggml_rms_norm(ctx, x, eps);
h = ggml_mul(ctx, h, layer.input_norm_w);
// Q/K/V projections, batched over the N columns
struct ggml_tensor * q = ggml_mul_mat(ctx, layer.attn.q_proj_w, h); // [n_q_heads*hd, N]
struct ggml_tensor * k = ggml_mul_mat(ctx, layer.attn.k_proj_w, h); // [n_kv*hd, N]
struct ggml_tensor * v = ggml_mul_mat(ctx, layer.attn.v_proj_w, h); // [n_kv*hd, N]
q = ggml_reshape_3d(ctx, q, hd, n_q_heads, N); // [hd, n_q_heads, N]
k = ggml_reshape_3d(ctx, k, hd, n_kv, N);
v = ggml_reshape_3d(ctx, v, hd, n_kv, N);
// Per-head QK-norm: RMS over hd, then multiply by [hd] gain.
q = ggml_rms_norm(ctx, q, eps);
q = ggml_mul(ctx, q, layer.attn.q_norm_w);
k = ggml_rms_norm(ctx, k, eps);
k = ggml_mul(ctx, k, layer.attn.k_norm_w);
// RoPE NEOX: positions [N] map to dim 2 of [hd, heads, N], one
// absolute position per set.
q = ggml_rope_ext(ctx, q, positions, NULL, hd, GGML_ROPE_TYPE_NEOX, 0, tw->rope_theta, 1.0f, 0.0f, 1.0f, 0.0f,
0.0f);
k = ggml_rope_ext(ctx, k, positions, NULL, hd, GGML_ROPE_TYPE_NEOX, 0, tw->rope_theta, 1.0f, 0.0f, 1.0f, 0.0f,
0.0f);
// Write the fresh K,V into the 4D cache via set_rows over the N
// consecutive sets: src reshapes [hd, n_kv, N] to [hd, 1, n_kv, N]
// (same layout) and kv_rows [1, 1, N] carries one destination row
// per set, broadcast across the n_kv head dim. The row ids travel
// as data so every step keeps an identical topology and the
// captured CUDA graph replays without an update.
struct ggml_tensor * k_sets = ggml_view_4d(ctx, k4, hd, k4->ne[1], n_kv, N, k4->nb[1], k4->nb[2], k4->nb[3], 0);
struct ggml_tensor * v_sets = ggml_view_4d(ctx, v4, hd, v4->ne[1], n_kv, N, v4->nb[1], v4->nb[2], v4->nb[3], 0);
struct ggml_tensor * k_new = ggml_reshape_4d(ctx, k, hd, 1, n_kv, N);
struct ggml_tensor * v_new = ggml_reshape_4d(ctx, v, hd, 1, n_kv, N);
ggml_build_forward_expand(gf, ggml_set_rows(ctx, k_sets, k_new, kv_rows));
ggml_build_forward_expand(gf, ggml_set_rows(ctx, v_sets, v_new, kv_rows));
// Batched read: [hd, n_kv_pad, n_kv, N] views of the 4D cache. The
// window covers every set's causal context and rounds up so the
// shape stays constant across 256 consecutive decode steps; the
// mask carries neg inf past each set's own context and the cache
// buffer is zero initialized, so the padded tail contributes
// nothing.
struct ggml_tensor * k_batch = ggml_view_4d(ctx, k4, hd, n_kv_pad, n_kv, N, k4->nb[1], k4->nb[2], k4->nb[3], 0);
struct ggml_tensor * v_batch = ggml_view_4d(ctx, v4, hd, n_kv_pad, n_kv, N, v4->nb[1], v4->nb[2], v4->nb[3], 0);
// Q: [hd, n_q_heads, N] -> [hd, 1, n_q_heads, N] (n_batch=1, ne3=N
// for batched flash_attn).
struct ggml_tensor * q4 = ggml_reshape_4d(ctx, q, hd, 1, n_q_heads, N);
// Clamp V before attention, same rationale as the prefill block:
// sub Ampere CUDA tensor cores accumulate in FP16 and a V overflow
// corrupts every subsequent attention.
if (clamp_fp16) {
v_batch = ggml_clamp(ctx, v_batch, -65504.0f, 65504.0f);
}
// Attention: fused flash kernel (set_prec(F32) promotes the
// accumulator) or the manual F32 chain, which broadcasts its
// mul_mat over dims 2 and 3 so the same helper covers 4D.
float scale = 1.0f / sqrtf((float) hd);
struct ggml_tensor * attn;
if (use_flash_attn) {
attn = ggml_flash_attn_ext(ctx, q4, k_batch, v_batch, mask, scale, 0.0f, 0.0f);
ggml_flash_attn_ext_set_prec(attn, GGML_PREC_F32);
} else {
attn = talker_attn_f32(ctx, q4, k_batch, v_batch, mask, scale);
}
// Output [hd, n_q_heads, 1, N] -> [n_q_heads*hd, N], flatten heads
// for o_proj.
attn = ggml_reshape_2d(ctx, attn, n_q_heads * hd, N);
struct ggml_tensor * o = ggml_mul_mat(ctx, layer.attn.o_proj_w, attn);
x = ggml_add(ctx, x, o);
if (clamp_fp16) {
x = ggml_clamp(ctx, x, -65504.0f, 65504.0f);
}
// MLP block: pre-norm + SwiGLU + residual
struct ggml_tensor * h2 = ggml_rms_norm(ctx, x, eps);
h2 = ggml_mul(ctx, h2, layer.post_attn_norm_w);
struct ggml_tensor * gate = ggml_mul_mat(ctx, layer.mlp.gate_proj_w, h2);
struct ggml_tensor * up = ggml_mul_mat(ctx, layer.mlp.up_proj_w, h2);
gate = ggml_silu(ctx, gate);
struct ggml_tensor * gu = ggml_mul(ctx, gate, up);
struct ggml_tensor * mlp = ggml_mul_mat(ctx, layer.mlp.down_proj_w, gu);
x = ggml_add(ctx, x, mlp);
if (clamp_fp16) {
x = ggml_clamp(ctx, x, -65504.0f, 65504.0f);
}
return x;
}
// Build one static batched decode graph for the given attention window
// and batch width N. Each slot's previous frame code ids gather and
// sum in graph with its overlay row on top; the batch of last position
// hiddens copies into columns [0, N) of hidden_bridge and the codec
// head logits [vocab, N] are the single output. Positions, kv rows,
// and masks stay plain inputs re-uploaded before every replay since
// each slot's n_past moves every step.
static bool talker_decode_graph_build(const TalkerWeights * tw, static bool talker_decode_graph_build(const TalkerWeights * tw,
KVCache * kv, KVCache * kv,
ggml_backend_t backend, ggml_backend_t backend,
@@ -494,6 +625,7 @@ static bool talker_decode_graph_build(const TalkerWeights * tw,
struct ggml_tensor * const * acoustic_embd, struct ggml_tensor * const * acoustic_embd,
int n_acoustic, int n_acoustic,
int n_kv_pad, int n_kv_pad,
int N,
bool use_flash_attn, bool use_flash_attn,
bool clamp_fp16, bool clamp_fp16,
TalkerDecodeGraph * tg) { TalkerDecodeGraph * tg) {
@@ -511,26 +643,29 @@ static bool talker_decode_graph_build(const TalkerWeights * tw,
} }
struct ggml_context * gctx = tg->ctx; struct ggml_context * gctx = tg->ctx;
struct ggml_tensor * pos_in = ggml_new_tensor_1d(gctx, GGML_TYPE_I32, 1); struct ggml_tensor * pos_in = ggml_new_tensor_1d(gctx, GGML_TYPE_I32, N);
struct ggml_tensor * mask_in = ggml_new_tensor_2d(gctx, GGML_TYPE_F16, n_kv_pad, 1); struct ggml_tensor * mask_in = ggml_new_tensor_4d(gctx, GGML_TYPE_F16, n_kv_pad, 1, 1, N);
struct ggml_tensor * rows_in = ggml_new_tensor_1d(gctx, GGML_TYPE_I64, 1); struct ggml_tensor * rows_in = ggml_new_tensor_3d(gctx, GGML_TYPE_I64, 1, 1, N);
struct ggml_tensor * ids_in = ggml_new_tensor_1d(gctx, GGML_TYPE_I32, 1 + n_acoustic); struct ggml_tensor * ids_in = ggml_new_tensor_1d(gctx, GGML_TYPE_I32, (1 + n_acoustic) * N);
struct ggml_tensor * overlay = ggml_new_tensor_2d(gctx, GGML_TYPE_F32, hidden, 1); struct ggml_tensor * overlay = ggml_new_tensor_2d(gctx, GGML_TYPE_F32, hidden, N);
ggml_set_name(pos_in, "positions"); ggml_set_name(pos_in, "positions");
ggml_set_name(mask_in, "causal_mask"); ggml_set_name(mask_in, "causal_mask");
ggml_set_name(rows_in, "kv_rows"); ggml_set_name(rows_in, "kv_rows");
ggml_set_name(ids_in, "frame_code_ids"); ggml_set_name(ids_in, "frame_code_ids");
ggml_set_name(overlay, "overlay_row"); ggml_set_name(overlay, "overlay_rows");
ggml_set_input(pos_in); ggml_set_input(pos_in);
ggml_set_input(mask_in); ggml_set_input(mask_in);
ggml_set_input(rows_in); ggml_set_input(rows_in);
ggml_set_input(ids_in); ggml_set_input(ids_in);
ggml_set_input(overlay); ggml_set_input(overlay);
struct ggml_tensor * id0 = ggml_view_1d(gctx, ids_in, 1, 0); // ids_in is group major: entry g * N + slot. Each group slices a
// contiguous [N] view and gathers its own table, summed across the
// 16 codebooks, plus the per slot overlay row.
struct ggml_tensor * id0 = ggml_view_1d(gctx, ids_in, N, 0);
struct ggml_tensor * x_in = ggml_get_rows(gctx, tw->codec_embedding, id0); struct ggml_tensor * x_in = ggml_get_rows(gctx, tw->codec_embedding, id0);
for (int g = 0; g < n_acoustic; g++) { for (int g = 0; g < n_acoustic; g++) {
struct ggml_tensor * idg = ggml_view_1d(gctx, ids_in, 1, (size_t) (g + 1) * sizeof(int32_t)); struct ggml_tensor * idg = ggml_view_1d(gctx, ids_in, N, (size_t) (g + 1) * (size_t) N * sizeof(int32_t));
x_in = ggml_add(gctx, x_in, ggml_get_rows(gctx, acoustic_embd[g], idg)); x_in = ggml_add(gctx, x_in, ggml_get_rows(gctx, acoustic_embd[g], idg));
} }
x_in = ggml_add(gctx, x_in, overlay); x_in = ggml_add(gctx, x_in, overlay);
@@ -540,16 +675,17 @@ static bool talker_decode_graph_build(const TalkerWeights * tw,
struct ggml_tensor * h = x_in; struct ggml_tensor * h = x_in;
for (int l = 0; l < n_layers; l++) { for (int l = 0; l < n_layers; l++) {
h = talker_layer_forward(gctx, tw, tw->layers[(size_t) l], h, pos_in, mask_in, rows_in, kv->k[(size_t) l], h = talker_layer_forward_batch(gctx, tw, tw->layers[(size_t) l], h, pos_in, mask_in, rows_in,
kv->v[(size_t) l], 1, n_kv_pad, use_flash_attn, clamp_fp16, gf); kv->k4[(size_t) l], kv->v4[(size_t) l], N, n_kv_pad, use_flash_attn, clamp_fp16,
gf);
} }
struct ggml_tensor * h_final = ggml_rms_norm(gctx, h, tw->rms_norm_eps); struct ggml_tensor * h_final = ggml_rms_norm(gctx, h, tw->rms_norm_eps);
h_final = ggml_mul(gctx, h_final, tw->norm_w); h_final = ggml_mul(gctx, h_final, tw->norm_w);
ggml_set_name(h_final, "hidden_final"); ggml_set_name(h_final, "hidden_final");
struct ggml_tensor * h_last = ggml_view_1d(gctx, h_final, hidden, 0); struct ggml_tensor * bridge_cols = ggml_view_2d(gctx, hidden_bridge, hidden, N, hidden_bridge->nb[1], 0);
struct ggml_tensor * bridge_cpy = ggml_cpy(gctx, h_last, hidden_bridge); struct ggml_tensor * bridge_cpy = ggml_cpy(gctx, h_final, bridge_cols);
struct ggml_tensor * logits = ggml_mul_mat(gctx, tw->codec_head_w, h_final); struct ggml_tensor * logits = ggml_mul_mat(gctx, tw->codec_head_w, h_final);
ggml_set_name(logits, "logits"); ggml_set_name(logits, "logits");
@@ -571,61 +707,83 @@ static bool talker_decode_graph_build(const TalkerWeights * tw,
tg->rows_in = rows_in; tg->rows_in = rows_in;
tg->mask_in = mask_in; tg->mask_in = mask_in;
tg->logits = logits; tg->logits = logits;
tg->mask.resize((size_t) n_kv_pad); tg->mask.resize((size_t) n_kv_pad * (size_t) N);
tg->pos_data.resize((size_t) N);
tg->rows_data.resize((size_t) N);
tg->n_kv_pad = n_kv_pad; tg->n_kv_pad = n_kv_pad;
tg->N = N;
return true; return true;
} }
// Decode: append one position from the previous frame's codes over the // Batched decode: append one position per KV set over the static graph
// static graph of the current window class, built lazily on the first // of the current window class and batch width, built lazily on the
// step entering the span. frame_ids holds [c0, c1..c15], acoustic_embd // first step entering the (span, N) pair. frame_ids holds the previous
// the 15 group tables owned by the code predictor, overlay the trailing // frame codes of every slot, group major: entry g * N + slot.
// text / pad row summed on top. Reads positions [0, kv->cur_len + 1); // acoustic_embd holds the 15 group tables owned by the code predictor,
// caller ensures kv->cur_len + 1 <= kv->max_seq_len holds by cache // overlays the [hidden, N] trailing text / pad rows summed on top. The
// sizing. read_hidden_host pulls the bridge back for dump paths. // window covers max over sets of (cur_len + 1); caller ensures every
static bool talker_forward_decode(const TalkerWeights * tw, // cur_len + 1 <= max_seq_len holds by cache sizing. Logits fill
KVCache * kv, // out->logits_last as [vocab, N] column blocks; read_hidden_host pulls
ggml_backend_t backend, // the bridge back as [hidden, N] for dump paths.
TalkerDecodeGraph * graphs, static bool talker_forward_decode(const TalkerWeights * tw,
struct ggml_tensor * hidden_bridge, KVCache * kv,
const int32_t * frame_ids, ggml_backend_t backend,
struct ggml_tensor * const * acoustic_embd, std::vector<TalkerDecodeGraph> & graphs,
int n_acoustic, struct ggml_tensor * hidden_bridge,
const float * overlay, const int32_t * frame_ids,
bool use_flash_attn, struct ggml_tensor * const * acoustic_embd,
bool clamp_fp16, int n_acoustic,
bool read_hidden_host, const float * overlays,
TalkerForwardOutput * out) { int N,
if (kv->cur_len + 1 > kv->max_seq_len) { bool use_flash_attn,
fprintf(stderr, "[TalkerForward] FATAL: decode would overflow cache (%d + 1 > %d)\n", kv->cur_len, bool clamp_fp16,
kv->max_seq_len); bool read_hidden_host,
return false; TalkerForwardOutput * out) {
int max_len = 0;
for (int i = 0; i < N; i++) {
const int len = kv->cur_len[(size_t) i] + 1;
if (len > kv->max_seq_len) {
fprintf(stderr, "[TalkerForward] FATAL: decode would overflow cache (%d > %d, set %d)\n", len,
kv->max_seq_len, i);
return false;
}
if (len > max_len) {
max_len = len;
}
} }
const int n_past = kv->cur_len; const int kv_pad_raw = (int) GGML_PAD(max_len, 256);
const int kv_pad_raw = (int) GGML_PAD(n_past + 1, 256);
const int n_kv_pad = kv_pad_raw < kv->max_seq_len ? kv_pad_raw : kv->max_seq_len; const int n_kv_pad = kv_pad_raw < kv->max_seq_len ? kv_pad_raw : kv->max_seq_len;
TalkerDecodeGraph * tg = &graphs[(n_kv_pad + 255) / 256 - 1]; TalkerDecodeGraph * tg = &graphs[(size_t) ((n_kv_pad + 255) / 256 - 1)];
if (!tg->ctx && !talker_decode_graph_build(tw, kv, backend, hidden_bridge, acoustic_embd, n_acoustic, n_kv_pad, if (tg->ctx && tg->N != N) {
talker_decode_graph_free(tg);
}
if (!tg->ctx && !talker_decode_graph_build(tw, kv, backend, hidden_bridge, acoustic_embd, n_acoustic, n_kv_pad, N,
use_flash_attn, clamp_fp16, tg)) { use_flash_attn, clamp_fp16, tg)) {
return false; return false;
} }
ggml_backend_tensor_set(tg->ids_in, frame_ids, 0, (size_t) (1 + n_acoustic) * sizeof(int32_t)); ggml_backend_tensor_set(tg->ids_in, frame_ids, 0, (size_t) (1 + n_acoustic) * (size_t) N * sizeof(int32_t));
ggml_backend_tensor_set(tg->overlay, overlay, 0, (size_t) tw->hidden_size * sizeof(float)); ggml_backend_tensor_set(tg->overlay, overlays, 0, (size_t) tw->hidden_size * (size_t) N * sizeof(float));
const int32_t pos = n_past; for (int i = 0; i < N; i++) {
ggml_backend_tensor_set(tg->pos_in, &pos, 0, sizeof(int32_t)); tg->pos_data[(size_t) i] = kv->cur_len[(size_t) i];
const int64_t row = (int64_t) n_past; tg->rows_data[(size_t) i] = (int64_t) kv->cur_len[(size_t) i];
ggml_backend_tensor_set(tg->rows_in, &row, 0, sizeof(int64_t)); }
ggml_backend_tensor_set(tg->pos_in, tg->pos_data.data(), 0, (size_t) N * sizeof(int32_t));
ggml_backend_tensor_set(tg->rows_in, tg->rows_data.data(), 0, (size_t) N * sizeof(int64_t));
// Causal mask: keys [0, n_past] carry 0, the padded tail neg inf. // Causal masks: for slot i keys [0, cur_len[i]] carry 0, the rest
// of the padded window neg inf.
{ {
std::vector<ggml_fp16_t> & mask = tg->mask; std::vector<ggml_fp16_t> & mask = tg->mask;
const ggml_fp16_t zero = ggml_fp32_to_fp16(0.0f); const ggml_fp16_t zero = ggml_fp32_to_fp16(0.0f);
const ggml_fp16_t neg_inf = ggml_fp32_to_fp16(-INFINITY); const ggml_fp16_t neg_inf = ggml_fp32_to_fp16(-INFINITY);
for (size_t i = 0; i < mask.size(); i++) { for (int i = 0; i < N; i++) {
mask[i] = (int) i <= n_past ? zero : neg_inf; const int n_past = kv->cur_len[(size_t) i];
for (int j = 0; j < n_kv_pad; j++) {
mask[(size_t) i * (size_t) n_kv_pad + (size_t) j] = j <= n_past ? zero : neg_inf;
}
} }
ggml_backend_tensor_set(tg->mask_in, mask.data(), 0, mask.size() * sizeof(ggml_fp16_t)); ggml_backend_tensor_set(tg->mask_in, mask.data(), 0, mask.size() * sizeof(ggml_fp16_t));
} }
@@ -637,13 +795,17 @@ static bool talker_forward_decode(const TalkerWeights * tw,
out->hidden = tw->hidden_size; out->hidden = tw->hidden_size;
out->vocab = tw->vocab_size; out->vocab = tw->vocab_size;
out->logits_last.assign((size_t) tw->vocab_size, 0.0f); out->logits_last.assign((size_t) tw->vocab_size * (size_t) N, 0.0f);
ggml_backend_tensor_get(tg->logits, out->logits_last.data(), 0, (size_t) tw->vocab_size * sizeof(float)); ggml_backend_tensor_get(tg->logits, out->logits_last.data(), 0,
(size_t) tw->vocab_size * (size_t) N * sizeof(float));
if (read_hidden_host) { if (read_hidden_host) {
out->hidden_last.assign((size_t) tw->hidden_size, 0.0f); out->hidden_last.assign((size_t) tw->hidden_size * (size_t) N, 0.0f);
ggml_backend_tensor_get(hidden_bridge, out->hidden_last.data(), 0, (size_t) tw->hidden_size * sizeof(float)); ggml_backend_tensor_get(hidden_bridge, out->hidden_last.data(), 0,
(size_t) tw->hidden_size * (size_t) N * sizeof(float));
} }
kv->cur_len = n_past + 1; for (int i = 0; i < N; i++) {
kv->cur_len[(size_t) i]++;
}
return true; return true;
} }
+2 -2
View File
@@ -393,8 +393,8 @@ static struct ggml_tensor * tok_trans_forward_stream(struct ggml_context *
h = ggml_add(ctx, h, tr->input_proj_b); h = ggml_add(ctx, h, tr->input_proj_b);
for (int l = 0; l < tr->num_layers; l++) { for (int l = 0; l < tr->num_layers; l++) {
h = tok_trans_layer_forward_stream(ctx, gf, tr, tr->layers[l], h, positions, mask, kv_rows, kv->k[(size_t) l], h = tok_trans_layer_forward_stream(ctx, gf, tr, tr->layers[l], h, positions, mask, kv_rows,
kv->v[(size_t) l], T, ring); kv_cache_k(kv, 0, l), kv_cache_v(kv, 0, l), T, ring);
} }
h = ggml_rms_norm(ctx, h, tr->rms_norm_eps); h = ggml_rms_norm(ctx, h, tr->rms_norm_eps);
+73 -23
View File
@@ -26,14 +26,18 @@
#include "audio-io.h" #include "audio-io.h"
#include "yyjson.h" #include "yyjson.h"
#include <atomic>
#include <cfloat> #include <cfloat>
#include <cmath> #include <cmath>
#include <condition_variable>
#include <csignal> #include <csignal>
#include <cstdint> #include <cstdint>
#include <cstdio> #include <cstdio>
#include <functional> #include <functional>
#include <memory>
#include <mutex> #include <mutex>
#include <string> #include <string>
#include <thread>
#include <vector> #include <vector>
// One synthesis request parsed from the OAI JSON body. // One synthesis request parsed from the OAI JSON body.
@@ -95,8 +99,9 @@ struct server_config {
int port = 8080; int port = 8080;
}; };
// Single GPU context : synthesis is serialised FIFO across connections. // Concurrency lives behind the ABI: qt_synthesize is thread safe and
static std::mutex g_synth_mutex; // batches concurrent requests when the context was initialised with
// max_batch > 1, so connection threads call the backend directly.
static httplib::Server * g_svr = nullptr; static httplib::Server * g_svr = nullptr;
static void tts_on_signal(int) { static void tts_on_signal(int) {
@@ -252,18 +257,18 @@ static void tts_handle_speech(const tts_backend & be, const httplib::Request & h
} }
if (req.format == "wav") { if (req.format == "wav") {
// One-shot : collect the whole utterance, then emit a RIFF file. // One-shot : collect the whole utterance, then emit a RIFF
// file. The backend call blocks this connection thread; in
// batch mode the chunks arrive from the ABI worker thread while
// this thread waits inside qt_synthesize, so the plain append
// stays single producer and the blocking return synchronizes.
std::vector<float> buf; std::vector<float> buf;
tts_sink sink = [&buf](const float * s, int n) { tts_sink sink = [&buf](const float * s, int n) {
buf.insert(buf.end(), s, s + n); buf.insert(buf.end(), s, s + n);
return true; return true;
}; };
std::string synth_err; std::string synth_err;
int rc; int rc = be.synthesize(req, sink, synth_err);
{
std::lock_guard<std::mutex> lock(g_synth_mutex);
rc = be.synthesize(req, sink, synth_err);
}
if (rc != 0) { if (rc != 0) {
tts_json_error(res, tts_status_to_http(rc), "server_error", tts_json_error(res, tts_status_to_http(rc), "server_error",
synth_err.empty() ? "synthesis failed" : synth_err.c_str()); synth_err.empty() ? "synthesis failed" : synth_err.c_str());
@@ -274,26 +279,71 @@ static void tts_handle_speech(const tts_backend & be, const httplib::Request & h
return; return;
} }
// Streaming : run synthesis inside the chunked provider on the connection // Streaming : the synthesis runs on its own thread pushing s16le
// thread, pushing s16le frames as the codec produces them. A failed // bytes into a queue; the connection thread drains the queue into
// sink.write means the client disconnected, which aborts generation and // the chunked sink. The decoupling keeps a slow client from
// frees the GPU instead of finishing a stream nobody reads. // stalling the ABI worker (head of line blocking across the whole
res.set_header("Cache-Control", "no-cache"); // batch), and a client disconnect flips client_gone so the next
res.set_header("X-Accel-Buffering", "no"); // chunk callback aborts generation and frees the GPU instead of
res.set_chunked_content_provider("audio/pcm", [&be, req](size_t, httplib::DataSink & sink) mutable -> bool { // finishing a stream nobody reads. Backpressure is the utterance
tts_sink push = [&sink](const float * s, int n) { // itself: pending grows at most to the full PCM of one synthesis.
struct stream_state {
std::mutex mu;
std::condition_variable cv;
std::string pending;
bool done = false;
std::atomic<bool> client_gone{ false };
std::thread th;
};
auto st = std::make_shared<stream_state>();
st->th = std::thread([&be, req, st]() {
tts_sink push = [st](const float * s, int n) {
if (st->client_gone.load(std::memory_order_acquire)) {
return false;
}
std::string bytes; std::string bytes;
tts_append_s16le(bytes, s, n); tts_append_s16le(bytes, s, n);
return sink.write(bytes.data(), bytes.size()); std::lock_guard<std::mutex> lk(st->mu);
st->pending += bytes;
st->cv.notify_all();
return true;
}; };
std::string synth_err; std::string synth_err;
{ be.synthesize(req, push, synth_err);
std::lock_guard<std::mutex> lock(g_synth_mutex); std::lock_guard<std::mutex> lk(st->mu);
be.synthesize(req, push, synth_err); st->done = true;
} st->cv.notify_all();
sink.done();
return true;
}); });
res.set_header("Cache-Control", "no-cache");
res.set_header("X-Accel-Buffering", "no");
res.set_chunked_content_provider(
"audio/pcm",
[st](size_t, httplib::DataSink & sink) -> bool {
std::string chunk;
{
std::unique_lock<std::mutex> lk(st->mu);
st->cv.wait(lk, [&] { return st->done || !st->pending.empty(); });
chunk.swap(st->pending);
}
if (!chunk.empty()) {
if (!sink.write(chunk.data(), chunk.size())) {
st->client_gone.store(true, std::memory_order_release);
return false;
}
return true;
}
sink.done();
return true;
},
[st](bool) {
st->client_gone.store(true, std::memory_order_release);
if (st->th.joinable()) {
st->th.join();
}
});
} }
// Decode standard base64 (with optional padding) into out. Returns // Decode standard base64 (with optional padding) into out. Returns
+44 -19
View File
@@ -30,8 +30,11 @@ struct voice_entry {
}; };
// Registered voices, name keyed. Every access happens under // Registered voices, name keyed. Every access happens under
// g_synth_mutex: registration touches the GPU through the extraction // g_voices_mutex; the GPU side of a registration is serialized inside
// path and lookups run inside the already serialized synthesize. // the ABI (qt_extract_voice_ref slips between batch frames), and the
// synthesize lookup copies the latents out so a concurrent replace or
// delete never frees buffers a running synthesis still reads.
static std::mutex g_voices_mutex;
static std::unordered_map<std::string, voice_entry> g_voices; static std::unordered_map<std::string, voice_entry> g_voices;
static void print_usage(const char * prog) { static void print_usage(const char * prog) {
@@ -46,6 +49,7 @@ static void print_usage(const char * prog) {
" --host <ip> Listen address (default: 127.0.0.1)\n" " --host <ip> Listen address (default: 127.0.0.1)\n"
" --port <n> Listen port (default: 8080)\n" " --port <n> Listen port (default: 8080)\n"
" --lang <name> Language label (default: auto)\n" " --lang <name> Language label (default: auto)\n"
" --max-batch <n> Concurrent requests batched on the GPU (default: 1)\n"
" --no-fa Disable flash attention\n" " --no-fa Disable flash attention\n"
" --clamp-fp16 Clamp hidden states to FP16 range\n", " --clamp-fp16 Clamp hidden states to FP16 range\n",
prog); prog);
@@ -66,6 +70,7 @@ int main(int argc, char ** argv) {
server_config cfg; server_config cfg;
bool use_fa = true; bool use_fa = true;
bool clamp_fp16 = false; bool clamp_fp16 = false;
int max_batch = 1;
for (int i = 1; i < argc; i++) { for (int i = 1; i < argc; i++) {
const char * arg = argv[i]; const char * arg = argv[i];
@@ -85,6 +90,8 @@ int main(int argc, char ** argv) {
use_fa = false; use_fa = false;
} else if (!std::strcmp(arg, "--clamp-fp16")) { } else if (!std::strcmp(arg, "--clamp-fp16")) {
clamp_fp16 = true; clamp_fp16 = true;
} else if (!std::strcmp(arg, "--max-batch") && i + 1 < argc) {
max_batch = std::atoi(argv[++i]);
} else if (!std::strcmp(arg, "--help") || !std::strcmp(arg, "-h")) { } else if (!std::strcmp(arg, "--help") || !std::strcmp(arg, "-h")) {
print_usage(argv[0]); print_usage(argv[0]);
return 0; return 0;
@@ -106,6 +113,7 @@ int main(int argc, char ** argv) {
iparams.codec_path = codec_path; iparams.codec_path = codec_path;
iparams.use_fa = use_fa; iparams.use_fa = use_fa;
iparams.clamp_fp16 = clamp_fp16; iparams.clamp_fp16 = clamp_fp16;
iparams.max_batch = max_batch;
struct qt_context * q = qt_init(&iparams); struct qt_context * q = qt_init(&iparams);
if (!q) { if (!q) {
@@ -136,11 +144,7 @@ int main(int argc, char ** argv) {
err = "cannot decode the WAV payload"; err = "cannot decode the WAV payload";
return false; return false;
} }
enum qt_status rc; enum qt_status rc = qt_extract_voice_ref(q, pcm, T, &entry.ref);
{
std::lock_guard<std::mutex> lock(g_synth_mutex);
rc = qt_extract_voice_ref(q, pcm, T, &entry.ref);
}
free(pcm); free(pcm);
if (rc != QT_STATUS_OK) { if (rc != QT_STATUS_OK) {
err = qt_last_error(); err = qt_last_error();
@@ -167,7 +171,7 @@ int main(int argc, char ** argv) {
std::memcpy(entry.ref.ref_codes, codes.data(), codes.size() * sizeof(int32_t)); std::memcpy(entry.ref.ref_codes, codes.data(), codes.size() * sizeof(int32_t));
} }
std::lock_guard<std::mutex> lock(g_synth_mutex); std::lock_guard<std::mutex> lock(g_voices_mutex);
auto it = g_voices.find(up.name); auto it = g_voices.find(up.name);
if (it != g_voices.end()) { if (it != g_voices.end()) {
qt_voice_ref_free(&it->second.ref); qt_voice_ref_free(&it->second.ref);
@@ -180,7 +184,7 @@ int main(int argc, char ** argv) {
}; };
be.remove_voice = [](const std::string & name) -> bool { be.remove_voice = [](const std::string & name) -> bool {
std::lock_guard<std::mutex> lock(g_synth_mutex); std::lock_guard<std::mutex> lock(g_voices_mutex);
auto it = g_voices.find(name); auto it = g_voices.find(name);
if (it == g_voices.end()) { if (it == g_voices.end()) {
return false; return false;
@@ -191,7 +195,7 @@ int main(int argc, char ** argv) {
}; };
be.registered_voices = []() -> std::vector<std::string> { be.registered_voices = []() -> std::vector<std::string> {
std::lock_guard<std::mutex> lock(g_synth_mutex); std::lock_guard<std::mutex> lock(g_voices_mutex);
std::vector<std::string> names; std::vector<std::string> names;
names.reserve(g_voices.size()); names.reserve(g_voices.size());
for (const auto & kv : g_voices) { for (const auto & kv : g_voices) {
@@ -214,15 +218,36 @@ int main(int argc, char ** argv) {
p.text = req.input.c_str(); p.text = req.input.c_str();
p.lang = lang.c_str(); p.lang = lang.c_str();
auto vit = req.voice.empty() ? g_voices.end() : g_voices.find(req.voice); // Copy the registered voice latents out under the lock: the
if (vit != g_voices.end()) { // synthesis may run for seconds while another connection
const voice_entry & v = vit->second; // replaces or deletes the entry.
p.ref_spk_emb = v.ref.ref_spk_emb; std::vector<float> voice_spk;
p.ref_spk_dim = v.ref.ref_spk_dim; std::vector<int32_t> voice_codes;
if (!v.ref_text.empty() && v.ref.ref_codes) { std::string voice_ref_text;
p.ref_codes = v.ref.ref_codes; int voice_ref_T = 0;
p.ref_T = v.ref.ref_T; bool have_voice = false;
p.ref_text = v.ref_text.c_str(); if (!req.voice.empty()) {
std::lock_guard<std::mutex> lock(g_voices_mutex);
auto vit = g_voices.find(req.voice);
if (vit != g_voices.end()) {
const voice_entry & v = vit->second;
have_voice = true;
voice_spk.assign(v.ref.ref_spk_emb, v.ref.ref_spk_emb + v.ref.ref_spk_dim);
if (!v.ref_text.empty() && v.ref.ref_codes) {
voice_codes.assign(v.ref.ref_codes,
v.ref.ref_codes + (size_t) v.ref.num_codebooks * (size_t) v.ref.ref_T);
voice_ref_T = v.ref.ref_T;
voice_ref_text = v.ref_text;
}
}
}
if (have_voice) {
p.ref_spk_emb = voice_spk.data();
p.ref_spk_dim = (int) voice_spk.size();
if (!voice_codes.empty()) {
p.ref_codes = voice_codes.data();
p.ref_T = voice_ref_T;
p.ref_text = voice_ref_text.c_str();
} }
} else if (!req.voice.empty() && qt_n_speakers(q) > 0) { } else if (!req.voice.empty() && qt_n_speakers(q) > 0) {
p.speaker = req.voice.c_str(); p.speaker = req.voice.c_str();