diff --git a/docs/community_models/audio8_tts.md b/docs/community_models/audio8_tts.md index 845dbcb87..a3fa60f8f 100644 --- a/docs/community_models/audio8_tts.md +++ b/docs/community_models/audio8_tts.md @@ -4,7 +4,7 @@ Audio8 TTS Preview 0.6B (Qwen backbone) and 0.1B (Falcon-H1 hybrid Mamba2+attent S2 Pro: a slow semantic transformer generates speech semantics, a fast codebook transformer expands each semantic step into a full codec frame, and a neural codec renders 44.1 kHz audio. The native path executes all three stages directly on ggml with no Python dependency. -> **Status 2026-08-29:** `0.6B` Qwen is fully native, CPU-validated via SenseVoice ASR round-trip (`The quick brown fox…`, `你好,欢迎使用audio8。`, `Artificial intelligence…`). `0.1B` Falcon-H1 is weight-complete and builds natively (GGUF `slow.embed_tokens` + `24× mamba/attention` + `semantic_output`), but the slow AR forward is a documented stub pending the Mamba2 port — see `docs/FALCON_H1_0.1B_PORT_PLAN.md` and `src/community_models/audio8_tts/ar.cpp:861` `TODO(Falcon-H1)`. +> **Status 2026-09-01:** `0.6B` Qwen is fully native, CPU-validated via SenseVoice ASR round-trip (`The quick brown fox…`, `你好,欢迎使用audio8。`, `Artificial intelligence…`). `0.1B` Falcon-H1 now has a **stateful native slow-AR forward** (Mamba2 + hybrid GQA attention, branch `feat/audio8-tts-falcon-h1-mamba2`), but it is **not yet correct for synthesis** — two open issues (logits argmax mismatch vs transformers reference, and recurrent SSM state blow-up on long sequences). See [audio8_tts_falcon_h1_status.md](audio8_tts_falcon_h1_status.md) for details and next steps. | Field | Value | |---|---| diff --git a/docs/community_models/audio8_tts_falcon_h1_status.md b/docs/community_models/audio8_tts_falcon_h1_status.md new file mode 100644 index 000000000..10c9a54c0 --- /dev/null +++ b/docs/community_models/audio8_tts_falcon_h1_status.md @@ -0,0 +1,116 @@ +# Audio8 TTS 0.1B (Falcon-H1) — Port Status + +> **Status 2026-09-01 (resolved):** The Falcon-H1 slow-AR path is a **stateful +> native implementation** (Mamba2 + hybrid GQA attention) and now matches the +> transformers reference: first-frame semantic argmax = 2732, per-step argmax +> parity over the whole prompt, ASR round-trip of synthesized "你好" returns +> "你好。", and long generation (~600 positions) is numerically stable. The +> 0.6B Qwen path is unaffected and fully functional. + +## What has been done + +Branch `feat/audio8-tts-falcon-h1-mamba2`. + +- `src/community_models/audio8_tts/ar.cpp` + - `FalconH1StepState` + `init_falcon_step_state`: per-layer conv/SSM states + and attention KV cache. + - `falcon_forward_step`: stateful single-token forward — + `RMSNorm -> (Mamba2 || GQA attention) -> residual -> gated FFN -> RMSNorm + -> semantic_output`, matching `transformers.models.falcon_h1`. + - `build_falcon_embedding_step`: `(text_emb + codebook_sum) * + embedding_multiplier` (multiplier applies to the whole sum). + - `generate()` falcon branch: token-by-token prefill + generation. +- `src/community_models/audio8_tts/falcon_kv_cache.h` + - `append_falcon_kv_token`: host KV-cache append with per-head re-stride + (see "Root causes" below). Covered by `audio8_tts_falcon_kv_cache_test`. +- `external/ggml/src/ggml-metal/ggml-metal.metal` + - Fixed `kernel_ssm_scan_f32` reduction: the old + `simd_sum(shared_sums[sgitg*NW + tiisg])` read garbage columns when + `sgptg < NW` (happens for `d_state=64` with `n_t=1`). Replaced with an + explicit loop summing `shared_sums[(i2+sgitg)*NW + g]` over `g < sgptg`. + +## Resolved Issue 1 — logits argmax mismatch vs transformers + +**Symptom (before):** first generated semantic code was wrong (argmax 3620 +instead of 2732; ASR round-trip said "三星" instead of "你好"). + +**Root causes (three, all fixed):** + +1. **conv1d kernel flip was wrong.** `ggml_ssm_conv` computes + `y[c] = sum_k w[k,c]*window[k,c]` with `window[0]` the OLDEST frame — + the same orientation as HF (`nn.Conv1d` prefill and the cached + `torch.sum(conv_states * w, dim=-1)` decode are both cross-correlation). + The GGUF tensor `[d_conv,1,conv_dim]` is the HF `[conv_dim,1,d_conv]` + weight with unchanged flat bytes, i.e. already in the layout ssm_conv + wants. An earlier "fix" that flipped the kernel taps corrupted the x/B/C + split every step. Fix: feed the kernel unflipped (`load_falcon_layer`, + `conv1d_kernel`). +2. **Unprotected host read-back of intermediate tensors.** `sx` (conv + window) and `k_r`/`v` (fresh K/V) are graph intermediates whose buffers + gallocr reuses; reading them back without `ggml_set_output` returned + garbage and corrupted conv state / KV cache every step. Fix: + `ggml_set_output` on exactly those three tensors per layer (pinning ~300 + tensors corrupts the whole graph — pin only what is read back). +3. **KV cache head-stride bug (the decisive one).** The host cache used the + *current* sequence length as the per-head stride while appending only the + new token: at step 1 the new head-0 token was written over token 0's + head-1 block, so every head past the first read corrupted context from + the second token on (head 0 was always correct, which masked the bug). + Fix: `append_falcon_kv_token` re-lays existing entries into the new + stride before appending. Regression test: + `tests/unittests/test_audio8_tts_falcon_kv_cache.cpp` (fails with the old + algorithm at the second append, passes after). + +**Verification:** bf16 GGUF vs f32 HF reference (`transformers==4.57.6`, +recurrent path forced for every token): per-layer conv/SSM/K/V states match +within bf16 rounding over the full 23-token prompt; argmax matches at every +prompt position except one knife-edge tie (ref top-2 margin 0.03 vs bf16 +logit noise 0.19). First frame: argmax 2732 (logit 26.524 vs ref 26.557). +End-to-end: synthesized "你好" transcribes back as "你好。" (Qwen3-ASR), on +both CPU and Metal backends. + +## Resolved Issue 2 — recurrent SSM state blow-up on long sequences + +**Symptom (before):** ~180 tokens in, per-layer states reached 1e15..1e18, +then logits went to zero / NaN. + +**Root cause:** not the recurrent scan itself — the corrupted KV cache +(Issue 1, cause 3) fed garbage attention output into the residual stream, +which drove `x`/`dt` of the Mamba2 branch into regime where the state +exploded. The f32 HF reference running the same recurrent math stays bounded +(states ~270 over the prompt), which ruled out the "inherent weak-decay" +theory previously recorded here. + +**Verification:** 140-character text → 525 generated frames (position 605): +max per-layer SSM state ≈ 1.1e3, zero NaN, clean EOS, and the audio +transcribes back to the input text verbatim. No chunked-scan or dt-clamp +mitigation was needed; HF's recurrent fallback does not apply the +`time_step_min/max` clamp either (`time_step_limit` is hardcoded +`(0.0, inf)` in `modeling_falcon_h1.py`). + +## Notes for the next agent + +- The parity harness used for the fix (an HF golden-dump script forcing the + recurrent path token-by-token, plus a differ) was session tooling and is not + committed. To rebuild it: run `modeling_arktts` under `transformers==4.57.6` + with each `layer.mamba.forward` replaced by the `use_precomputed_states` + recurrent branch, dump `cache.conv_states/ssm_states/key_cache` per step, + and diff against `state.ssm_states/conv_states/k_cache/v_cache` read back in + `falcon_forward_step` (same flat layouts: ssm `[s + 64d + 2048h]`, conv rows + oldest→newest, kv `d + 64*(t + T*h)`). +- `A_log` / `D` / `dt_bias` / `conv1d.weight` must load with + `assets::TensorStorageType::F32` (GGUF stores them quantized; `Native` + keeps the quantized type and `ggml_backend_tensor_get` then reads out of + bounds). +- The GGUF layout for `conv1d.weight` is `[d_conv, 1, conv_dim]`, which is + the HF `[conv_dim, 1, d_conv]` weight with unchanged flat bytes. Feed it to + `ggml_ssm_conv` **unflipped** (see Resolved Issue 1, cause 1). +- `ggml_set_output` on a graph intermediate pins its buffer so host + read-back is safe — but mass-pinning hundreds of tensors corrupts the + whole graph (all-zero logits). Pin only the tensors actually read back. +- Reference environment: `uv venv` + `uv pip install torch "transformers>=4.57,<5"`; + `transformers>=5` renames `FalconHybridMambaAttentionDynamicCache` and + breaks the 0.1B remote code. +- Known remaining gap: the fast-AR codebooks during generation mean the C++ + rollout cannot be compared token-by-token against a reference that feeds + zero codebook rows; parity was established over the prompt + first frame. diff --git a/external/ggml/include/ggml.h b/external/ggml/include/ggml.h index 12ad92b45..0098c3185 100644 --- a/external/ggml/include/ggml.h +++ b/external/ggml/include/ggml.h @@ -604,6 +604,9 @@ extern "C" { GGML_OP_MUL_MAT_ADD, GGML_OP_MUL_MAT_ADD_RELU, GGML_OP_IM2COL_ASYM, + // audio8_tts codec per-tap accumulation and fused snake (audio8 PR). + GGML_OP_MUL_MAT_ACC, + GGML_OP_SNAKE_1D, GGML_OP_COUNT, }; @@ -1447,6 +1450,21 @@ extern "C" { struct ggml_context * ctx, struct ggml_tensor * a, struct ggml_tensor * b); + + // accumulate matrix multiplication in-place: acc += a * b + // result is a view of acc (which must have the shape of a * b), so the + // accumulation lands directly in acc's memory without a separate add pass + GGML_API struct ggml_tensor * ggml_mul_mat_acc( + struct ggml_context * ctx, + struct ggml_tensor * a, + struct ggml_tensor * b, + struct ggml_tensor * acc); + + // fused snake activation: dst = a + sin(a * alpha)^2 / alpha, alpha broadcast per channel + GGML_API struct ggml_tensor * ggml_snake_1d( + struct ggml_context * ctx, + struct ggml_tensor * a, + struct ggml_tensor * alpha); GGML_API struct ggml_tensor * ggml_mul_mat_pack4( struct ggml_context * ctx, diff --git a/external/ggml/src/ggml-cpu/ggml-cpu.c b/external/ggml/src/ggml-cpu/ggml-cpu.c index 3b102aa0c..1d6ac3373 100644 --- a/external/ggml/src/ggml-cpu/ggml-cpu.c +++ b/external/ggml/src/ggml-cpu/ggml-cpu.c @@ -1708,6 +1708,93 @@ static void ggml_compute_forward_mul_mat_id( } } +// reference implementation of the accumulate-in-place matmul (dst aliases src[2]): +// dst += src0 * src1. The op is only exercised on Metal; this plain single-threaded +// loop exists so the CPU backend stays correct if a graph containing it is ever run. +static void ggml_compute_forward_mul_mat_acc( + const struct ggml_compute_params * params, + struct ggml_tensor * dst) { + const struct ggml_tensor * src0 = dst->src[0]; // a [K, M] + const struct ggml_tensor * src1 = dst->src[1]; // b [K, N] + + GGML_ASSERT(dst->type == GGML_TYPE_F32); + GGML_ASSERT(src0->type == GGML_TYPE_F32); + GGML_ASSERT(src1->type == GGML_TYPE_F32); + + if (params->ith != 0) { + return; + } + + const int64_t K = src0->ne[0]; + const int64_t M = src0->ne[1]; + const int64_t N = src1->ne[1]; + + GGML_ASSERT(src1->ne[0] == K); + GGML_ASSERT(dst->ne[0] == M && dst->ne[1] == N); + GGML_ASSERT(src0->ne[2] == 1 && src0->ne[3] == 1); + GGML_ASSERT(src1->ne[2] == 1 && src1->ne[3] == 1); + GGML_ASSERT(dst->ne[2] == 1 && dst->ne[3] == 1); + + const char * A = (const char *) src0->data; + const char * B = (const char *) src1->data; + char * C = (char *) dst->data; + + for (int64_t n = 0; n < N; ++n) { + for (int64_t m = 0; m < M; ++m) { + float sum = 0.0f; + for (int64_t k = 0; k < K; ++k) { + const float av = *(const float *) (A + k*src0->nb[0] + m*src0->nb[1]); + const float bv = *(const float *) (B + k*src1->nb[0] + n*src1->nb[1]); + sum += av * bv; + } + float * cv = (float *) (C + m*dst->nb[0] + n*dst->nb[1]); + *cv += sum; + } + } +} + +// ggml_compute_forward_snake_1d +// +// fused snake activation: y = x + sin(x * alpha)^2 / alpha, alpha broadcast per channel. +// naive single-threaded reference so the CPU backend stays correct if a graph containing +// this op is ever run there. +static void ggml_compute_forward_snake_1d( + const struct ggml_compute_params * params, + struct ggml_tensor * dst) { + const struct ggml_tensor * src0 = dst->src[0]; // x [C, T], channels on the fast axis + const struct ggml_tensor * src1 = dst->src[1]; // alpha [C, 1] + + if (params->ith != 0) { + return; + } + + GGML_ASSERT(dst->type == GGML_TYPE_F32); + GGML_ASSERT(src0->type == GGML_TYPE_F32); + GGML_ASSERT(src1->type == GGML_TYPE_F32); + GGML_ASSERT(ggml_is_contiguous(src0)); + GGML_ASSERT(ggml_is_contiguous(src1)); + GGML_ASSERT(ggml_is_contiguous(dst)); + GGML_ASSERT(src0->ne[2] == 1 && src0->ne[3] == 1); + GGML_ASSERT(src1->ne[0] == src0->ne[0] && src1->ne[1] == 1); + + const int64_t nc = src0->ne[0]; + const int64_t nt = src0->ne[1]; + + const float * x = (const float *) src0->data; + const float * a = (const float *) src1->data; + float * y = (float *) dst->data; + + for (int64_t t = 0; t < nt; t++) { + for (int64_t c = 0; c < nc; c++) { + const float av = a[c]; + const float xv = x[t*nc + c]; + const float ax = xv * av; + const float s = sinf(ax); + y[t*nc + c] = xv + (s*s)/av; + } + } +} + ///////////////////////////////// static void ggml_compute_forward(struct ggml_compute_params * params, struct ggml_tensor * tensor) { @@ -1840,6 +1927,14 @@ static void ggml_compute_forward(struct ggml_compute_params * params, struct ggm { ggml_compute_forward_mul_mat(params, tensor); } break; + case GGML_OP_MUL_MAT_ACC: + { + ggml_compute_forward_mul_mat_acc(params, tensor); + } break; + case GGML_OP_SNAKE_1D: + { + ggml_compute_forward_snake_1d(params, tensor); + } break; case GGML_OP_MUL_MAT_ID: { ggml_compute_forward_mul_mat_id(params, tensor); @@ -2339,6 +2434,16 @@ static int ggml_get_n_tasks(struct ggml_tensor * node, int n_threads) { { n_tasks = n_threads; } break; + case GGML_OP_MUL_MAT_ACC: + { + // reference implementation is single-threaded + n_tasks = 1; + } break; + case GGML_OP_SNAKE_1D: + { + // reference implementation is single-threaded + n_tasks = 1; + } break; case GGML_OP_GET_ROWS: case GGML_OP_SET_ROWS: { diff --git a/external/ggml/src/ggml-metal/ggml-metal-device.cpp b/external/ggml/src/ggml-metal/ggml-metal-device.cpp index 8f11f92a2..b8c544e60 100644 --- a/external/ggml/src/ggml-metal/ggml-metal-device.cpp +++ b/external/ggml/src/ggml-metal/ggml-metal-device.cpp @@ -352,6 +352,26 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_unary(ggml_metal return res; } +ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_snake_1d(ggml_metal_library_t lib, const ggml_tensor * op) { + GGML_ASSERT(op->op == GGML_OP_SNAKE_1D); + GGML_ASSERT(op->src[0]->type == GGML_TYPE_F32); + GGML_ASSERT(op->src[1]->type == GGML_TYPE_F32); + GGML_ASSERT(op->type == GGML_TYPE_F32); + + char base[256]; + char name[256]; + + snprintf(base, 256, "kernel_snake_1d_f32"); + snprintf(name, 256, "%s", base); + + ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name); + if (!res.pipeline) { + res = ggml_metal_library_compile_pipeline(lib, base, name, nullptr); + } + + return res; +} + ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_glu(ggml_metal_library_t lib, const ggml_tensor * op) { GGML_ASSERT(ggml_is_contiguous_1(op->src[0])); @@ -814,6 +834,72 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mm(ggml_meta return res; } +// accumulate-in-place variant of kernel_mul_mm (tensor-core path only): identical tiling and +// threadgroup usage to ggml_metal_library_get_pipeline_mul_mm, just a different kernel that adds +// the result tile into the destination instead of overwriting it. +ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mm_acc(ggml_metal_library_t lib, const ggml_tensor * op) { + char base[256]; + char name[256]; + + const ggml_type tsrc0 = op->src[0]->type; + const ggml_type tsrc1 = op->src[1]->type; + + const bool bc_inp = op->src[0]->ne[0] % 32 != 0; + + constexpr int NRA = SZ_SIMDGROUP * N_MM_BLOCK_Y * N_MM_SIMD_GROUP_Y; + constexpr int NRB = SZ_SIMDGROUP * N_MM_BLOCK_X * N_MM_SIMD_GROUP_X; + + const bool bc_out = (op->ne[0] % NRA != 0 || op->ne[1] % NRB != 0); + + GGML_ASSERT(op->src[1]->ne[2] <= INT16_MAX && op->src[1]->ne[3] <= INT16_MAX); + const int16_t ne12 = (int16_t) op->src[1]->ne[2]; + const int16_t ne13 = (int16_t) op->src[1]->ne[3]; + const int16_t r2 = (int16_t) (ne12 / op->src[0]->ne[2]); + const int16_t r3 = (int16_t) (ne13 / op->src[0]->ne[3]); + + snprintf(base, 256, "kernel_mul_mm_acc_%s_%s", ggml_type_name(tsrc0), ggml_type_name(tsrc1)); + snprintf(name, 256, "%s_bci=%d_bco=%d_ne12=%d_ne13=%d_r2=%d_r3=%d", + base, bc_inp, bc_out, ne12, ne13, r2, r3); + + ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name); + if (!res.pipeline) { + ggml_metal_cv_t cv = ggml_metal_cv_init(); + + ggml_metal_cv_set_bool(cv, bc_inp, FC_MUL_MM + 0); + ggml_metal_cv_set_bool(cv, bc_out, FC_MUL_MM + 1); + ggml_metal_cv_set_int16(cv, ne12, FC_MUL_MM + 2); + ggml_metal_cv_set_int16(cv, ne13, FC_MUL_MM + 3); + ggml_metal_cv_set_int16(cv, r2, FC_MUL_MM + 4); + ggml_metal_cv_set_int16(cv, r3, FC_MUL_MM + 5); + + res = ggml_metal_library_compile_pipeline(lib, base, name, cv); + + ggml_metal_cv_free(cv); + } + + const bool has_tensor = ggml_metal_device_get_props(ggml_metal_library_get_device(lib))->has_tensor; + + if (has_tensor) { + res.nr0 = NRA; + res.nr1 = NRB; + + // threadgroup memory holds the dequantized A tile only (the epilogue accumulates + // through per-thread registers, no extra shared memory) + res.smem = NRA * N_MM_NK_TOTAL * sizeof(ggml_fp16_t); + } else { + res.nr0 = 64; + res.nr1 = 32; + + // the accumulate epilogue always stages the result tile through threadgroup memory + // (NR0 * NR1 floats), which subsumes the sa/sb region + res.smem = 8192; + } + + res.nsg = N_MM_SIMD_GROUP_X * N_MM_SIMD_GROUP_Y; + + return res; +} + ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv(ggml_metal_library_t lib, const ggml_tensor * op) { GGML_TENSOR_LOCALS( int32_t, ne0, op->src[0], ne); GGML_TENSOR_LOCALS( int32_t, ne1, op->src[1], ne); diff --git a/external/ggml/src/ggml-metal/ggml-metal-device.h b/external/ggml/src/ggml-metal/ggml-metal-device.h index 71e7aea5f..16eb89dfc 100644 --- a/external/ggml/src/ggml-metal/ggml-metal-device.h +++ b/external/ggml/src/ggml-metal/ggml-metal-device.h @@ -121,6 +121,7 @@ struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_diag struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_diag_mask_inf (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_repeat (ggml_metal_library_t lib, enum ggml_type tsrc); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_unary (ggml_metal_library_t lib, const struct ggml_tensor * op); +struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_snake_1d (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_glu (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_sum (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_sum_rows (ggml_metal_library_t lib, const struct ggml_tensor * op); @@ -136,6 +137,7 @@ struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_gated_del struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_solve_tri (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv_ext (ggml_metal_library_t lib, const struct ggml_tensor * op, int nsg, int nxpsg, int r1ptg); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mm (ggml_metal_library_t lib, const struct ggml_tensor * op); +struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mm_acc (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mm_id_map0 (ggml_metal_library_t lib, int ne02, int ne20); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mm_id (ggml_metal_library_t lib, const struct ggml_tensor * op); diff --git a/external/ggml/src/ggml-metal/ggml-metal-device.m b/external/ggml/src/ggml-metal/ggml-metal-device.m index dca96bc5c..773a9b33e 100644 --- a/external/ggml/src/ggml-metal/ggml-metal-device.m +++ b/external/ggml/src/ggml-metal/ggml-metal-device.m @@ -1252,6 +1252,21 @@ bool ggml_metal_device_supports_op(ggml_metal_device_t dev, const struct ggml_te case GGML_OP_MUL_MAT: case GGML_OP_MUL_MAT_ID: return has_simdgroup_reduction && op->src[0]->type != GGML_TYPE_NVFP4; + case GGML_OP_SNAKE_1D: + // fused snake activation: elementwise F32, nothing exotic required + return op->src[0]->type == GGML_TYPE_F32 && op->src[1]->type == GGML_TYPE_F32 && + op->type == GGML_TYPE_F32; + case GGML_OP_MUL_MAT_ACC: + // accumulate-in-place matmul: mirrors the has_simdgroup_mm branch of mul_mat + // (the encode always takes the mm kernel), contiguous F32 x F32 -> F32 + return has_simdgroup_mm && + op->src[0]->type == GGML_TYPE_F32 && + op->src[1]->type == GGML_TYPE_F32 && + op->type == GGML_TYPE_F32 && + op->src[0]->ne[0] >= 64 && + op->src[1]->ne[1] > 8 && + !ggml_is_transposed(op->src[0]) && + !ggml_is_transposed(op->src[1]); case GGML_OP_SET: case GGML_OP_CPY: case GGML_OP_DUP: diff --git a/external/ggml/src/ggml-metal/ggml-metal-impl.h b/external/ggml/src/ggml-metal/ggml-metal-impl.h index a6ad1eec5..b567dd215 100644 --- a/external/ggml/src/ggml-metal/ggml-metal-impl.h +++ b/external/ggml/src/ggml-metal/ggml-metal-impl.h @@ -205,6 +205,17 @@ typedef struct { float max; } ggml_metal_kargs_unary; +typedef struct { + int32_t ne00; + int32_t ne01; + uint64_t nb00; + uint64_t nb01; + int32_t ne0; + int32_t ne1; + uint64_t nb0; + uint64_t nb1; +} ggml_metal_kargs_snake_1d; + typedef struct { int32_t ne00; int32_t ne01; diff --git a/external/ggml/src/ggml-metal/ggml-metal-ops.cpp b/external/ggml/src/ggml-metal/ggml-metal-ops.cpp index 40a5dca40..5395e1d58 100644 --- a/external/ggml/src/ggml-metal/ggml-metal-ops.cpp +++ b/external/ggml/src/ggml-metal/ggml-metal-ops.cpp @@ -390,6 +390,10 @@ static int ggml_metal_op_encode_impl(ggml_metal_op_t ctx, int idx) { { n_fuse = ggml_metal_op_unary(ctx, idx); } break; + case GGML_OP_SNAKE_1D: + { + n_fuse = ggml_metal_op_snake_1d(ctx, idx); + } break; case GGML_OP_GLU: { n_fuse = ggml_metal_op_glu(ctx, idx); @@ -436,6 +440,10 @@ static int ggml_metal_op_encode_impl(ggml_metal_op_t ctx, int idx) { { n_fuse = ggml_metal_op_mul_mat(ctx, idx); } break; + case GGML_OP_MUL_MAT_ACC: + { + n_fuse = ggml_metal_op_mul_mat_acc(ctx, idx); + } break; case GGML_OP_MUL_MAT_ID: { n_fuse = ggml_metal_op_mul_mat_id(ctx, idx); @@ -942,6 +950,51 @@ int ggml_metal_op_unary(ggml_metal_op_t ctx, int idx) { return 1; } +int ggml_metal_op_snake_1d(ggml_metal_op_t ctx, int idx) { + ggml_tensor * op = ctx->node(idx); + + ggml_metal_library_t lib = ctx->lib; + ggml_metal_encoder_t enc = ctx->enc; + + GGML_TENSOR_LOCALS( int32_t, ne0, op->src[0], ne); + GGML_TENSOR_LOCALS(uint64_t, nb0, op->src[0], nb); + GGML_TENSOR_LOCALS( int32_t, ne, op, ne); + GGML_TENSOR_LOCALS(uint64_t, nb, op, nb); + + GGML_ASSERT(ggml_is_contiguous(op->src[0])); + GGML_ASSERT(op->src[1]->ne[1] == 1); + + ggml_metal_kargs_snake_1d args = { + /*.ne00 =*/ ne00, + /*.ne01 =*/ ne01, + /*.nb00 =*/ nb00, + /*.nb01 =*/ nb01, + /*.ne0 =*/ ne0, + /*.ne1 =*/ ne1, + /*.nb0 =*/ nb0, + /*.nb1 =*/ nb1, + }; + + auto pipeline = ggml_metal_library_get_pipeline_snake_1d(lib, op); + + ggml_metal_encoder_set_pipeline(enc, pipeline); + ggml_metal_encoder_set_bytes (enc, &args, sizeof(args), 0); + ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[0]), 1); + ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[1]), 2); + ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op), 3); + + const int64_t n = int64_t(ne00)*ne01; + + const int nth = MIN(1024, ggml_metal_pipeline_max_theads_per_threadgroup(pipeline)); + + const int nk0 = int((n + nth - 1)/nth); + + ggml_metal_encoder_dispatch_threadgroups(enc, nk0, 1, 1, nth, 1, 1); + + return 1; +} + + int ggml_metal_op_glu(ggml_metal_op_t ctx, int idx) { ggml_tensor * op = ctx->node(idx); @@ -2469,6 +2522,71 @@ int ggml_metal_op_mul_mat(ggml_metal_op_t ctx, int idx) { return 1; } +// accumulate-in-place matmul (dst aliases src[2]): encodes the tensor-core mm kernel that +// adds the product tile into the destination. Mirrors the has_simdgroup_mm branch of +// ggml_metal_op_mul_mat exactly; only the pipeline (accumulate variant) differs. +int ggml_metal_op_mul_mat_acc(ggml_metal_op_t ctx, int idx) { + ggml_tensor * op = ctx->node(idx); + + ggml_metal_library_t lib = ctx->lib; + ggml_metal_encoder_t enc = ctx->enc; + + GGML_TENSOR_LOCALS( int32_t, ne0, op->src[0], ne); + GGML_TENSOR_LOCALS(uint64_t, nb0, op->src[0], nb); + GGML_TENSOR_LOCALS( int32_t, ne1, op->src[1], ne); + GGML_TENSOR_LOCALS(uint64_t, nb1, op->src[1], nb); + GGML_TENSOR_LOCALS( int32_t, ne, op, ne); + GGML_TENSOR_LOCALS(uint64_t, nb, op, nb); + + GGML_ASSERT(ne00 == ne10); + + GGML_ASSERT(ne12 % ne02 == 0); + GGML_ASSERT(ne13 % ne03 == 0); + + // the kernel assumes a contiguous [ne0, ne1] destination tile (dst stride {1, ne0}) + GGML_ASSERT(ggml_is_contiguous(op)); + + const int16_t r2 = ne12/ne02; + const int16_t r3 = ne13/ne03; + + auto pipeline = ggml_metal_library_get_pipeline_mul_mm_acc(lib, op); + + ggml_metal_kargs_mul_mm args = { + /*.ne00 =*/ ne00, + /*.ne02 =*/ ne02, + /*.nb01 =*/ nb01, + /*.nb02 =*/ nb02, + /*.nb03 =*/ nb03, + /*.ne12 =*/ ne12, + /*.nb10 =*/ nb10, + /*.nb11 =*/ nb11, + /*.nb12 =*/ nb12, + /*.nb13 =*/ nb13, + /*.ne0 =*/ ne0, + /*.ne1 =*/ ne1, + /*.r2 =*/ r2, + /*.r3 =*/ r3, + }; + + ggml_metal_encoder_set_pipeline(enc, pipeline); + ggml_metal_encoder_set_bytes (enc, &args, sizeof(args), 0); + ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[0]), 1); + ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[1]), 2); + ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op), 3); + + const size_t smem = pipeline.smem; + + ggml_metal_encoder_set_threadgroup_memory_size(enc, smem, 0); + + const int nr0 = pipeline.nr0; + const int nr1 = pipeline.nr1; + const int nsg = pipeline.nsg; + + ggml_metal_encoder_dispatch_threadgroups(enc, ((ne11 + nr1 - 1) / nr1), ((ne01 + nr0 - 1) / nr0), ne12 * ne13, 32, nsg, 1); + + return 1; +} + size_t ggml_metal_op_mul_mat_id_extra_tpe(const ggml_tensor * op) { assert(op->op == GGML_OP_MUL_MAT_ID); diff --git a/external/ggml/src/ggml-metal/ggml-metal-ops.h b/external/ggml/src/ggml-metal/ggml-metal-ops.h index 5dc229e26..ef1274349 100644 --- a/external/ggml/src/ggml-metal/ggml-metal-ops.h +++ b/external/ggml/src/ggml-metal/ggml-metal-ops.h @@ -47,6 +47,7 @@ int ggml_metal_op_concat (ggml_metal_op_t ctx, int idx); int ggml_metal_op_repeat (ggml_metal_op_t ctx, int idx); int ggml_metal_op_acc (ggml_metal_op_t ctx, int idx); int ggml_metal_op_unary (ggml_metal_op_t ctx, int idx); +int ggml_metal_op_snake_1d (ggml_metal_op_t ctx, int idx); int ggml_metal_op_glu (ggml_metal_op_t ctx, int idx); int ggml_metal_op_sum (ggml_metal_op_t ctx, int idx); int ggml_metal_op_sum_rows (ggml_metal_op_t ctx, int idx); @@ -66,6 +67,7 @@ int ggml_metal_op_cpy (ggml_metal_op_t ctx, int idx); int ggml_metal_op_pool_1d (ggml_metal_op_t ctx, int idx); int ggml_metal_op_pool_2d (ggml_metal_op_t ctx, int idx); int ggml_metal_op_mul_mat (ggml_metal_op_t ctx, int idx); +int ggml_metal_op_mul_mat_acc (ggml_metal_op_t ctx, int idx); int ggml_metal_op_mul_mat_id (ggml_metal_op_t ctx, int idx); int ggml_metal_op_add_id (ggml_metal_op_t ctx, int idx); int ggml_metal_op_flash_attn_ext (ggml_metal_op_t ctx, int idx); diff --git a/external/ggml/src/ggml-metal/ggml-metal.metal b/external/ggml/src/ggml-metal/ggml-metal.metal index b71b83c68..a090f6440 100644 --- a/external/ggml/src/ggml-metal/ggml-metal.metal +++ b/external/ggml/src/ggml-metal/ggml-metal.metal @@ -1190,6 +1190,42 @@ kernel void kernel_unary_impl( #undef FC_OP #undef FC_CNT +} + +// fused snake activation: y = x + sin(x * alpha[c])^2 / alpha[c], with c = i0 the +// channel (fast axis). One elementwise pass replacing the mul -> sin -> mul -> div +// -> add chain; the per-element op sequence matches the chain bit-for-bit (f32 ops, +// precise sin), so outputs are identical. +kernel void kernel_snake_1d_f32( + constant ggml_metal_kargs_snake_1d & args, + device const char * src0, + device const char * src1, + device char * dst, + uint3 tgpig[[threadgroup_position_in_grid]], + ushort3 tpitg[[thread_position_in_threadgroup]], + ushort3 ntg[[threads_per_threadgroup]]) { + const int n = args.ne00*args.ne01; + + const int ith = tgpig.x*ntg.x + tpitg.x; + + if (ith >= n) { + return; + } + + const int i0 = ith % args.ne00; + const int i1 = ith / args.ne00; + + device const float * x = (device const float *)(src0 + i0*args.nb00 + i1*args.nb01); + device const float * a = (device const float *)(src1 + i0*4); // alpha [C,1] contiguous f32 + device float * y = (device float *)(dst + i0*args.nb0 + i1*args.nb1); + + const float xv = x[0]; + const float av = a[0]; + const float ax = xv * av; + const float s = sin(ax); + const float s2 = s * s; + + y[0] = xv + s2/av; } typedef decltype(kernel_unary_impl) kernel_unary_t; @@ -2386,7 +2422,17 @@ kernel void kernel_ssm_scan_f32( threadgroup_barrier(mem_flags::mem_threadgroup); - const float sumf = simd_sum(shared_sums[sgitg*NW + tiisg]); + // Each token's full output is the sum over all simdgroups of that token's + // partial sums (shared_sums[t*NW + g] for g in 0..sgptg-1). The previous + // simd_sum(shared_sums[sgitg*NW + tiisg]) read garbage columns whenever + // sgptg < NW (e.g. d_state=64 -> sgptg=2) with few tokens, corrupting the + // SSM state. Compute the token sum redundantly on every thread instead. + float sumf = 0.0f; + if (i2 + sgitg < n_t) { + for (int g = 0; g < sgptg; g++) { + sumf += shared_sums[(i2 + sgitg)*NW + g]; + } + } if (tiisg == 0 && i2 + sgitg < n_t) { y[sgitg*nh*nr] = sumf; @@ -10206,6 +10252,146 @@ kernel void kernel_mul_mm( cT.store(tD.slice(ra, rb)); } +// Accumulate-in-place variant of kernel_mul_mm: dst += A * B. +// The MMA body is identical to kernel_mul_mm (same K-reduction order, so the partial +// products are bit-identical); only the epilogue differs, folding the result tile into +// the destination instead of overwriting it. Tensor-core path only (GGML_METAL_HAS_TENSOR). + +template< + typename SA, typename SA_4x4, typename SA_8x8, + typename SB, typename SB_2x4, typename SB_8x8, + typename block_q, short nl, void (*dequantize_func)(device const block_q *, short, thread SA_4x4 &), + typename T0, typename T0_4x4, typename T1, typename T1_2x4> +kernel void kernel_mul_mm_acc( + constant ggml_metal_kargs_mul_mm & args, + device const char * srcA, + device const char * srcB, + device char * dst, + threadgroup char * shmem [[threadgroup(0)]], + uint3 tgpig [[threadgroup_position_in_grid]], + ushort tiitg [[thread_index_in_threadgroup]], + ushort sgitg [[simdgroup_index_in_threadgroup]]) { + (void) sgitg; + + // Matrix dimensions: A(M,K) x B(K,N) -> C(M,N) + const int K = args.ne00; + const int M = args.ne0; + const int N = args.ne1; + + // Batch dimension handling + const int im = tgpig.z; + const int i12 = im % FC_mul_mm_ne12; + const int i13 = im / FC_mul_mm_ne12; + + // Batch offsets for srcA and srcB + const uint64_t offset0 = (i12/FC_mul_mm_r2)*args.nb02 + (i13/FC_mul_mm_r3)*args.nb03; + + // Tile dimensions + constexpr int NRB = SZ_SIMDGROUP * N_MM_BLOCK_X * N_MM_SIMD_GROUP_X; + constexpr int NRA = SZ_SIMDGROUP * N_MM_BLOCK_Y * N_MM_SIMD_GROUP_Y; + + // Tile offsets in output matrix + const int ra = tgpig.y * NRA; + const int rb = tgpig.x * NRB; + + // Threadgroup memory for dequantized A tile only + threadgroup SA * sa = (threadgroup SA *)(shmem); + + // Work-item count for A loading + constexpr int A_WORK_ITEMS = NRA * N_MM_NK; + constexpr int NUM_THREADS = N_SIMDWIDTH * N_MM_SIMD_GROUP_X * N_MM_SIMD_GROUP_Y; + + // tA wraps threadgroup memory + auto tA = tensor(sa, dextents(N_MM_NK_TOTAL, NRA)); + + // tB wraps device memory directly + device T1 * ptrB = (device T1 *)(srcB + args.nb12*i12 + args.nb13*i13); + const int strideB = args.nb11 / sizeof(T1); + auto tB = tensor(ptrB, dextents(K, N), array({1, strideB})); + + // Configure matmul operation + mpp::tensor_ops::matmul2d< + mpp::tensor_ops::matmul2d_descriptor( + NRB, NRA, N_MM_NK_TOTAL, false, true, true, + mpp::tensor_ops::matmul2d_descriptor::mode::multiply_accumulate), + execution_simdgroups> mm; + + auto cT = mm.get_destination_cooperative_tensor(); + + // Accumulate partial results over K dimension + for (int loop_k = 0; loop_k < K; loop_k += N_MM_NK_TOTAL) { + // === PHASE 1: Dequantization of A into threadgroup memory === + for (int work = tiitg; work < A_WORK_ITEMS; work += NUM_THREADS) { + const int row = work / N_MM_NK; + const int k_chunk = work % N_MM_NK; + const int k_pos = loop_k + k_chunk * 16; + const short k_base = k_chunk * 16; + + // Bounds check: skip device read if row is out of matrix bounds + if (ra + row < M) { + if (is_same::value && FC_mul_mm_bc_inp) { + // Element-wise reads when K is not aligned (nb01 not aligned for half4x4/float4x4). + // MSL spec Table 2.5: half4x4 requires 8-byte alignment. When K is odd, + // nb01 = K*2 is not 8-byte aligned, so odd-row pointers are misaligned. + // Mirrors the legacy kernel's existing guard. + device const T0 * row_ptr = (device const T0 *)(srcA + args.nb01 * (ra + row) + offset0); + + FOR_UNROLL (short i = 0; i < 16; i++) { + sa[row * N_MM_NK_TOTAL + (k_base + i)] = (k_pos + i < K) ? (SA) row_ptr[k_pos + i] : (SA)0; + } + } else { + const int block_idx = k_pos / (16 * nl); + const short il = (k_pos / 16) % nl; + + device const block_q * row_ptr = (device const block_q *)(srcA + args.nb01 * (ra + row) + offset0); + + SA_4x4 temp_a; + dequantize_func(row_ptr + block_idx, il, temp_a); + + FOR_UNROLL (short i = 0; i < 16; i++) { + // Zero-pad A for K positions beyond valid range (handles partial K iterations) + sa[row * N_MM_NK_TOTAL + (k_base + i)] = (k_pos + i < K) ? temp_a[i/4][i%4] : (SA)0; + } + } + } else { + // Zero-pad rows beyond matrix bounds + FOR_UNROLL (short i = 0; i < 16; i++) { + sa[row * N_MM_NK_TOTAL + (k_base + i)] = (SA)0; + } + } + } + + threadgroup_barrier(mem_flags::mem_threadgroup); + + // === PHASE 2: Tensor matmul === + auto mA = tA.slice(0, 0); + auto mB = tB.slice(loop_k, rb); + + mm.run(mB, mA, cT); + + threadgroup_barrier(mem_flags::mem_threadgroup); + } + + // Accumulate the result tile onto the output matrix (with batch offset). + // The dst tile already holds the running sum; load it, add the freshly computed + // tile element-wise (one F32 add per element, matching ggml_add), and store back. + // cAcc shares cT's layout, so per-thread element i maps to the same (row, col) and + // each thread reads/adds/stores only its own elements -- no extra synchronization. + device float * dstBatch = (device float *)dst + im * N * M; + + auto tD = tensor(dstBatch, dextents(M, N), array({1, M})); + auto tDst = tD.slice(ra, rb); + + auto cAcc = mm.get_destination_cooperative_tensor(); + cAcc.load(tDst); + + for (auto it = cT.begin(), jt = cAcc.begin(); it != cT.end(); ++it, ++jt) { + *it += *jt; + } + + cT.store(tDst); +} + #else template< @@ -10418,6 +10604,217 @@ kernel void kernel_mul_mm( for (; i < nr0; i++) { *(D + i) = *(C + i); } + } + } + } +} + +// Accumulate-in-place variant of kernel_mul_mm (simdgroup path): dst += A * B. +// The MMA body is identical to kernel_mul_mm (same K-reduction order, so the partial products +// are bit-identical); only the epilogue differs, folding the result tile into the destination +// via one F32 add per element instead of overwriting it. + +template< + typename S0, typename S0_4x4, typename S0_8x8, + typename S1, typename S1_2x4, typename S1_8x8, + typename block_q, short nl, void (*dequantize_func)(device const block_q *, short, thread S0_4x4 &), + typename T0, typename T0_4x4, typename T1, typename T1_2x4> +kernel void kernel_mul_mm_acc( + constant ggml_metal_kargs_mul_mm & args, + device const char * src0, + device const char * src1, + device char * dst, + threadgroup char * shmem [[threadgroup(0)]], + uint3 tgpig[[threadgroup_position_in_grid]], + ushort tiitg[[thread_index_in_threadgroup]], + ushort sgitg[[simdgroup_index_in_threadgroup]]) { + + threadgroup S0 * sa = (threadgroup S0 *)(shmem); + threadgroup S1 * sb = (threadgroup S1 *)(shmem + 4096); + + constexpr int NR0 = 64; + constexpr int NR1 = 32; + + constexpr int NK = 32; + constexpr int NL0 = NK/16; + constexpr int NL1 = NK/8; + + const int im = tgpig.z; + const int r0 = tgpig.y*NR0; + const int r1 = tgpig.x*NR1; + + // if this block is of 64x32 shape or smaller + const short nr0 = (args.ne0 - r0 < NR0) ? (args.ne0 - r0) : NR0; + const short nr1 = (args.ne1 - r1 < NR1) ? (args.ne1 - r1) : NR1; + + // a thread shouldn't load data outside of the matrix + const short lr0 = ((short)tiitg/NL0) < nr0 ? ((short)tiitg/NL0) : nr0 - 1; // 0 .. 63 + const short lr1 = ((short)tiitg/NL1) < nr1 ? ((short)tiitg/NL1) : nr1 - 1; // 0 .. 31 + + const short il0 = (tiitg % NL0); + + short il = il0; + + const int i12 = im % FC_mul_mm_ne12; + const int i13 = im / FC_mul_mm_ne12; + + const uint64_t offset0 = (i12/FC_mul_mm_r2)*args.nb02 + (i13/FC_mul_mm_r3)*args.nb03; + const short offset1 = il0/nl; + + device const block_q * x = (device const block_q *)(src0 + args.nb01*(r0 + lr0) + offset0) + offset1; + + const short iy = 8*(tiitg % NL1); + + device const T1 * y = (device const T1 *)(src1 + + args.nb13*i13 + + args.nb12*i12 + + args.nb11*(r1 + lr1) + + args.nb10*iy); + + S0_8x8 ma[4]; + S1_8x8 mb[2]; + + simdgroup_float8x8 mc[8]; + + for (short i = 0; i < 8; i++){ + mc[i] = make_filled_simdgroup_matrix(0.f); + } + + for (int loop_k = 0; loop_k < args.ne00; loop_k += NK) { + // load data and store to threadgroup memory + if (is_same::value && FC_mul_mm_bc_inp) { + threadgroup_barrier(mem_flags::mem_threadgroup); + + // no need for dequantization + for (short i = 0; i < 16; i++) { + const short sx = 2*il0 + i/8; + const short sy = (tiitg/NL0)/8; + + //const short lx = i%8; + //const short ly = (tiitg/NL0)%8; + const short lx = (tiitg/NL0)%8; + const short ly = i%8; + + const short ib = 8*sx + sy; + + *(sa + 64*ib + 8*ly + lx) = loop_k + 16*il + i < args.ne00 ? *((device T0 *) x + i) : 0; + } + } else { + S0_4x4 temp_a; + dequantize_func(x, il, temp_a); + + threadgroup_barrier(mem_flags::mem_threadgroup); + + FOR_UNROLL (short i = 0; i < 16; i++) { + const short sx = 2*il0 + i/8; + const short sy = (tiitg/NL0)/8; + + //const short lx = i%8; + //const short ly = (tiitg/NL0)%8; + const short lx = (tiitg/NL0)%8; + const short ly = i%8; + + const short ib = 8*sx + sy; + + // NOTE: this is massively slower.. WTF? + //sa[64*ib + 8*ly + lx] = temp_a[i/4][i%4]; + + *(sa + 64*ib + 8*ly + lx) = temp_a[i/4][i%4]; + } + } + + if (FC_mul_mm_bc_inp) { + for (short i = 0; i < 8; ++i) { + const short sx = (tiitg%NL1); + const short sy = (tiitg/NL1)/8; + + const short lx = i; + const short ly = (tiitg/NL1)%8; + //const short lx = (tiitg/NL1)%8; + //const short ly = i; + + const short ib = 4*sx + sy; + + *(sb + 64*ib + 8*ly + lx) = loop_k + iy + i < args.ne00 ? (S1) *((device T1 *) y + i) : 0; + } + } else { + const short sx = (tiitg%NL1); + const short sy = (tiitg/NL1)/8; + + //const short dx = sx; + //const short dy = sy; + + const short ly = (tiitg/NL1)%8; + + const short ib = 4*sx + sy; + + *(threadgroup S1_2x4 *)(sb + 64*ib + 8*ly) = (S1_2x4)(*((device T1_2x4 *) y)); + } + + il = (il + 2 < nl) ? il + 2 : il % 2; + x = (il < 2) ? x + (2 + nl - 1)/nl : x; + + y += NK; + + threadgroup_barrier(mem_flags::mem_threadgroup); + + // load matrices from threadgroup memory and conduct outer products + threadgroup const S0 * lsma = (sa + 4*64*(sgitg%2)); + threadgroup const S1 * lsmb = (sb + 2*64*(sgitg/2)); + + FOR_UNROLL (short ik = 0; ik < NK/8; ik++) { + simdgroup_barrier(mem_flags::mem_none); + + FOR_UNROLL (short i = 0; i < 4; i++) { + simdgroup_load(ma[i], lsma + 64*i, 8, 0, false); + } + + simdgroup_barrier(mem_flags::mem_none); + + FOR_UNROLL (short i = 0; i < 2; i++) { + simdgroup_load(mb[i], lsmb + 64*i, 8, 0, false); + } + + simdgroup_barrier(mem_flags::mem_none); + + FOR_UNROLL (short i = 0; i < 8; i++){ + simdgroup_multiply_accumulate(mc[i], mb[i/4], ma[i%4], mc[i]); + } + + lsma += 8*64; + lsmb += 4*64; + } + } + + // accumulate: stage the result tiles to threadgroup memory, then add onto dst in-place. + // one F32 add per element, matching ggml_add -> bit-identical to mul_mm + add. Always + // staged (no direct-store fast path) so partial tiles are clipped the same way. + threadgroup_barrier(mem_flags::mem_threadgroup); + + threadgroup float * temp_str = ((threadgroup float *) shmem) + 32*(sgitg&1) + (16*(sgitg >> 1))*NR0; + + for (short i = 0; i < 8; i++) { + simdgroup_store(mc[i], temp_str + 8*(i%4) + 8*NR0*(i/4), NR0, 0, false); + } + + threadgroup_barrier(mem_flags::mem_threadgroup); + + if (sgitg == 0) { + for (int j = tiitg; j < nr1; j += NR1) { + device float * D = (device float *) dst + r0 + (r1 + j)*args.ne0 + im*args.ne1*args.ne0; + device float4 * D4 = (device float4 *) D; + + threadgroup float * C = temp_str + (j*NR0); + threadgroup float4 * C4 = (threadgroup float4 *) C; + + int i = 0; + for (; i < nr0/4; i++) { + *(D4 + i) += *(C4 + i); + } + + i *= 4; + for (; i < nr0; i++) { + *(D + i) += *(C + i); } } } @@ -10872,6 +11269,10 @@ template [[host_name("kernel_set_rows_iq4_nl_i32")]] kernel set_rows_q32_t kerne typedef decltype(kernel_mul_mm) mul_mm_t; template [[host_name("kernel_mul_mm_f32_f32")]] kernel mul_mm_t kernel_mul_mm; + +typedef decltype(kernel_mul_mm_acc) mul_mm_acc_t; + +template [[host_name("kernel_mul_mm_acc_f32_f32")]] kernel mul_mm_acc_t kernel_mul_mm_acc; template [[host_name("kernel_mul_mm_f16_f32")]] kernel mul_mm_t kernel_mul_mm; #if defined(GGML_METAL_HAS_BF16) template [[host_name("kernel_mul_mm_bf16_f32")]] kernel mul_mm_t kernel_mul_mm; diff --git a/external/ggml/src/ggml.c b/external/ggml/src/ggml.c index 39319bfc7..233edec8a 100644 --- a/external/ggml/src/ggml.c +++ b/external/ggml/src/ggml.c @@ -1131,9 +1131,11 @@ static const char * GGML_OP_NAME[GGML_OP_COUNT] = { "MUL_MAT_ADD", "MUL_MAT_ADD_RELU", "IM2COL_ASYM", + "MUL_MAT_ACC", + "SNAKE_1D", }; -static_assert(GGML_OP_COUNT == 107, "GGML_OP_COUNT != 107"); +static_assert(GGML_OP_COUNT == 109, "GGML_OP_COUNT != 109"); static const char * GGML_OP_SYMBOL[GGML_OP_COUNT] = { "none", @@ -1253,9 +1255,11 @@ static const char * GGML_OP_SYMBOL[GGML_OP_COUNT] = { "a*b+bias", "relu(a*b+bias)", "im2col_asym(x)", + "mul_mat_acc(a, b, acc)", + "snake_1d(a, alpha)", }; -static_assert(GGML_OP_COUNT == 107, "GGML_OP_COUNT != 107"); +static_assert(GGML_OP_COUNT == 109, "GGML_OP_COUNT != 109"); static_assert(GGML_OP_POOL_COUNT == 2, "GGML_OP_POOL_COUNT != 2"); @@ -3337,6 +3341,51 @@ struct ggml_tensor * ggml_mul_mat( result->op = GGML_OP_MUL_MAT; result->src[0] = a; result->src[1] = b; + + return result; +} + +struct ggml_tensor * ggml_mul_mat_acc( + struct ggml_context * ctx, + struct ggml_tensor * a, + struct ggml_tensor * b, + struct ggml_tensor * acc) { + GGML_ASSERT(ggml_can_mul_mat(a, b)); + GGML_ASSERT(!ggml_is_transposed(a)); + GGML_ASSERT(acc->type == GGML_TYPE_F32); + // acc must have the shape of the a * b product + GGML_ASSERT(acc->ne[0] == a->ne[1] && acc->ne[1] == b->ne[1] && + acc->ne[2] == b->ne[2] && acc->ne[3] == b->ne[3]); + + // result is a view of acc: the accumulation is written in-place into acc's memory + struct ggml_tensor * result = ggml_view_tensor(ctx, acc); + + result->op = GGML_OP_MUL_MAT_ACC; + result->src[0] = a; + result->src[1] = b; + result->src[2] = acc; + + return result; +} + +// ggml_snake_1d + +struct ggml_tensor * ggml_snake_1d( + struct ggml_context * ctx, + struct ggml_tensor * a, + struct ggml_tensor * alpha) { + GGML_ASSERT(a->type == GGML_TYPE_F32); + GGML_ASSERT(alpha->type == GGML_TYPE_F32); + GGML_ASSERT(ggml_is_contiguous(a)); + GGML_ASSERT(ggml_is_contiguous(alpha)); + GGML_ASSERT(alpha->ne[0] == a->ne[0]); + GGML_ASSERT(alpha->ne[1] == 1 && alpha->ne[2] == 1 && alpha->ne[3] == 1); + + struct ggml_tensor * result = ggml_dup_tensor(ctx, a); + + result->op = GGML_OP_SNAKE_1D; + result->src[0] = a; + result->src[1] = alpha; return result; } diff --git a/include/engine/framework/modules/conv_modules.h b/include/engine/framework/modules/conv_modules.h index 0a04766ae..df19a03f1 100644 --- a/include/engine/framework/modules/conv_modules.h +++ b/include/engine/framework/modules/conv_modules.h @@ -40,6 +40,20 @@ class Conv1dModule { Conv1dConfig config_; }; +// Raw Metal fast paths on channel-fast activations (ne = [channels, frames], contiguous +// F32) for the audio codec decoder's chained regions. Unlike the module build() methods +// these take and return raw ggml tensors without the canonical [frames, channels] +// orientation, so consecutive convolutions chain without paying two transposes per conv. +// The caller owns layout conversion at region edges and causal padding. +// +// conv1d_pertap_channel_fast: requires padding=0, stride=1; per-tap GEMM decomposition; +// returns [out_channels, output_frames] with bias broadcast-added when use_bias. +ggml_tensor * conv1d_pertap_channel_fast( + core::ModuleBuildContext & ctx, + const Conv1dWeights & weights, + ggml_tensor * input_cf, + const Conv1dConfig & config); + struct Conv2dConfig { int64_t in_channels = 0; int64_t out_channels = 0; @@ -278,6 +292,16 @@ bool is_conv_transpose1d_col2im_fast_path_eligible( const core::ModuleBuildContext & ctx, const ConvTranspose1dConfig & config) noexcept; +// conv_transpose1d_col2im_channel_fast: same col2im math as ConvTranspose1dModule's fast +// path but consumes channel-fast input directly (skipping its internal transpose) and +// returns the raw time-fast [frames_out, out_channels] tensor with bias included; the +// caller owns causal trimming and layout conversion. +ggml_tensor * conv_transpose1d_col2im_channel_fast( + core::ModuleBuildContext & ctx, + const ConvTranspose1dWeights & weights, + ggml_tensor * input_cf, + const ConvTranspose1dConfig & config); + class ConvTranspose1dModule { public: explicit ConvTranspose1dModule(ConvTranspose1dConfig config); diff --git a/src/community_models/audio8_tts/ar.cpp b/src/community_models/audio8_tts/ar.cpp index bdd03d2b2..ba11865ca 100644 --- a/src/community_models/audio8_tts/ar.cpp +++ b/src/community_models/audio8_tts/ar.cpp @@ -1,5 +1,7 @@ #include "engine/community_models/audio8_tts/ar.h" +#include "falcon_kv_cache.h" + #include "engine/framework/core/backend_weight_store.h" #include "engine/framework/core/backend.h" #include "engine/framework/debug/profiler.h" @@ -61,6 +63,13 @@ struct ArkttsARProfile { double sample_main_ms = 0.0; double sample_high_ms = 0.0; double sample_fast_ms = 0.0; + double falcon_step_init_ms = 0.0; + double falcon_step_build_ms = 0.0; + double falcon_step_gallocr_ms = 0.0; + double falcon_step_upload_ms = 0.0; + double falcon_step_compute_ms = 0.0; + double falcon_step_download_ms = 0.0; + int64_t falcon_step_runs = 0; int64_t prefill_runs = 0; int64_t step_runs = 0; int64_t fast_runs = 0; @@ -103,6 +112,7 @@ struct FalconH1LayerWeights { assets::TensorDataF32 input_layernorm; // slow.layers.*.input_layernorm.weight [512] core::TensorValue ssm_in; // slow.layers.*.mamba.in_proj.weight [1688,512] core::TensorValue ssm_conv1d; // slow.layers.*.mamba.conv1d.weight [896,1,4] -> [4,896] after convert + assets::TensorDataF32 conv1d_kernel; // host [d_conv, conv_dim] GGUF layout for ggml ssm_conv assets::TensorDataF32 ssm_conv1d_b; // slow.layers.*.mamba.conv1d.bias [896] core::TensorValue ssm_dt_b; // slow.layers.*.mamba.dt_bias [24] core::TensorValue ssm_A; // slow.layers.*.mamba.A_log [24] -> [1,24] @@ -431,21 +441,40 @@ FalconH1LayerWeights load_falcon_layer( w.ssm_in = store.load_tensor(source, prefix + ".mamba.in_proj.weight", storage_type, meta.shape); } { + // ssm_conv Metal pipeline requires contiguous F32 conv weights. auto meta = source.require_metadata(prefix + ".mamba.conv1d.weight"); - w.ssm_conv1d = store.load_tensor(source, prefix + ".mamba.conv1d.weight", storage_type, meta.shape); + w.ssm_conv1d = store.load_tensor(source, prefix + ".mamba.conv1d.weight", assets::TensorStorageType::F32, meta.shape); + { + // ggml ssm_conv computes y[c] = sum_k w[k,c]*window[k,c] with window[0] + // the OLDEST frame. HF (both nn.Conv1d prefill and the cached + // torch.sum(conv_states * w, dim=-1) decode) uses the identical + // orientation: w[...,0] multiplies the oldest frame. The GGUF tensor + // is the HF [conv_dim,1,d_conv] weight with reversed dims + // [d_conv,1,conv_dim] and unchanged flat bytes, i.e. + // raw[k + d_conv*c] == hf_w[c,k] — exactly what ssm_conv wants. + // Feed it through UNFLIPPED (an earlier kernel flip here reversed the + // tap order and corrupted the x/B/C split on every step). + auto raw = source.require_f32_tensor(prefix + ".mamba.conv1d.weight"); + const int64_t d_conv = raw.shape.dims[0]; + const int64_t conv_dim = raw.shape.dims[2]; + w.conv1d_kernel.shape = core::TensorShape::from_dims({d_conv, conv_dim}); + w.conv1d_kernel.values = raw.values; + } } w.ssm_conv1d_b = source.require_f32_tensor(prefix + ".mamba.conv1d.bias"); + // Per-head Mamba params must stay unquantized (Native): they are consumed as + // raw F32 scalars by the SSM path (A = -exp(A_log), D, dt bias). { auto meta = source.require_metadata(prefix + ".mamba.dt_bias"); - w.ssm_dt_b = store.load_tensor(source, prefix + ".mamba.dt_bias", storage_type, meta.shape); + w.ssm_dt_b = store.load_tensor(source, prefix + ".mamba.dt_bias", assets::TensorStorageType::F32, meta.shape); } { auto meta = source.require_metadata(prefix + ".mamba.A_log"); - w.ssm_A = store.load_tensor(source, prefix + ".mamba.A_log", storage_type, meta.shape); + w.ssm_A = store.load_tensor(source, prefix + ".mamba.A_log", assets::TensorStorageType::F32, meta.shape); } { auto meta = source.require_metadata(prefix + ".mamba.D"); - w.ssm_D = store.load_tensor(source, prefix + ".mamba.D", storage_type, meta.shape); + w.ssm_D = store.load_tensor(source, prefix + ".mamba.D", assets::TensorStorageType::F32, meta.shape); } { auto meta = source.require_metadata(prefix + ".mamba.out_proj.weight"); @@ -846,135 +875,1137 @@ std::vector build_falcon_embeddings( for (int64_t step = 0; step < steps; ++step) { const int32_t token = matrix[step]; auto row = lookup_row(weights.text_embedding_host, token, hidden); - for (auto & v : row) v *= config.text.embedding_multiplier; if (is_semantic_token(config, token)) { for (int64_t codebook = 0; codebook < config.fast.num_codebooks; ++codebook) { const int32_t code = matrix[(codebook + 1) * steps + step]; add_row(weights.codebook_embedding_host, codebook * config.fast.vocab_size + code, hidden, row); } } + for (auto & v : row) v *= config.text.embedding_multiplier; std::copy(row.begin(), row.end(), out.begin() + static_cast(step * hidden)); } return out; } -// TODO(Falcon-H1): Replace with full Mamba2 port (ggml_ssm_conv + B/C/dt/A/D -// + ggml_ssm_scan + recurrent conv/ssm state + hybrid attention). -// See docs/FALCON_H1_0.1B_PORT_PLAN.md M2/M3 and -// ../llama.cpp/src/models/mamba-base.cpp:151 / falcon-h1.cpp:132. -// Current stub keeps weight loading native but omits the SSM core and -// hybrid attention (attn_out = 0), recomputes full sequence each step, -// and only applies ssm_out/lm_head multipliers — tracked for follow-up. -SlowForwardOutput falcon_forward_stateless( +// ============================================================================ +// Falcon-H1 (Mamba2 + hybrid GQA attention) stateful single-token forward. +// Replaces the documented stub (TODO(Falcon-H1)) with a full Mamba2 port +// mirroring transformers.models.falcon_h1 FalconH1DecoderLayer and +// llama.cpp mamba-base.cpp build_mamba2_layer. Verified shapes against +// Audio8-TTS-Preview-0.1b (dim 512, d_ssm 768, d_state 64, d_conv 4, +// mamba heads 24 x head 32, GQA 8/2 x 64, RoPE NEOX base 1e11). +// ============================================================================ + +// One baked Falcon-H1 step graph (zero-copy path), valid for a fixed KV +// capacity bucket. All weights and recurrent state are bound as external views +// of host memory owned by the weights/state structs, so a plan can be reused +// for every step whose sequence fits the bucket. +struct FalconStepPlan { + int64_t cap = 0; + std::unique_ptr ctx; + ggml_cgraph * gf = nullptr; // owned by ctx + ggml_gallocr_t gallocr = nullptr; + ggml_tensor * logits_out = nullptr; + ggml_tensor * hidden_out = nullptr; + ggml_backend_t backend = nullptr; // borrowed; the runtime outlives any generation + ~FalconStepPlan() { + if (backend != nullptr && gf != nullptr) { + core::release_backend_graph_resources(backend, gf); + } + if (gallocr != nullptr) { + ggml_gallocr_free(gallocr); + } + } + FalconStepPlan() = default; + FalconStepPlan(const FalconStepPlan &) = delete; + FalconStepPlan & operator=(const FalconStepPlan &) = delete; +}; + +struct FalconH1StepState { + int64_t n_layer = 0; + int64_t d_inner = 0; // mamba_d_ssm (768) + int64_t d_state = 0; // mamba_d_state (64) + int64_t d_conv = 0; // mamba_d_conv (4) + int64_t n_groups = 0; // mamba_n_groups (1) + int64_t n_mamba_heads = 0; // mamba_n_heads (24) + int64_t conv_dim = 0; // d_inner + 2*ng*d_state (896) + int64_t kv_dim = 0; // n_local_heads * head_dim (128) + std::vector> conv_states; // [layer][(d_conv-1)*conv_dim] + std::vector> ssm_states; // [layer][d_state*d_inner] + std::vector> k_cache; // [layer][seq*kv_dim] + std::vector> v_cache; // [layer][seq*kv_dim] + int64_t seq_len = 0; + // Padded KV caches for the zero-copy path: flat [head_dim, kv_cap, n_kv] + // (head stride = head_dim*kv_cap), grown geometrically so each step's graph + // can write the new token's k/v in-place and flash-attend over a strided + // prefix view — no per-step host upload, read-back, or concat copy. The + // exact-size k_cache/v_cache vectors above are only used on the fallback + // (non-host backend) path. + std::vector> k_pad; + std::vector> v_pad; + int64_t kv_cap = 0; + // Zero-copy constant cache (filled once per generation on host backends): + // A = -exp(A_log) per mamba head, D expanded per channel, plus fallback + // buffers for absent norm/bias weights and the scalar graph inputs. + std::vector> pre_A; // [layer][n_mamba_heads] + std::vector> pre_D; // [layer][d_inner] + std::vector ones_dim; // [dim] + std::vector zeros_conv; // [conv_dim] + std::vector zeros_conv_kernel; // [d_conv*conv_dim] + int32_t ids_value = 0; + int32_t pos_value = 0; + bool constants_ready = false; + // Bucketed plan reuse (zero-copy path): staging buffer for the per-step + // embedding (the graph input tensor is baked to its address), the flash + // mask scratch (slot seq is the last visible position), the dynamic + // set_rows slot index, and the baked graph itself. + std::vector emb_stage; // [dim] + std::vector kv_mask; // [kv_cap] + int32_t kv_slot = 0; + std::unique_ptr plan; +}; + +FalconH1StepState init_falcon_step_state(const Audio8TtsConfig & config) { + FalconH1StepState st; + st.n_layer = config.text.n_layer; + st.d_inner = config.text.mamba_d_ssm; + st.d_state = config.text.mamba_d_state; + st.d_conv = config.text.mamba_d_conv; + st.n_groups = config.text.mamba_n_groups; + st.n_mamba_heads = config.text.mamba_n_heads; + st.conv_dim = st.d_inner + 2 * st.n_groups * st.d_state; + st.kv_dim = config.text.n_local_heads * config.text.head_dim; + st.conv_states.resize(static_cast(st.n_layer)); + st.ssm_states.resize(static_cast(st.n_layer)); + st.k_cache.resize(static_cast(st.n_layer)); + st.v_cache.resize(static_cast(st.n_layer)); + st.k_pad.resize(static_cast(st.n_layer)); + st.v_pad.resize(static_cast(st.n_layer)); + const size_t conv_sz = static_cast((st.d_conv - 1) * st.conv_dim); + const size_t ssm_sz = static_cast(st.d_state * st.d_inner); + for (int64_t i = 0; i < st.n_layer; ++i) { + st.conv_states[static_cast(i)].assign(conv_sz, 0.0F); + st.ssm_states[static_cast(i)].assign(ssm_sz, 0.0F); + } + st.seq_len = 0; + return st; +} + +// Grow the zero-copy padded KV caches. The head stride changes with capacity, +// so live slots are re-laid-out per head; called between steps, and per-step +// graphs re-bind their views from scratch afterwards. +void grow_falcon_kv_pad(FalconH1StepState & state, int64_t head_dim, int64_t n_kv, int64_t need) { + if (need <= state.kv_cap) return; + // 128-slot buckets: fine-grained enough to keep padded flash cheap, and + // they keep nek1 below the CPU flash kernel's split-KV threshold (512) for + // as long as possible so the masked padded reduction stays bitwise equal + // to the exact-prefix one. + int64_t new_cap = std::max(64, ((need + 127) / 128) * 128); + const size_t live = static_cast(head_dim * state.seq_len); // valid floats per head + for (auto * caches : {&state.k_pad, &state.v_pad}) { + for (auto & cache : *caches) { + std::vector next(static_cast(head_dim * new_cap * n_kv), 0.0F); + if (!cache.empty() && live > 0) { + const size_t old_stride = static_cast(head_dim * state.kv_cap); + const size_t new_stride = static_cast(head_dim * new_cap); + for (int64_t h = 0; h < n_kv; ++h) { + std::memcpy(next.data() + h * new_stride, + cache.data() + h * old_stride, + live * sizeof(float)); + } + } + cache.swap(next); + } + } + state.kv_cap = new_cap; +} + +// Loop-invariant SSM constants, resolved once per generation: A = -exp(A_log) +// per mamba head, D expanded per channel. +void precompute_falcon_constants( + const Audio8TtsConfig & config, + const ArkttsARWeights & weights, + FalconH1StepState & state) { + if (state.constants_ready) return; + const int64_t n_layer = config.text.n_layer; + const int64_t d_inner = config.text.mamba_d_ssm; + const int64_t n_mamba_heads = config.text.mamba_n_heads; + const int64_t mamba_head_dim = config.text.mamba_d_head; + state.pre_A.resize(static_cast(n_layer)); + state.pre_D.resize(static_cast(n_layer)); + for (int64_t li = 0; li < n_layer; ++li) { + const auto & layer = weights.falcon_layers[static_cast(li)]; + auto & a_vals = state.pre_A[static_cast(li)]; + a_vals.resize(static_cast(n_mamba_heads)); + std::vector a_log(static_cast(n_mamba_heads)); + ggml_backend_tensor_get(layer.ssm_A.tensor, a_log.data(), 0, a_log.size() * sizeof(float)); + for (int64_t h = 0; h < n_mamba_heads; ++h) { + a_vals[static_cast(h)] = -std::exp(a_log[static_cast(h)]); + } + auto & d_vals = state.pre_D[static_cast(li)]; + d_vals.resize(static_cast(d_inner)); + std::vector d_raw(static_cast(n_mamba_heads)); + ggml_backend_tensor_get(layer.ssm_D.tensor, d_raw.data(), 0, d_raw.size() * sizeof(float)); + for (int64_t h = 0; h < n_mamba_heads; ++h) { + for (int64_t d = 0; d < mamba_head_dim; ++d) { + d_vals[static_cast(d + h * mamba_head_dim)] = d_raw[static_cast(h)]; + } + } + } + state.constants_ready = true; +} + +// Builds the reusable per-bucket Falcon step graph for the zero-copy path. +// Every weight and recurrent-state tensor is bound as an external view of host +// memory owned by `weights`/`state` (gallocr skips tensors whose data is set +// externally), conv/ssm state write-backs and the KV slot writes run in-graph +// (ggml_cpy / ggml_set_rows), and flash attention reads the full padded cache +// under a mask so the graph topology does not depend on the sequence length. +std::unique_ptr build_falcon_step_plan( + ggml_backend_t backend, + size_t arena_bytes, + const Audio8TtsConfig & config, + const ArkttsARWeights & weights, + FalconH1StepState & state, + ArkttsARProfile * profile) { + const int64_t dim = config.text.dim; + const int64_t n_layer = config.text.n_layer; + const int64_t d_inner = config.text.mamba_d_ssm; + const int64_t d_state = config.text.mamba_d_state; + const int64_t d_conv = config.text.mamba_d_conv; + const int64_t n_groups = config.text.mamba_n_groups; + const int64_t n_mamba_heads = config.text.mamba_n_heads; + const int64_t mamba_head_dim = config.text.mamba_d_head; + const int64_t conv_dim = d_inner + 2 * n_groups * d_state; + const int64_t n_head = config.text.n_head; + const int64_t n_kv = config.text.n_local_heads; + const int64_t head_dim = config.text.head_dim; + const float norm_eps = config.text.norm_eps; + const float rope_base = config.text.rope_base; + const int64_t cap = state.kv_cap; + + auto plan = std::make_unique(); + plan->cap = cap; + plan->backend = backend; + + auto t_init = Clock::now(); + ggml_init_params params{arena_bytes, nullptr, true}; + plan->ctx.reset(ggml_init(params)); + if (!plan->ctx) throw std::runtime_error("build_falcon_step_plan: ggml_init failed"); + ggml_context * ctx = plan->ctx.get(); + auto t_build = Clock::now(); + + // Mask scratch: slots [0, seq_len] visible, the rest -inf; the runner + // reveals one more slot per step. + state.kv_mask.assign(static_cast(cap), ggml_fp32_to_fp16(-INFINITY)); + for (int64_t i = 0; i <= state.seq_len && i < cap; ++i) { + state.kv_mask[static_cast(i)] = ggml_fp32_to_fp16(0.0F); + } + ggml_tensor * mask_t = ggml_new_tensor_4d(ctx, GGML_TYPE_F16, cap, 1, 1, 1); + mask_t->data = state.kv_mask.data(); + + auto bind_const = [&](ggml_tensor * t, const std::vector & values, + std::vector & fallback, float fallback_fill) { + if (!values.empty()) { + t->data = const_cast(values.data()); + return; + } + if (fallback.empty()) { + fallback.assign(static_cast(ggml_nelements(t)), fallback_fill); + } + t->data = fallback.data(); + }; + + ggml_tensor * cur = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, dim); + ggml_set_name(cur, "falcon_input"); + cur->data = state.emb_stage.data(); + + std::vector writebacks; // conv/ssm state tails (dependency-ordered) + std::vector kv_writebacks; // KV slot writes alias the flash reads: expanded first + writebacks.reserve(static_cast(n_layer) * 2); + kv_writebacks.reserve(static_cast(n_layer) * 2); + + for (int64_t li = 0; li < n_layer; ++li) { + const auto & layer = weights.falcon_layers[static_cast(li)]; + + // input_layernorm (RMS) + ggml_tensor * ln_w = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, dim); + bind_const(ln_w, layer.input_layernorm.values, state.ones_dim, 1.0F); + ggml_tensor * normed = ggml_rms_norm(ctx, cur, norm_eps); + normed = ggml_mul(ctx, normed, ln_w); + + // ---- Mamba2 branch ---- + // zxBCdt = in_proj(normed) -> [d_inner + conv_dim + n_mamba_heads] + ggml_tensor * zxBCdt = ggml_mul_mat(ctx, layer.ssm_in.tensor, normed); + ggml_tensor * z = ggml_view_1d(ctx, zxBCdt, d_inner, 0); + ggml_tensor * xBC = ggml_view_1d(ctx, zxBCdt, conv_dim, d_inner * ggml_element_size(zxBCdt)); + ggml_tensor * dt = ggml_view_1d(ctx, zxBCdt, n_mamba_heads, (d_inner + conv_dim) * ggml_element_size(zxBCdt)); + + // conv: state (d_conv-1 rows) + current xBC -> [d_conv, conv_dim, 1] + ggml_tensor * st_t = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, conv_dim, d_conv - 1); + st_t->data = state.conv_states[static_cast(li)].data(); + ggml_tensor * stT = ggml_cont(ctx, ggml_transpose(ctx, st_t)); // [d_conv-1, conv_dim] + ggml_tensor * xBC_r = ggml_cont(ctx, ggml_transpose(ctx, ggml_reshape_2d(ctx, xBC, conv_dim, 1))); // [1, conv_dim] + ggml_tensor * sx = ggml_concat(ctx, stT, xBC_r, 0); // [d_conv, conv_dim] + // Next conv state = sx rows 1..d_conv-1, written back into the host + // state vector in-graph (the cpy depends on sx, hence on the cont() + // that consumed st_t, so the old state is read before it is + // overwritten). + ggml_tensor * tail = ggml_view_2d(ctx, sx, d_conv - 1, conv_dim, + sx->nb[1], ggml_element_size(sx)); + writebacks.push_back(ggml_cpy(ctx, ggml_transpose(ctx, tail), st_t)); + ggml_tensor * sx3 = ggml_reshape_3d(ctx, sx, d_conv, conv_dim, 1); + // ggml ssm_conv computes y[c] = sum_k w[k,c]*sx[k,c] with sx row 0 the + // oldest frame — the same orientation as the HF conv1d/cached decode. + ggml_tensor * conv_w2 = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, d_conv, conv_dim); + bind_const(conv_w2, layer.conv1d_kernel.values, state.zeros_conv_kernel, 0.0F); + ggml_tensor * xBC_conv = ggml_ssm_conv(ctx, sx3, conv_w2); // [conv_dim, 1, 1] + ggml_tensor * conv_b = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, conv_dim); + bind_const(conv_b, layer.ssm_conv1d_b.values, state.zeros_conv, 0.0F); + xBC_conv = ggml_add(ctx, xBC_conv, ggml_reshape_3d(ctx, conv_b, conv_dim, 1, 1)); + xBC_conv = ggml_silu(ctx, xBC_conv); + + // split x / B / C (conv output is contiguous [conv_dim,1,1]) + ggml_tensor * x = ggml_view_1d(ctx, xBC_conv, d_inner, 0); + ggml_tensor * B = ggml_view_1d(ctx, xBC_conv, d_state * n_groups, d_inner * ggml_element_size(xBC_conv)); + ggml_tensor * C = ggml_view_1d(ctx, xBC_conv, d_state * n_groups, (d_inner + d_state * n_groups) * ggml_element_size(xBC_conv)); + + // x -> [head_dim, n_mamba_heads, 1, 1] + ggml_tensor * x4 = ggml_view_4d(ctx, x, mamba_head_dim, n_mamba_heads, 1, 1, + mamba_head_dim * ggml_element_size(x), + mamba_head_dim * n_mamba_heads * ggml_element_size(x), + mamba_head_dim * n_mamba_heads * ggml_element_size(x), 0); + ggml_tensor * B4 = ggml_view_4d(ctx, B, d_state, n_groups, 1, 1, + d_state * ggml_element_size(B), + d_state * n_groups * ggml_element_size(B), + d_state * n_groups * ggml_element_size(B), 0); + ggml_tensor * C4 = ggml_view_4d(ctx, C, d_state, n_groups, 1, 1, + d_state * ggml_element_size(C), + d_state * n_groups * ggml_element_size(C), + d_state * n_groups * ggml_element_size(C), 0); + + // dt = dt + dt_bias -> [n_mamba_heads, 1, 1] + ggml_tensor * dt_eff = ggml_add(ctx, dt, layer.ssm_dt_b.tensor); + ggml_tensor * dt3 = ggml_view_3d(ctx, dt_eff, n_mamba_heads, 1, 1, + n_mamba_heads * ggml_element_size(dt_eff), + n_mamba_heads * ggml_element_size(dt_eff), 0); + + // A = -exp(A_log), precomputed: [1, n_mamba_heads] + ggml_tensor * A_t = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 1, n_mamba_heads); + A_t->data = state.pre_A[static_cast(li)].data(); + + // ssm state: [d_state, mamba_head_dim, n_mamba_heads] + ggml_tensor * ssm_t = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, d_state, mamba_head_dim, n_mamba_heads); + ssm_t->data = state.ssm_states[static_cast(li)].data(); + + // ids for scan (1 sequence) + ggml_tensor * ids = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, 1); + ids->data = &state.ids_value; + + ggml_tensor * scan = ggml_ssm_scan(ctx, ssm_t, x4, dt3, A_t, B4, C4, ids); + // New ssm state = scan tail (after the d_inner y values), written back + // into the host state vector in-graph (ordered after the scan read). + ggml_tensor * next_state = ggml_view_3d(ctx, scan, + d_state, mamba_head_dim, n_mamba_heads, + d_state * ggml_element_size(scan), + d_state * mamba_head_dim * ggml_element_size(scan), + d_inner * ggml_element_size(scan)); + writebacks.push_back(ggml_cpy(ctx, next_state, ssm_t)); + ggml_tensor * y = ggml_view_1d(ctx, scan, d_inner, 0); + + // y += x * D (D precomputed) + ggml_tensor * D_t = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, d_inner); + D_t->data = state.pre_D[static_cast(li)].data(); + y = ggml_add(ctx, y, ggml_mul(ctx, ggml_view_1d(ctx, x4, d_inner, 0), D_t)); + + // z gate: y *= silu(z) + y = ggml_mul(ctx, y, ggml_silu(ctx, z)); + + ggml_tensor * out_mamba = ggml_mul_mat(ctx, layer.ssm_out.tensor, y); // [dim] + + // ---- GQA attention branch ---- + // ggml_flash_attn_ext layout: [head_dim, n_tokens, n_head, batch]. + ggml_tensor * q = ggml_mul_mat(ctx, layer.attn_q_proj.tensor, normed); // [n_head*head_dim] + ggml_tensor * k = ggml_mul_mat(ctx, layer.attn_k_proj.tensor, normed); // [n_kv*head_dim] + ggml_tensor * v = ggml_mul_mat(ctx, layer.attn_v_proj.tensor, normed); // [n_kv*head_dim] + + ggml_tensor * q4 = ggml_view_4d(ctx, q, head_dim, n_head, 1, 1, + head_dim * ggml_element_size(q), + head_dim * n_head * ggml_element_size(q), + head_dim * n_head * ggml_element_size(q), 0); + ggml_tensor * k4 = ggml_view_4d(ctx, k, head_dim, n_kv, 1, 1, + head_dim * ggml_element_size(k), + head_dim * n_kv * ggml_element_size(k), + head_dim * n_kv * ggml_element_size(k), 0); + ggml_tensor * v4 = ggml_view_4d(ctx, v, head_dim, n_kv, 1, 1, + head_dim * ggml_element_size(v), + head_dim * n_kv * ggml_element_size(v), + head_dim * n_kv * ggml_element_size(v), 0); + + // RoPE (NEOX / HF default half rotation), base 1e11; ne2 = n_tokens = 1. + ggml_tensor * pos_t = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, 1); + pos_t->data = &state.pos_value; + ggml_tensor * q_r = ggml_rope_ext(ctx, q4, pos_t, nullptr, head_dim, + GGML_ROPE_TYPE_NEOX, config.text.max_seq_len, + rope_base, 1.0F, 1.0F, 1.0F, 32.0F, 1.0F); + ggml_tensor * k_r = ggml_rope_ext(ctx, k4, pos_t, nullptr, head_dim, + GGML_ROPE_TYPE_NEOX, config.text.max_seq_len, + rope_base, 1.0F, 1.0F, 1.0F, 32.0F, 1.0F); + + // [head_dim, n_head, 1, 1] -> permute(0,2,1,3) -> [head_dim, 1, n_head, 1] + ggml_tensor * q_p = ggml_permute(ctx, q_r, 0, 2, 1, 3); + ggml_tensor * k_p = ggml_permute(ctx, k_r, 0, 2, 1, 3); + ggml_tensor * v_p = ggml_permute(ctx, v4, 0, 2, 1, 3); + + // Padded caches as external leaves [head_dim, cap, n_kv, 1]; the fresh + // k/v land in slot `kv_slot` via in-graph set_rows, and flash attends + // the whole padded cache under the mask (slots > seq are -inf, which + // contributes exactly zero to the softmax). + ggml_tensor * k_pad_t = ggml_new_tensor_4d(ctx, GGML_TYPE_F32, head_dim, cap, n_kv, 1); + ggml_tensor * v_pad_t = ggml_new_tensor_4d(ctx, GGML_TYPE_F32, head_dim, cap, n_kv, 1); + k_pad_t->data = state.k_pad[static_cast(li)].data(); + v_pad_t->data = state.v_pad[static_cast(li)].data(); + ggml_tensor * slot_t = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, 1); + slot_t->data = &state.kv_slot; + kv_writebacks.push_back(ggml_set_rows(ctx, k_pad_t, k_p, slot_t)); + kv_writebacks.push_back(ggml_set_rows(ctx, v_pad_t, v_p, slot_t)); + + // Single-token causal: all unmasked cache slots are visible. + ggml_tensor * attn = ggml_flash_attn_ext(ctx, q_p, k_pad_t, v_pad_t, mask_t, + 1.0F / std::sqrt(static_cast(head_dim)), + 0.0F, 0.0F); + ggml_flash_attn_ext_set_prec(attn, GGML_PREC_F32); + ggml_tensor * attn_flat = ggml_cont(ctx, ggml_reshape_1d(ctx, attn, n_head * head_dim)); + ggml_tensor * attn_out = ggml_mul_mat(ctx, layer.attn_o_proj.tensor, attn_flat); // [dim] + + // ---- merge + residual ---- + ggml_tensor * h = ggml_add(ctx, out_mamba, attn_out); + h = ggml_add(ctx, cur, h); + + // ---- FFN ---- + ggml_tensor * pre_w = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, dim); + bind_const(pre_w, layer.pre_ff_layernorm.values, state.ones_dim, 1.0F); + ggml_tensor * h2 = ggml_rms_norm(ctx, h, norm_eps); + h2 = ggml_mul(ctx, h2, pre_w); + ggml_tensor * gate_ff = ggml_mul_mat(ctx, layer.ffn_gate.tensor, h2); + ggml_tensor * up_ff = ggml_mul_mat(ctx, layer.ffn_up.tensor, h2); + ggml_tensor * gated = ggml_mul(ctx, ggml_silu(ctx, gate_ff), up_ff); + ggml_tensor * down = ggml_mul_mat(ctx, layer.ffn_down.tensor, gated); + cur = ggml_add(ctx, h, down); + } + + // final norm + semantic head + ggml_tensor * final_w = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, dim); + bind_const(final_w, weights.slow_norm.values, state.ones_dim, 1.0F); + ggml_tensor * final_norm = ggml_rms_norm(ctx, cur, norm_eps); + final_norm = ggml_mul(ctx, final_norm, final_w); + ggml_tensor * logits = ggml_mul_mat(ctx, weights.falcon_lm_head.tensor, final_norm); // [vocab] + // NOTE: lm_head_multiplier applies only to FalconH1ForCausalLM's full-vocab head. + // ArkttsModel uses the compact semantic_output head and does NOT scale logits. + plan->logits_out = ggml_dup(ctx, logits); + plan->hidden_out = ggml_dup(ctx, final_norm); + ggml_set_name(plan->logits_out, "logits_out"); + ggml_set_name(plan->hidden_out, "hidden_out"); + // Pin the outputs: the writeback nodes expanded after them allocate no + // memory of their own, but the flag keeps gallocr from ever reusing the + // output buffers while the plan is reused across steps. + ggml_set_output(plan->logits_out); + ggml_set_output(plan->hidden_out); + + ggml_cgraph * gf = ggml_new_graph_custom(ctx, 8192, false); + // KV slot writes first: they alias the cache the attention reads, so they + // must precede those nodes in graph order (backends execute sequentially). + for (ggml_tensor * wb : kv_writebacks) { + ggml_build_forward_expand(gf, wb); + } + ggml_build_forward_expand(gf, plan->logits_out); + ggml_build_forward_expand(gf, plan->hidden_out); + for (ggml_tensor * wb : writebacks) { + ggml_build_forward_expand(gf, wb); + } + plan->gf = gf; + auto t_graph = Clock::now(); + plan->gallocr = ggml_gallocr_new(ggml_backend_get_default_buffer_type(backend)); + if (!plan->gallocr || !ggml_gallocr_reserve(plan->gallocr, gf) || !ggml_gallocr_alloc_graph(plan->gallocr, gf)) { + throw std::runtime_error("build_falcon_step_plan: gallocr failed"); + } + auto t_alloc = Clock::now(); + if (profile != nullptr) { + profile->falcon_step_init_ms += engine::debug::elapsed_ms(t_init, t_build); + profile->falcon_step_build_ms += engine::debug::elapsed_ms(t_build, t_graph); + profile->falcon_step_gallocr_ms += engine::debug::elapsed_ms(t_graph, t_alloc); + } + return plan; +} + +// Zero-copy runner: feed the staged inputs, run the baked bucket graph, read +// logits/hidden. No per-step graph construction, allocation, or state I/O. +SlowForwardOutput falcon_forward_step_zero_copy( ggml_backend_t backend, int threads, size_t arena_bytes, const Audio8TtsConfig & config, const ArkttsARWeights & weights, - const std::vector & embeddings, - int64_t seq_len) { - if (seq_len <= 0) throw std::runtime_error("falcon_forward: zero seq"); - if (weights.falcon_layers.empty()) throw std::runtime_error("falcon_forward: no falcon layers"); + const std::vector & embedding, + FalconH1StepState & state, + int64_t position, + ArkttsARProfile * profile) { const int64_t dim = config.text.dim; - const float eps = config.text.norm_eps; - const float lm_mult = config.text.lm_head_multiplier; + const int64_t head_dim = config.text.head_dim; + const int64_t n_kv = config.text.n_local_heads; + const int64_t vocab = config.fast.vocab_size + 1; + const int64_t seq = state.seq_len; + + precompute_falcon_constants(config, weights, state); + grow_falcon_kv_pad(state, head_dim, n_kv, seq + 1); + if (state.emb_stage.empty()) { + state.emb_stage.assign(static_cast(dim), 0.0F); + } + if (!state.plan || state.plan->cap != state.kv_cap) { + state.plan = build_falcon_step_plan(backend, arena_bytes, config, weights, state, profile); + } + auto t_feed = Clock::now(); + std::memcpy(state.emb_stage.data(), embedding.data(), embedding.size() * sizeof(float)); + state.pos_value = static_cast(position); + state.kv_slot = static_cast(seq); + state.kv_mask[static_cast(seq)] = ggml_fp32_to_fp16(0.0F); + auto t_compute = Clock::now(); + core::set_backend_threads(backend, threads); + const ggml_status status = core::compute_backend_graph(backend, state.plan->gf, nullptr, "falcon_forward_step"); + ggml_backend_synchronize(backend); + auto t_read = Clock::now(); + if (status != GGML_STATUS_SUCCESS) { + throw std::runtime_error("falcon_forward_step compute failed"); + } + SlowForwardOutput out; + out.logits.resize(static_cast(vocab)); + out.hidden.resize(static_cast(dim)); + ggml_backend_tensor_get(state.plan->logits_out, out.logits.data(), 0, out.logits.size() * sizeof(float)); + ggml_backend_tensor_get(state.plan->hidden_out, out.hidden.data(), 0, out.hidden.size() * sizeof(float)); + state.seq_len = seq + 1; + if (profile != nullptr) { + profile->falcon_step_upload_ms += engine::debug::elapsed_ms(t_feed, t_compute); + profile->falcon_step_compute_ms += engine::debug::elapsed_ms(t_compute, t_read); + profile->falcon_step_download_ms += engine::debug::elapsed_ms(t_read, Clock::now()); + profile->falcon_step_runs += 1; + } + return out; +} + +// Single-token Falcon-H1 forward. `embedding` is the pre-multiplied token +// embedding (text embedding * embedding_multiplier + codebook sum). +SlowForwardOutput falcon_forward_step( + ggml_backend_t backend, + int threads, + size_t arena_bytes, + const Audio8TtsConfig & config, + const ArkttsARWeights & weights, + const std::vector & embedding, // [dim] + FalconH1StepState & state, + int64_t position, + ArkttsARProfile * profile = nullptr) { + const int64_t dim = config.text.dim; + const int64_t n_layer = config.text.n_layer; + const int64_t d_inner = config.text.mamba_d_ssm; + const int64_t d_state = config.text.mamba_d_state; + const int64_t d_conv = config.text.mamba_d_conv; + const int64_t n_groups = config.text.mamba_n_groups; + const int64_t n_mamba_heads = config.text.mamba_n_heads; + const int64_t mamba_head_dim = config.text.mamba_d_head; + const int64_t conv_dim = d_inner + 2 * n_groups * d_state; + const int64_t n_head = config.text.n_head; + const int64_t n_kv = config.text.n_local_heads; + const int64_t head_dim = config.text.head_dim; + const float norm_eps = config.text.norm_eps; + const float rope_base = config.text.rope_base; + const int64_t vocab = config.fast.vocab_size + 1; + const int64_t seq = state.seq_len; + + if (embedding.size() != static_cast(dim)) { + throw std::runtime_error("falcon_forward_step: embedding size mismatch"); + } + + // Host backends run the reusable zero-copy plan graph (one baked graph per + // KV capacity bucket; weights/state bound as external host views, state + // write-backs in-graph). Other backends keep the per-call explicit + // upload/download path below. + const bool zero_copy = core::backend_type(backend) == core::BackendType::Cpu; + if (zero_copy) { + return falcon_forward_step_zero_copy(backend, threads, arena_bytes, config, weights, + embedding, state, position, profile); + } + + auto t_init = Clock::now(); ggml_init_params params{arena_bytes, nullptr, true}; std::unique_ptr ctx(ggml_init(params)); - if (!ctx) throw std::runtime_error("falcon_forward: ggml_init failed"); - ggml_tensor * cur = ggml_new_tensor_2d(ctx.get(), GGML_TYPE_F32, dim, seq_len); + if (!ctx) throw std::runtime_error("falcon_forward_step: ggml_init failed"); + auto t_build = Clock::now(); + + ggml_tensor * cur = ggml_new_tensor_1d(ctx.get(), GGML_TYPE_F32, dim); ggml_set_name(cur, "falcon_input"); - std::vector ln_ws; - std::vector bias_ws; - std::vector pre_ws; - ln_ws.reserve(weights.falcon_layers.size()); - bias_ws.reserve(weights.falcon_layers.size()); - pre_ws.reserve(weights.falcon_layers.size()); - for (size_t li = 0; li < weights.falcon_layers.size(); ++li) { - const auto & layer = weights.falcon_layers[li]; + if (zero_copy) { + cur->data = const_cast(embedding.data()); + } else { + ggml_set_input(cur); + } + + // Bind a loop-invariant F32 tensor to its host values with zero copies when + // the backend reads host memory directly; otherwise mark it for the upload + // pass below. `fallback` (filled once with `fallback_fill`) covers weights + // that are absent from the checkpoint. + auto bind_const = [&](ggml_tensor * t, const std::vector & values, + std::vector & fallback, float fallback_fill) { + if (!zero_copy) { + ggml_set_input(t); + return; + } + if (!values.empty()) { + t->data = const_cast(values.data()); + return; + } + if (fallback.empty()) { + fallback.assign(static_cast(ggml_nelements(t)), fallback_fill); + } + t->data = fallback.data(); + }; + auto bind_state = [&](ggml_tensor * t, const std::vector & values) { + if (!zero_copy) { + ggml_set_input(t); + return; + } + t->data = const_cast(values.data()); + }; + std::vector writebacks; + writebacks.reserve(static_cast(n_layer) * 2); + // KV slot writes have no data dependency protecting them from the + // attention read of the same cache, so they are expanded into the graph + // FIRST below — backends execute nodes sequentially in graph order. + std::vector kv_writebacks; + kv_writebacks.reserve(static_cast(n_layer) * 2); + + std::vector ln_w_ts; + std::vector pre_w_ts; + std::vector conv_st_ts; + std::vector ssm_st_ts; + std::vector k_cur_ts; + std::vector v_cur_ts; + std::vector conv_b_ts; + std::vector conv_w2_ts; + std::vector k_cache_ts; + std::vector v_cache_ts; + std::vector sx_ts; + std::vector scan_ts; + std::vector ids_ts; + std::vector pos_ts; + ln_w_ts.reserve(static_cast(n_layer)); + pre_w_ts.reserve(static_cast(n_layer)); + conv_st_ts.reserve(static_cast(n_layer)); + ssm_st_ts.reserve(static_cast(n_layer)); + k_cur_ts.reserve(static_cast(n_layer)); + v_cur_ts.reserve(static_cast(n_layer)); + conv_b_ts.reserve(static_cast(n_layer)); + conv_w2_ts.reserve(static_cast(n_layer)); + k_cache_ts.reserve(static_cast(n_layer)); + v_cache_ts.reserve(static_cast(n_layer)); + sx_ts.reserve(static_cast(n_layer)); + scan_ts.reserve(static_cast(n_layer)); + ids_ts.reserve(static_cast(n_layer)); + pos_ts.reserve(static_cast(n_layer)); + + // A = -exp(A_log) per layer + std::vector A_ts; + A_ts.reserve(static_cast(n_layer)); + // D expanded to [d_inner]: D[h] repeated mamba_head_dim times + std::vector D_ts; + + for (int64_t li = 0; li < n_layer; ++li) { + const auto & layer = weights.falcon_layers[static_cast(li)]; + + // input_layernorm (RMS) ggml_tensor * ln_w = ggml_new_tensor_1d(ctx.get(), GGML_TYPE_F32, dim); - ln_ws.push_back(ln_w); - ggml_tensor * normed = ggml_rms_norm(ctx.get(), cur, eps); + bind_const(ln_w, layer.input_layernorm.values, state.ones_dim, 1.0F); + ln_w_ts.push_back(ln_w); + ggml_tensor * normed = ggml_rms_norm(ctx.get(), cur, norm_eps); normed = ggml_mul(ctx.get(), normed, ln_w); - ggml_tensor * proj = ggml_mul_mat(ctx.get(), layer.ssm_in.tensor, normed); - ggml_tensor * gate = ggml_view_2d(ctx.get(), proj, 768, seq_len, proj->nb[1], 0); - ggml_tensor * xBC = ggml_view_2d(ctx.get(), proj, 896, seq_len, proj->nb[1], 768 * sizeof(float)); - ggml_tensor * bias = ggml_new_tensor_1d(ctx.get(), GGML_TYPE_F32, 896); - bias_ws.push_back(bias); - ggml_tensor * bias_bcast = ggml_repeat(ctx.get(), bias, xBC); - ggml_tensor * xBC_b = ggml_add(ctx.get(), xBC, bias_bcast); - ggml_tensor * xBC_silu = ggml_silu(ctx.get(), xBC_b); - ggml_tensor * x = ggml_view_2d(ctx.get(), xBC_silu, 768, seq_len, xBC_silu->nb[1], 0); - ggml_tensor * gate_silu = ggml_silu(ctx.get(), gate); - ggml_tensor * y_gated = ggml_mul(ctx.get(), x, gate_silu); - ggml_tensor * out_mamba = ggml_mul_mat(ctx.get(), layer.ssm_out.tensor, y_gated); - if (std::abs(config.text.ssm_out_multiplier - 1.0f) > 1e-6) out_mamba = ggml_scale(ctx.get(), out_mamba, config.text.ssm_out_multiplier); - ggml_tensor * attn_out = ggml_scale(ctx.get(), cur, 0.0f); - ggml_tensor * hybrid = ggml_add(ctx.get(), out_mamba, attn_out); - ggml_tensor * cur_res = ggml_add(ctx.get(), cur, hybrid); + + // ---- Mamba2 branch ---- + // zxBCdt = in_proj(normed) -> [d_inner + conv_dim + n_mamba_heads] + ggml_tensor * zxBCdt = ggml_mul_mat(ctx.get(), layer.ssm_in.tensor, normed); + ggml_tensor * z = ggml_view_1d(ctx.get(), zxBCdt, d_inner, 0); + ggml_tensor * xBC = ggml_view_1d(ctx.get(), zxBCdt, conv_dim, d_inner * ggml_element_size(zxBCdt)); + ggml_tensor * dt = ggml_view_1d(ctx.get(), zxBCdt, n_mamba_heads, (d_inner + conv_dim) * ggml_element_size(zxBCdt)); + + // conv: state (d_conv-1 rows) + current xBC -> [d_conv, conv_dim, 1] + ggml_tensor * st_t = ggml_new_tensor_2d(ctx.get(), GGML_TYPE_F32, conv_dim, d_conv - 1); + bind_state(st_t, state.conv_states[static_cast(li)]); + conv_st_ts.push_back(st_t); + ggml_tensor * stT = ggml_cont(ctx.get(), ggml_transpose(ctx.get(), st_t)); // [d_conv-1, conv_dim] + ggml_tensor * xBC_r = ggml_cont(ctx.get(), ggml_transpose(ctx.get(), ggml_reshape_2d(ctx.get(), xBC, conv_dim, 1))); // [1, conv_dim] + ggml_tensor * sx = ggml_concat(ctx.get(), stT, xBC_r, 0); // [d_conv, conv_dim] + sx_ts.push_back(sx); + if (zero_copy) { + // Next conv state = sx rows 1..d_conv-1, written back into the host + // state vector in-graph. The cpy depends on sx (hence on the cont() + // that consumed st_t), so the old state is read before it is + // overwritten on every backend. + ggml_tensor * tail = ggml_view_2d(ctx.get(), sx, d_conv - 1, conv_dim, + sx->nb[1], ggml_element_size(sx)); + writebacks.push_back(ggml_cpy(ctx.get(), ggml_transpose(ctx.get(), tail), st_t)); + } else { + ggml_set_output(sx); // host reads the conv window back — pin the buffer + } + ggml_tensor * sx3 = ggml_reshape_3d(ctx.get(), sx, d_conv, conv_dim, 1); + // ggml ssm_conv computes y[c] = sum_k w[k,c]*sx[k,c] with sx row 0 the oldest + // frame — the same orientation as the HF conv1d/cached decode, so the GGUF + // kernel is fed as loaded (see load_falcon_layer). + ggml_tensor * conv_w2 = ggml_new_tensor_2d(ctx.get(), GGML_TYPE_F32, d_conv, conv_dim); + bind_const(conv_w2, layer.conv1d_kernel.values, state.zeros_conv_kernel, 0.0F); + conv_w2_ts.push_back(conv_w2); + ggml_tensor * xBC_conv = ggml_ssm_conv(ctx.get(), sx3, conv_w2); // [conv_dim, 1, 1] + ggml_tensor * conv_b = ggml_new_tensor_1d(ctx.get(), GGML_TYPE_F32, conv_dim); + bind_const(conv_b, layer.ssm_conv1d_b.values, state.zeros_conv, 0.0F); + conv_b_ts.push_back(conv_b); + xBC_conv = ggml_add(ctx.get(), xBC_conv, ggml_reshape_3d(ctx.get(), conv_b, conv_dim, 1, 1)); + xBC_conv = ggml_silu(ctx.get(), xBC_conv); + + // split x / B / C (conv output is contiguous [conv_dim,1,1]) + ggml_tensor * x = ggml_view_1d(ctx.get(), xBC_conv, d_inner, 0); + ggml_tensor * B = ggml_view_1d(ctx.get(), xBC_conv, d_state * n_groups, d_inner * ggml_element_size(xBC_conv)); + ggml_tensor * C = ggml_view_1d(ctx.get(), xBC_conv, d_state * n_groups, (d_inner + d_state * n_groups) * ggml_element_size(xBC_conv)); + + // x -> [head_dim, n_mamba_heads, 1, 1] + ggml_tensor * x4 = ggml_view_4d(ctx.get(), x, mamba_head_dim, n_mamba_heads, 1, 1, + mamba_head_dim * ggml_element_size(x), + mamba_head_dim * n_mamba_heads * ggml_element_size(x), + mamba_head_dim * n_mamba_heads * ggml_element_size(x), 0); + ggml_tensor * B4 = ggml_view_4d(ctx.get(), B, d_state, n_groups, 1, 1, + d_state * ggml_element_size(B), + d_state * n_groups * ggml_element_size(B), + d_state * n_groups * ggml_element_size(B), 0); + ggml_tensor * C4 = ggml_view_4d(ctx.get(), C, d_state, n_groups, 1, 1, + d_state * ggml_element_size(C), + d_state * n_groups * ggml_element_size(C), + d_state * n_groups * ggml_element_size(C), 0); + + // dt = dt + dt_bias -> [n_mamba_heads, 1, 1] + ggml_tensor * dt_eff = ggml_add(ctx.get(), dt, layer.ssm_dt_b.tensor); + ggml_tensor * dt3 = ggml_view_3d(ctx.get(), dt_eff, n_mamba_heads, 1, 1, + n_mamba_heads * ggml_element_size(dt_eff), + n_mamba_heads * ggml_element_size(dt_eff), 0); + + // A = -exp(A_log): [1, n_mamba_heads] — precomputed once per generation + // into state.pre_A (zero-copy) or uploaded per step below. + ggml_tensor * A_t = ggml_new_tensor_2d(ctx.get(), GGML_TYPE_F32, 1, n_mamba_heads); + if (zero_copy) { + A_t->data = state.pre_A[static_cast(li)].data(); + } else { + ggml_set_input(A_t); + } + A_ts.push_back(A_t); + + // ssm state: [d_state, mamba_head_dim, n_mamba_heads] + ggml_tensor * ssm_t = ggml_new_tensor_3d(ctx.get(), GGML_TYPE_F32, d_state, mamba_head_dim, n_mamba_heads); + bind_state(ssm_t, state.ssm_states[static_cast(li)]); + ssm_st_ts.push_back(ssm_t); + + // ids for scan (1 sequence) + ggml_tensor * ids = ggml_new_tensor_1d(ctx.get(), GGML_TYPE_I32, 1); + if (zero_copy) { + ids->data = &state.ids_value; + } else { + ggml_set_input(ids); + } + ids_ts.push_back(ids); + + ggml_tensor * scan = ggml_ssm_scan(ctx.get(), ssm_t, x4, dt3, A_t, B4, C4, ids); + scan_ts.push_back(scan); + if (zero_copy) { + // New ssm state = scan tail (after the d_inner y values), written + // back into the host state vector in-graph. The cpy depends on the + // scan that read ssm_t, so the old state is fully consumed first. + ggml_tensor * next_state = ggml_view_3d(ctx.get(), scan, + d_state, mamba_head_dim, n_mamba_heads, + d_state * ggml_element_size(scan), + d_state * mamba_head_dim * ggml_element_size(scan), + d_inner * ggml_element_size(scan)); + writebacks.push_back(ggml_cpy(ctx.get(), next_state, ssm_t)); + } else { + ggml_set_output(scan); // keep state tail alive for host read-back + } + ggml_tensor * y = ggml_view_1d(ctx.get(), scan, d_inner, 0); + + // y += x * D (D precomputed per generation / uploaded per step) + ggml_tensor * D_t = ggml_new_tensor_1d(ctx.get(), GGML_TYPE_F32, d_inner); + if (zero_copy) { + D_t->data = state.pre_D[static_cast(li)].data(); + } else { + ggml_set_input(D_t); + } + D_ts.push_back(D_t); + y = ggml_add(ctx.get(), y, ggml_mul(ctx.get(), ggml_view_1d(ctx.get(), x4, d_inner, 0), D_t)); + + // z gate: y *= silu(z) + y = ggml_mul(ctx.get(), y, ggml_silu(ctx.get(), z)); + + ggml_tensor * out_mamba = ggml_mul_mat(ctx.get(), layer.ssm_out.tensor, y); // [dim] + + // ---- GQA attention branch ---- + // ggml_flash_attn_ext layout: [head_dim, n_tokens, n_head, batch]. + // Project -> [head_dim, n_head, 1, 1] -> RoPE (ne2 = tokens) -> permute(0,2,1,3) + // -> [head_dim, 1, n_head, 1]; KV cache kept in the same layout. + ggml_tensor * q = ggml_mul_mat(ctx.get(), layer.attn_q_proj.tensor, normed); // [n_head*head_dim] + ggml_tensor * k = ggml_mul_mat(ctx.get(), layer.attn_k_proj.tensor, normed); // [n_kv*head_dim] + ggml_tensor * v = ggml_mul_mat(ctx.get(), layer.attn_v_proj.tensor, normed); // [n_kv*head_dim] + + ggml_tensor * q4 = ggml_view_4d(ctx.get(), q, head_dim, n_head, 1, 1, + head_dim * ggml_element_size(q), + head_dim * n_head * ggml_element_size(q), + head_dim * n_head * ggml_element_size(q), 0); + ggml_tensor * k4 = ggml_view_4d(ctx.get(), k, head_dim, n_kv, 1, 1, + head_dim * ggml_element_size(k), + head_dim * n_kv * ggml_element_size(k), + head_dim * n_kv * ggml_element_size(k), 0); + ggml_tensor * v4 = ggml_view_4d(ctx.get(), v, head_dim, n_kv, 1, 1, + head_dim * ggml_element_size(v), + head_dim * n_kv * ggml_element_size(v), + head_dim * n_kv * ggml_element_size(v), 0); + + // RoPE (NEOX / HF default half rotation), base 1e11; ne2 = n_tokens = 1. + ggml_tensor * pos_t = ggml_new_tensor_1d(ctx.get(), GGML_TYPE_I32, 1); + if (zero_copy) { + pos_t->data = &state.pos_value; + } else { + ggml_set_input(pos_t); + } + pos_ts.push_back(pos_t); + ggml_tensor * q_r = ggml_rope_ext(ctx.get(), q4, pos_t, nullptr, head_dim, + GGML_ROPE_TYPE_NEOX, config.text.max_seq_len, + rope_base, 1.0F, 1.0F, 1.0F, 32.0F, 1.0F); + ggml_tensor * k_r = ggml_rope_ext(ctx.get(), k4, pos_t, nullptr, head_dim, + GGML_ROPE_TYPE_NEOX, config.text.max_seq_len, + rope_base, 1.0F, 1.0F, 1.0F, 32.0F, 1.0F); + // k_p/v_p (and their k_r/v bases) are read back on the host for the KV + // cache on the fallback path — pin the underlying buffers there. + if (!zero_copy) { + ggml_set_output(k_r); + ggml_set_output(v); + } + + // [head_dim, n_head, 1, 1] -> permute(0,2,1,3) -> [head_dim, 1, n_head, 1] + ggml_tensor * q_p = ggml_permute(ctx.get(), q_r, 0, 2, 1, 3); + ggml_tensor * k_p = ggml_permute(ctx.get(), k_r, 0, 2, 1, 3); + ggml_tensor * v_p = ggml_permute(ctx.get(), v4, 0, 2, 1, 3); + k_cur_ts.push_back(k_p); + v_cur_ts.push_back(v_p); + + // KV cache in flash layout: [head_dim, n_tokens, n_kv, 1] + ggml_tensor * K_all = nullptr; + ggml_tensor * V_all = nullptr; + if (zero_copy) { + // Padded caches bound as external leaves; the fresh k/v are written + // into slot `seq` in-graph (kv_writebacks are expanded before the + // attention nodes, ordering the writes first), and flash attends + // over a strided prefix view of slots [0, seq+1). CPU + // flash_attn_ext only requires contiguous rows (nb0), so the + // padded head stride is fine. + ggml_tensor * k_pad_t = ggml_new_tensor_4d(ctx.get(), GGML_TYPE_F32, head_dim, state.kv_cap, n_kv, 1); + ggml_tensor * v_pad_t = ggml_new_tensor_4d(ctx.get(), GGML_TYPE_F32, head_dim, state.kv_cap, n_kv, 1); + k_pad_t->data = state.k_pad[static_cast(li)].data(); + v_pad_t->data = state.v_pad[static_cast(li)].data(); + const size_t slot_off = static_cast(seq * head_dim) * ggml_element_size(k_pad_t); + kv_writebacks.push_back(ggml_cpy(ctx.get(), k_p, + ggml_view_4d(ctx.get(), k_pad_t, head_dim, 1, n_kv, 1, + k_pad_t->nb[1], k_pad_t->nb[2], k_pad_t->nb[3], slot_off))); + kv_writebacks.push_back(ggml_cpy(ctx.get(), v_p, + ggml_view_4d(ctx.get(), v_pad_t, head_dim, 1, n_kv, 1, + v_pad_t->nb[1], v_pad_t->nb[2], v_pad_t->nb[3], slot_off))); + K_all = ggml_view_4d(ctx.get(), k_pad_t, head_dim, seq + 1, n_kv, 1, + k_pad_t->nb[1], k_pad_t->nb[2], k_pad_t->nb[3], 0); + V_all = ggml_view_4d(ctx.get(), v_pad_t, head_dim, seq + 1, n_kv, 1, + v_pad_t->nb[1], v_pad_t->nb[2], v_pad_t->nb[3], 0); + } else { + ggml_tensor * k_cache_t = ggml_new_tensor_4d(ctx.get(), GGML_TYPE_F32, head_dim, seq, n_kv, 1); + ggml_tensor * v_cache_t = ggml_new_tensor_4d(ctx.get(), GGML_TYPE_F32, head_dim, seq, n_kv, 1); + ggml_set_input(k_cache_t); + ggml_set_input(v_cache_t); + k_cache_ts.push_back(k_cache_t); + v_cache_ts.push_back(v_cache_t); + K_all = ggml_concat(ctx.get(), k_cache_t, k_p, 1); // [head_dim, seq+1, n_kv, 1] + V_all = ggml_concat(ctx.get(), v_cache_t, v_p, 1); + } + + // Single-token causal: current query attends to all cached keys (all visible). + ggml_tensor * attn = ggml_flash_attn_ext(ctx.get(), q_p, K_all, V_all, nullptr, + 1.0F / std::sqrt(static_cast(head_dim)), + 0.0F, 0.0F); + ggml_flash_attn_ext_set_prec(attn, GGML_PREC_F32); + ggml_tensor * attn_flat = ggml_cont(ctx.get(), ggml_reshape_1d(ctx.get(), attn, n_head * head_dim)); + ggml_tensor * attn_out = ggml_mul_mat(ctx.get(), layer.attn_o_proj.tensor, attn_flat); // [dim] + + // ---- merge + residual ---- + ggml_tensor * h = ggml_add(ctx.get(), out_mamba, attn_out); + h = ggml_add(ctx.get(), cur, h); + + // ---- FFN ---- ggml_tensor * pre_w = ggml_new_tensor_1d(ctx.get(), GGML_TYPE_F32, dim); - pre_ws.push_back(pre_w); - ggml_tensor * pre_norm = ggml_rms_norm(ctx.get(), cur_res, eps); - pre_norm = ggml_mul(ctx.get(), pre_norm, pre_w); - ggml_tensor * gate_ff = ggml_mul_mat(ctx.get(), layer.ffn_gate.tensor, pre_norm); - ggml_tensor * up_ff = ggml_mul_mat(ctx.get(), layer.ffn_up.tensor, pre_norm); - ggml_tensor * gate_silu2 = ggml_silu(ctx.get(), gate_ff); - ggml_tensor * gated = ggml_mul(ctx.get(), gate_silu2, up_ff); + bind_const(pre_w, layer.pre_ff_layernorm.values, state.ones_dim, 1.0F); + pre_w_ts.push_back(pre_w); + ggml_tensor * h2 = ggml_rms_norm(ctx.get(), h, norm_eps); + h2 = ggml_mul(ctx.get(), h2, pre_w); + ggml_tensor * gate_ff = ggml_mul_mat(ctx.get(), layer.ffn_gate.tensor, h2); + ggml_tensor * up_ff = ggml_mul_mat(ctx.get(), layer.ffn_up.tensor, h2); + ggml_tensor * gated = ggml_mul(ctx.get(), ggml_silu(ctx.get(), gate_ff), up_ff); ggml_tensor * down = ggml_mul_mat(ctx.get(), layer.ffn_down.tensor, gated); - cur = ggml_add(ctx.get(), cur_res, down); + cur = ggml_add(ctx.get(), h, down); } + + // final norm + semantic head ggml_tensor * final_w = ggml_new_tensor_1d(ctx.get(), GGML_TYPE_F32, dim); - ggml_tensor * final_norm = ggml_rms_norm(ctx.get(), cur, eps); + bind_const(final_w, weights.slow_norm.values, state.ones_dim, 1.0F); + ggml_tensor * final_norm = ggml_rms_norm(ctx.get(), cur, norm_eps); final_norm = ggml_mul(ctx.get(), final_norm, final_w); - ggml_tensor * last_hidden = ggml_view_2d(ctx.get(), final_norm, dim, 1, final_norm->nb[1], (seq_len - 1) * final_norm->nb[1]); - ggml_tensor * logits = ggml_mul_mat(ctx.get(), weights.falcon_lm_head.tensor, last_hidden); - if (std::abs(lm_mult - 1.0f) > 1e-6) logits = ggml_scale(ctx.get(), logits, lm_mult); + ggml_tensor * logits = ggml_mul_mat(ctx.get(), weights.falcon_lm_head.tensor, final_norm); // [vocab] + // NOTE: lm_head_multiplier applies only to FalconH1ForCausalLM's full-vocab head. + // ArkttsModel uses the compact semantic_output head and does NOT scale logits. ggml_tensor * logits_out = ggml_dup(ctx.get(), logits); - ggml_tensor * hidden_out = ggml_dup(ctx.get(), last_hidden); + ggml_tensor * hidden_out = ggml_dup(ctx.get(), final_norm); ggml_set_name(logits_out, "logits_out"); ggml_set_name(hidden_out, "hidden_out"); + ggml_cgraph * gf = ggml_new_graph_custom(ctx.get(), 8192, false); + // KV slot writes first: they alias the cache views the attention reads, so + // they must precede those nodes in graph order. + for (ggml_tensor * wb : kv_writebacks) { + ggml_build_forward_expand(gf, wb); + } ggml_build_forward_expand(gf, logits_out); ggml_build_forward_expand(gf, hidden_out); + for (ggml_tensor * wb : writebacks) { + ggml_build_forward_expand(gf, wb); + } + auto t_graph = Clock::now(); ggml_gallocr_t gallocr = ggml_gallocr_new(ggml_backend_get_default_buffer_type(backend)); - if (!gallocr || !ggml_gallocr_reserve(gallocr, gf) || !ggml_gallocr_alloc_graph(gallocr, gf)) throw std::runtime_error("falcon_forward: gallocr failed"); - for (size_t i = 0; i < ln_ws.size(); ++i) { - const auto & vals = weights.falcon_layers[i].input_layernorm.values; - if (!vals.empty()) ggml_backend_tensor_set(ln_ws[i], vals.data(), 0, vals.size() * sizeof(float)); - else { std::vector ones(static_cast(dim), 1.0f); ggml_backend_tensor_set(ln_ws[i], ones.data(), 0, ones.size() * sizeof(float)); } - } - for (size_t i = 0; i < bias_ws.size(); ++i) { - const auto & vals = weights.falcon_layers[i].ssm_conv1d_b.values; - if (!vals.empty()) ggml_backend_tensor_set(bias_ws[i], vals.data(), 0, vals.size() * sizeof(float)); - else { std::vector zeros(896, 0.0f); ggml_backend_tensor_set(bias_ws[i], zeros.data(), 0, zeros.size() * sizeof(float)); } - } - for (size_t i = 0; i < pre_ws.size(); ++i) { - const auto & vals = weights.falcon_layers[i].pre_ff_layernorm.values; - if (!vals.empty()) ggml_backend_tensor_set(pre_ws[i], vals.data(), 0, vals.size() * sizeof(float)); - else { std::vector ones(static_cast(dim), 1.0f); ggml_backend_tensor_set(pre_ws[i], ones.data(), 0, ones.size() * sizeof(float)); } - } - if (!weights.slow_norm.values.empty()) ggml_backend_tensor_set(final_w, weights.slow_norm.values.data(), 0, weights.slow_norm.values.size() * sizeof(float)); - else { std::vector ones(static_cast(dim), 1.0f); ggml_backend_tensor_set(final_w, ones.data(), 0, ones.size() * sizeof(float)); } - std::vector cur_data(static_cast(dim * seq_len)); - for (int64_t s = 0; s < seq_len; ++s) for (int64_t d = 0; d < dim; ++d) cur_data[static_cast(d + s * dim)] = embeddings[static_cast(s * dim + d)]; - ggml_backend_tensor_set(cur, cur_data.data(), 0, cur_data.size() * sizeof(float)); + if (!gallocr || !ggml_gallocr_reserve(gallocr, gf) || !ggml_gallocr_alloc_graph(gallocr, gf)) { + throw std::runtime_error("falcon_forward_step: gallocr failed"); + } + auto t_alloc = Clock::now(); + + // ---- feed host constants ---- + if (zero_copy) { + // Everything except the rope position is bound to host memory already; + // the graph reads the position straight from the state struct. + state.pos_value = static_cast(position); + } else { + ggml_backend_tensor_set(cur, embedding.data(), 0, embedding.size() * sizeof(float)); + { + const int32_t ids0 = 0; + const int32_t posv = static_cast(position); + for (int64_t li = 0; li < n_layer; ++li) { + ggml_backend_tensor_set(ids_ts[static_cast(li)], &ids0, 0, sizeof(int32_t)); + ggml_backend_tensor_set(pos_ts[static_cast(li)], &posv, 0, sizeof(int32_t)); + } + } + for (int64_t li = 0; li < n_layer; ++li) { + const auto & layer = weights.falcon_layers[static_cast(li)]; + if (!layer.input_layernorm.values.empty()) { + ggml_backend_tensor_set(ln_w_ts[static_cast(li)], layer.input_layernorm.values.data(), 0, + layer.input_layernorm.values.size() * sizeof(float)); + } else { + std::vector ones(static_cast(dim), 1.0F); + ggml_backend_tensor_set(ln_w_ts[static_cast(li)], ones.data(), 0, ones.size() * sizeof(float)); + } + if (!layer.pre_ff_layernorm.values.empty()) { + ggml_backend_tensor_set(pre_w_ts[static_cast(li)], layer.pre_ff_layernorm.values.data(), 0, + layer.pre_ff_layernorm.values.size() * sizeof(float)); + } else { + std::vector ones(static_cast(dim), 1.0F); + ggml_backend_tensor_set(pre_w_ts[static_cast(li)], ones.data(), 0, ones.size() * sizeof(float)); + } + // conv state [d_conv-1, conv_dim] (col-major: element (c,r) at r*conv_dim + c) + const auto & cstate = state.conv_states[static_cast(li)]; + ggml_backend_tensor_set(conv_st_ts[static_cast(li)], cstate.data(), 0, cstate.size() * sizeof(float)); + // ssm state [d_state, mamba_head_dim, n_mamba_heads] + const auto & sstate = state.ssm_states[static_cast(li)]; + ggml_backend_tensor_set(ssm_st_ts[static_cast(li)], sstate.data(), 0, sstate.size() * sizeof(float)); + // A = -exp(A_log) + { + std::vector a_log(static_cast(n_mamba_heads)); + ggml_backend_tensor_get(layer.ssm_A.tensor, a_log.data(), 0, a_log.size() * sizeof(float)); + std::vector av(static_cast(n_mamba_heads)); + for (int64_t h = 0; h < n_mamba_heads; ++h) { + av[static_cast(h)] = -std::exp(a_log[static_cast(h)]); + } + ggml_backend_tensor_set(A_ts[static_cast(li)], av.data(), 0, av.size() * sizeof(float)); + } + // D expanded + { + std::vector d_raw(static_cast(n_mamba_heads)); + ggml_backend_tensor_get(layer.ssm_D.tensor, d_raw.data(), 0, d_raw.size() * sizeof(float)); + std::vector dv(static_cast(d_inner)); + for (int64_t h = 0; h < n_mamba_heads; ++h) { + for (int64_t d = 0; d < mamba_head_dim; ++d) { + dv[static_cast(d + h * mamba_head_dim)] = d_raw[static_cast(h)]; + } + } + ggml_backend_tensor_set(D_ts[static_cast(li)], dv.data(), 0, dv.size() * sizeof(float)); + } + // conv bias + if (!layer.ssm_conv1d_b.values.empty()) { + ggml_backend_tensor_set(conv_b_ts[static_cast(li)], layer.ssm_conv1d_b.values.data(), 0, + layer.ssm_conv1d_b.values.size() * sizeof(float)); + } else { + std::vector zeros(static_cast(conv_dim), 0.0F); + ggml_backend_tensor_set(conv_b_ts[static_cast(li)], zeros.data(), 0, zeros.size() * sizeof(float)); + } + // conv1d weight: host [d_conv, conv_dim] kernel (GGUF layout, no flip) + { + const auto & cw = layer.conv1d_kernel; + ggml_backend_tensor_set(conv_w2_ts[static_cast(li)], cw.values.data(), 0, cw.values.size() * sizeof(float)); + } + // KV cache + const auto & kc = state.k_cache[static_cast(li)]; + const auto & vc = state.v_cache[static_cast(li)]; + if (!kc.empty()) ggml_backend_tensor_set(k_cache_ts[static_cast(li)], kc.data(), 0, kc.size() * sizeof(float)); + if (!vc.empty()) ggml_backend_tensor_set(v_cache_ts[static_cast(li)], vc.data(), 0, vc.size() * sizeof(float)); + } + if (!weights.slow_norm.values.empty()) { + ggml_backend_tensor_set(final_w, weights.slow_norm.values.data(), 0, weights.slow_norm.values.size() * sizeof(float)); + } else { + std::vector ones(static_cast(dim), 1.0F); + ggml_backend_tensor_set(final_w, ones.data(), 0, ones.size() * sizeof(float)); + } + } + core::set_backend_threads(backend, threads); - ggml_status status = core::compute_backend_graph(backend, gf, nullptr, "falcon_forward"); + auto t_upload = Clock::now(); + ggml_status status = core::compute_backend_graph(backend, gf, nullptr, "falcon_forward_step"); ggml_backend_synchronize(backend); - if (status != GGML_STATUS_SUCCESS) throw std::runtime_error("falcon_forward compute failed"); + auto t_compute = Clock::now(); + if (status != GGML_STATUS_SUCCESS) { + ggml_gallocr_free(gallocr); + throw std::runtime_error("falcon_forward_step compute failed"); + } + + // ---- read outputs + update states ---- SlowForwardOutput out; - size_t vocab = static_cast(logits_out->ne[0]); - if (vocab == 0) vocab = 4097; - out.logits.resize(vocab); + out.logits.resize(static_cast(vocab)); out.hidden.resize(static_cast(dim)); - ggml_backend_tensor_get(logits_out, out.logits.data(), 0, vocab * sizeof(float)); + ggml_backend_tensor_get(logits_out, out.logits.data(), 0, static_cast(vocab) * sizeof(float)); ggml_backend_tensor_get(hidden_out, out.hidden.data(), 0, static_cast(dim) * sizeof(float)); + + // conv/ssm states: updated in-graph on host backends; on GPU backends the + // host reads the new state tails back and shifts them into the vectors. + if (!zero_copy) { + // conv state: last d_conv-1 kernel rows of sx (per layer). sx is col-major + // [d_conv, conv_dim], element (k, c) at k + d_conv*c. + { + std::vector sx_vals(static_cast(d_conv * conv_dim)); + for (int64_t li = 0; li < n_layer; ++li) { + ggml_backend_tensor_get(sx_ts[static_cast(li)], sx_vals.data(), 0, sx_vals.size() * sizeof(float)); + auto & cstate = state.conv_states[static_cast(li)]; + for (int64_t r = 0; r < d_conv - 1; ++r) { + for (int64_t c = 0; c < conv_dim; ++c) { + cstate[r * conv_dim + c] = sx_vals[(r + 1) + d_conv * c]; + } + } + } + } + + // ssm state: tail of scan output (d_state*d_inner per layer) — per-layer tensors! + { + const size_t y_sz = static_cast(d_inner); + const size_t s_sz = static_cast(d_state * d_inner); + std::vector scan_vals(y_sz + s_sz); + for (int64_t li = 0; li < n_layer; ++li) { + ggml_backend_tensor_get(scan_ts[static_cast(li)], scan_vals.data(), 0, scan_vals.size() * sizeof(float)); + auto & sstate = state.ssm_states[static_cast(li)]; + for (size_t i = 0; i < s_sz; ++i) { + sstate[i] = scan_vals[y_sz + i]; + } + } + } + } + + // KV cache append (fallback path only; the zero-copy path wrote the fresh + // k/v into the padded caches in-graph). ggml col-major layout + // [head_dim, seq, n_kv, 1]: element (d, t, h) at + // d + head_dim*(t + new_seq_len*h). The freshly projected/roped k/v read + // back as [d + head_dim*h] (128 values). + if (!zero_copy) { + std::vector kv(static_cast(n_kv * head_dim)); + for (int64_t li = 0; li < n_layer; ++li) { + ggml_backend_tensor_get(k_cur_ts[static_cast(li)], kv.data(), 0, kv.size() * sizeof(float)); + append_falcon_kv_token(state.k_cache[static_cast(li)], seq, n_kv, head_dim, kv.data()); + ggml_backend_tensor_get(v_cur_ts[static_cast(li)], kv.data(), 0, kv.size() * sizeof(float)); + append_falcon_kv_token(state.v_cache[static_cast(li)], seq, n_kv, head_dim, kv.data()); + } + } + state.seq_len = seq + 1; + + if (profile != nullptr) { + profile->falcon_step_init_ms += engine::debug::elapsed_ms(t_init, t_build); + profile->falcon_step_build_ms += engine::debug::elapsed_ms(t_build, t_graph); + profile->falcon_step_gallocr_ms += engine::debug::elapsed_ms(t_graph, t_alloc); + profile->falcon_step_upload_ms += engine::debug::elapsed_ms(t_alloc, t_upload); + profile->falcon_step_compute_ms += engine::debug::elapsed_ms(t_upload, t_compute); + profile->falcon_step_download_ms += engine::debug::elapsed_ms(t_compute, Clock::now()); + profile->falcon_step_runs += 1; + } + ggml_gallocr_free(gallocr); core::release_backend_graph_resources(backend, gf); return out; } +// Single-token Falcon embedding: text_embedding(semantic) * embedding_multiplier +// + sum(codebook_embeddings) for semantic tokens. `matrix` is [codebook_rows][steps]. +std::vector build_falcon_embedding_step( + const Audio8TtsConfig & config, + const ArkttsARWeights & weights, + const int32_t * matrix, + int64_t steps, + int64_t step) { + const int64_t hidden = config.text.dim; + const int64_t codebook_rows = config.fast.num_codebooks + 1; + const int32_t token = matrix[step]; + std::vector out(static_cast(hidden), 0.0F); + auto row = lookup_row(weights.text_embedding_host, token, hidden); + if (token >= config.semantic_start_token_id && token <= config.semantic_end_token_id) { + for (int64_t cb = 0; cb < config.fast.num_codebooks; ++cb) { + const int32_t code = matrix[(cb + 1) * steps + step]; + add_row(weights.codebook_embedding_host, cb * config.fast.vocab_size + code, hidden, row); + } + } + // HF _embed + _slow_backbone: (text_emb + codebook_sum) * embedding_multiplier. + for (auto & v : row) v *= config.text.embedding_multiplier; + std::copy(row.begin(), row.end(), out.begin()); + return out; +} + + } // namespace +// Copies a backend-resident weight tensor byte-for-byte (same ggml type and +// dimensions) so it can live on a second backend: the fast AR graph runs on a +// dedicated CPU backend while the slow path and codec stay on the GPU (the +// per-step fast AR submit+sync latency dominates on GPU backends, while CPU +// computes the same graph several times faster). q8_0/f32/f16 all copy +// losslessly — the point is backend placement, not conversion. +core::TensorValue schedule_tensor_copy( + ggml_context * dst_ctx, + const core::TensorValue & src) { + ggml_tensor * dst = ggml_new_tensor( + dst_ctx, src.tensor->type, ggml_n_dims(src.tensor), src.tensor->ne); + return core::wrap_tensor(dst, src.shape, src.type); +} + +void copy_tensor_bytes(const core::TensorValue & src, const core::TensorValue & dst) { + const size_t bytes = static_cast(ggml_nbytes(src.tensor)); + std::vector host(bytes); + ggml_backend_tensor_get(src.tensor, host.data(), 0, bytes); + ggml_backend_tensor_set(dst.tensor, host.data(), 0, bytes); +} + class ArkttsARWeightsRuntime { public: ArkttsARWeightsRuntime( @@ -993,15 +2024,35 @@ class ArkttsARWeightsRuntime { backend_config.threads = threads_; backend_ = core::init_backend(backend_config); backend_type_ = core::backend_type(backend_); - weights_ = std::make_shared( - load_ar_weights(*assets_, backend_, backend_type_, weight_context_bytes, weight_storage_type)); + ArkttsARWeights loaded = + load_ar_weights(*assets_, backend_, backend_type_, weight_context_bytes, weight_storage_type); + if (backend_type_ != core::BackendType::Cpu) { + // Fast AR is submit+sync latency bound on GPU backends (one graph + // submission per generated codebook token); the same graph computes + // several times faster on CPU. Give it a dedicated CPU backend and + // move the fast-layer weights over, leaving slow path + codec on + // the GPU backend. + core::BackendConfig fast_backend_config; + fast_backend_config.type = core::BackendType::Cpu; + fast_backend_config.threads = threads_; + fast_backend_ = core::init_backend(fast_backend_config); + fast_backend_type_ = core::backend_type(fast_backend_); + retarget_fast_weights(loaded); + if (!loaded.falcon_layers.empty()) { + // The Falcon-H1 per-token step graph is ~600 tiny nodes; on Metal it + // is dispatch-latency bound (measured 5.9 ms/step vs 2.0 ms on CPU, + // same weights). Run it on the same dedicated CPU backend. + retarget_falcon_weights(loaded); + } + } + weights_ = std::make_shared(std::move(loaded)); slow_step_constants_ = std::make_unique( backend_, threads_, "audio8_tts.ar.step.constants", 256ull * 1024ull * 1024ull); fast_constants_ = std::make_unique( - backend_, + fast_backend(), threads_, "audio8_tts.ar.fast.constants", 256ull * 1024ull * 1024ull); @@ -1011,6 +2062,17 @@ class ArkttsARWeightsRuntime { fast_constants_.reset(); slow_step_constants_.reset(); weights_.reset(); + if (falcon_weight_buffer_ != nullptr) { + ggml_backend_buffer_free(falcon_weight_buffer_); + } + falcon_weight_ctx_.reset(); + if (fast_weight_buffer_ != nullptr) { + ggml_backend_buffer_free(fast_weight_buffer_); + } + fast_weight_ctx_.reset(); + if (fast_backend_ != nullptr) { + ggml_backend_free(fast_backend_); + } if (backend_ != nullptr) { ggml_backend_free(backend_); } @@ -1043,6 +2105,23 @@ class ArkttsARWeightsRuntime { return backend_type_; } + // Backend hosting the fast AR graph: a dedicated CPU backend when the main + // backend is a GPU, otherwise the main backend itself. + ggml_backend_t fast_backend() const noexcept { + return fast_backend_ != nullptr ? fast_backend_ : backend_; + } + + // Falcon-H1 per-token steps run on the dedicated CPU backend when the main + // backend is a GPU one (dispatch-latency bound there); identical tensors + // otherwise, so this is always the right backend for falcon_forward_step. + ggml_backend_t falcon_step_backend() const noexcept { + return fast_backend_ != nullptr ? fast_backend_ : backend_; + } + + core::BackendType fast_backend_type() const noexcept { + return fast_backend_ != nullptr ? fast_backend_type_ : backend_type_; + } + core::ConstantTensorCache & slow_step_constants() const noexcept { return *slow_step_constants_; } @@ -1052,12 +2131,143 @@ class ArkttsARWeightsRuntime { } private: + // Re-binds the fast AR layer weights (and fast_output) onto the dedicated + // CPU fast backend. The projections are byte copies of the tensors the + // weight store uploaded to the main backend; norms stay host TensorData + // and upload through the (CPU-backed) fast constants cache at graph build. + void retarget_fast_weights(ArkttsARWeights & weights) { + ggml_init_params params{8ull * 1024ull * 1024ull, nullptr, true}; + fast_weight_ctx_.reset(ggml_init(params)); + if (fast_weight_ctx_ == nullptr) { + throw std::runtime_error("failed to initialize Audio8 TTS fast AR CPU weight context"); + } + struct ScheduledCopy { + const core::TensorValue * source; + core::TensorValue target; + }; + std::vector copies; + copies.reserve(weights.fast_layers.size() * 5 + 1); + auto schedule = [&](const core::TensorValue & value) { + if (!value.valid()) { + throw std::runtime_error("Audio8 TTS fast AR weight tensor is missing"); + } + copies.push_back({&value, schedule_tensor_copy(fast_weight_ctx_.get(), value)}); + }; + for (const auto & layer : weights.fast_layers) { + schedule(layer.qkv_proj); + if (layer.qkv_bias.has_value()) { + schedule(*layer.qkv_bias); + } + schedule(layer.o_proj); + schedule(layer.gate_up_proj); + schedule(layer.down_proj); + } + schedule(weights.fast_output); + fast_weight_buffer_ = ggml_backend_alloc_ctx_tensors(fast_weight_ctx_.get(), fast_backend_); + if (fast_weight_buffer_ == nullptr) { + throw std::runtime_error("failed to allocate Audio8 TTS fast AR CPU weights"); + } + for (auto & copy : copies) { + copy_tensor_bytes(*copy.source, copy.target); + } + size_t index = 0; + auto commit = [&](core::TensorValue & value) { + if (!value.valid()) { + return; + } + value = std::move(copies[index].target); + ++index; + }; + for (auto & layer : weights.fast_layers) { + commit(layer.qkv_proj); + if (layer.qkv_bias.has_value()) { + commit(*layer.qkv_bias); + } + commit(layer.o_proj); + commit(layer.gate_up_proj); + commit(layer.down_proj); + } + commit(weights.fast_output); + } + + // Re-binds the Falcon-H1 layer weights (and the semantic head) onto the + // dedicated CPU backend, mirroring retarget_fast_weights. ssm_A / ssm_D are + // included so the per-step A=-exp(A_log) / D-expansion reads become plain + // CPU memcpys instead of GPU->host syncs. + void retarget_falcon_weights(ArkttsARWeights & weights) { + ggml_init_params params{8ull * 1024ull * 1024ull, nullptr, true}; + falcon_weight_ctx_.reset(ggml_init(params)); + if (falcon_weight_ctx_ == nullptr) { + throw std::runtime_error("failed to initialize Audio8 TTS Falcon step CPU weight context"); + } + struct ScheduledCopy { + const core::TensorValue * source; + core::TensorValue target; + }; + std::vector copies; + copies.reserve(weights.falcon_layers.size() * 12 + 1); + auto schedule = [&](const core::TensorValue & value) { + if (!value.valid()) { + throw std::runtime_error("Audio8 TTS Falcon step weight tensor is missing"); + } + copies.push_back({&value, schedule_tensor_copy(falcon_weight_ctx_.get(), value)}); + }; + for (const auto & layer : weights.falcon_layers) { + schedule(layer.ssm_in); + schedule(layer.ssm_dt_b); + schedule(layer.ssm_A); + schedule(layer.ssm_D); + schedule(layer.ssm_out); + schedule(layer.attn_q_proj); + schedule(layer.attn_k_proj); + schedule(layer.attn_v_proj); + schedule(layer.attn_o_proj); + schedule(layer.ffn_gate); + schedule(layer.ffn_up); + schedule(layer.ffn_down); + } + schedule(weights.falcon_lm_head); + falcon_weight_buffer_ = ggml_backend_alloc_ctx_tensors(falcon_weight_ctx_.get(), fast_backend_); + if (falcon_weight_buffer_ == nullptr) { + throw std::runtime_error("failed to allocate Audio8 TTS Falcon step CPU weights"); + } + for (auto & copy : copies) { + copy_tensor_bytes(*copy.source, copy.target); + } + size_t index = 0; + auto commit = [&](core::TensorValue & value) { + value = std::move(copies[index].target); + ++index; + }; + for (auto & layer : weights.falcon_layers) { + commit(layer.ssm_in); + commit(layer.ssm_dt_b); + commit(layer.ssm_A); + commit(layer.ssm_D); + commit(layer.ssm_out); + commit(layer.attn_q_proj); + commit(layer.attn_k_proj); + commit(layer.attn_v_proj); + commit(layer.attn_o_proj); + commit(layer.ffn_gate); + commit(layer.ffn_up); + commit(layer.ffn_down); + } + commit(weights.falcon_lm_head); + } + std::shared_ptr assets_; std::shared_ptr weights_; int threads_ = 1; size_t graph_arena_bytes_ = 0; ggml_backend_t backend_ = nullptr; core::BackendType backend_type_ = core::BackendType::Cpu; + ggml_backend_t fast_backend_ = nullptr; + core::BackendType fast_backend_type_ = core::BackendType::Cpu; + std::unique_ptr fast_weight_ctx_; + ggml_backend_buffer_t fast_weight_buffer_ = nullptr; + std::unique_ptr falcon_weight_ctx_; + ggml_backend_buffer_t falcon_weight_buffer_ = nullptr; std::unique_ptr slow_step_constants_; std::unique_ptr fast_constants_; }; @@ -1099,18 +2309,13 @@ class Audio8TtsARRuntime::Impl { const bool is_falcon = assets.config.text.slow_backbone == "falcon_h1" || assets.model_weights->has_tensor("slow.embed_tokens.weight"); if (is_falcon) { - // Falcon-H1 0.1B — native ggml path (see docs/FALCON_H1_0.1B_PORT_PLAN.md). - // Current limitation (drawback stub): falcon_forward_stateless is a - // simplified forward that implements RMSNorm + Mamba in_proj split - // (gate/xBC) + conv bias SiLU + gated out_proj + FFN, but stubs the - // SSM core (no ggml_ssm_conv / B/C / dt / A / D / ggml_ssm_scan / - // recurrent conv/ssm state, no hybrid attention). It recomputes the - // full sequence each step O(N^2) and only applies ssm_out/lm_head - // multipliers. This produces prompt-invariant logits and fails STT - // without the full Mamba2 port (see mamba-base.cpp:151, - // falcon-h1.cpp:132). The full port is tracked in the plan file and - // reuses vendored external/ggml ssm backends (cpu/cuda/metal/vulkan) - // — no Python dependency, no /tmp or system() calls. + // Falcon-H1 0.1B — native ggml path. Stateful Mamba2 + hybrid GQA + // attention single-token forward (falcon_forward_step), mirroring + // transformers.models.falcon_h1 FalconH1DecoderLayer and llama.cpp + // mamba-base.cpp build_mamba2_layer. Prefill runs each prompt token + // through the step graph to populate conv/ssm states and the KV + // cache; generation continues token by token (O(N) per step instead + // of the former O(N^2) stateless recompute). if (prompt.codebook_rows != assets.config.fast.num_codebooks + 1 || static_cast(prompt.matrix.size()) != prompt.codebook_rows * prompt.steps) { throw std::runtime_error("Audio8 TTS AR prompt shape mismatch"); @@ -1138,8 +2343,13 @@ class Audio8TtsARRuntime::Impl { } return full; }; - auto pre_emb = build_falcon_embeddings(assets.config, weights, full_matrix.data(), cur_steps); - auto pre_out = falcon_forward_stateless(runtime_->backend(), runtime_->threads(), runtime_->graph_arena_bytes(), assets.config, weights, pre_emb, cur_steps); + FalconH1StepState fstate = init_falcon_step_state(assets.config); + SlowForwardOutput pre_out; + for (int64_t p = 0; p < cur_steps; ++p) { + auto p_emb = build_falcon_embedding_step(assets.config, weights, full_matrix.data(), cur_steps, p); + pre_out = falcon_forward_step(runtime_->falcon_step_backend(), runtime_->threads(), runtime_->graph_arena_bytes(), + assets.config, weights, p_emb, fstate, p, &profile); + } auto pre_logits_full = expand_compact(pre_out.logits); auto frame = sample_frame(pre_logits_full, pre_out.hidden, options, sample, false, profile); if (frame.front() == im_end_id()) { @@ -1161,8 +2371,10 @@ class Audio8TtsARRuntime::Impl { } bool ended_by_im_end = false; for (int64_t step = 1; step < max_new_tokens; ++step) { - auto emb = build_falcon_embeddings(assets.config, weights, full_matrix.data(), cur_steps); - auto out = falcon_forward_stateless(runtime_->backend(), runtime_->threads(), runtime_->graph_arena_bytes(), assets.config, weights, emb, cur_steps); + const int64_t pos = cur_steps - 1; + auto emb = build_falcon_embedding_step(assets.config, weights, full_matrix.data(), cur_steps, pos); + auto out = falcon_forward_step(runtime_->falcon_step_backend(), runtime_->threads(), runtime_->graph_arena_bytes(), + assets.config, weights, emb, fstate, pos, &profile); auto logits_full = expand_compact(out.logits); auto next_frame = sample_frame(logits_full, out.hidden, options, sample, true, profile); if (next_frame.front() == im_end_id()) { ended_by_im_end = true; break; } @@ -1607,7 +2819,7 @@ class Audio8TtsARRuntime::Impl { cache_keys.reserve(weights.fast_layers.size()); cache_values.reserve(weights.fast_layers.size()); const ggml_type cache_type = - runtime_->backend_type() == core::BackendType::Vulkan ? GGML_TYPE_F32 : GGML_TYPE_BF16; + runtime_->fast_backend_type() == core::BackendType::Vulkan ? GGML_TYPE_F32 : GGML_TYPE_BF16; for (size_t layer = 0; layer < weights.fast_layers.size(); ++layer) { cache_keys.push_back(core::wrap_tensor( ggml_new_tensor_4d( @@ -1630,7 +2842,7 @@ class Audio8TtsARRuntime::Impl { core::TensorShape::from_dims({1, config.num_codebooks, config.n_local_heads, config.head_dim}), cache_type)); } - state_buffer_ = ggml_backend_alloc_ctx_tensors(state_ctx_.get(), runtime_->backend()); + state_buffer_ = ggml_backend_alloc_ctx_tensors(state_ctx_.get(), runtime_->fast_backend()); if (state_buffer_ == nullptr) { throw std::runtime_error("failed to allocate Audio8 TTS fast AR state tensors"); } @@ -1643,7 +2855,7 @@ class Audio8TtsARRuntime::Impl { ggml_backend_tensor_set(cache.tensor, zeros.data(), 0, zeros.size()); } - core::ModuleBuildContext ctx{graph_ctx_.get(), "audio8_tts.ar.fast", runtime_->backend_type()}; + core::ModuleBuildContext ctx{graph_ctx_.get(), "audio8_tts.ar.fast", runtime_->fast_backend_type()}; auto input = core::make_tensor(ctx, GGML_TYPE_F32, core::TensorShape::from_dims({1, 1, config.dim})); input = core::wrap_tensor(ggml_cpy(ctx.ggml, input_, input.tensor), input.shape, input.type); auto position_value = core::wrap_tensor(position_, core::TensorShape::from_dims({1}), GGML_TYPE_I32); @@ -1664,7 +2876,7 @@ class Audio8TtsARRuntime::Impl { input, position_value, decoder_weights, - make_fast_decoder_config(config, runtime_->backend_type()), + make_fast_decoder_config(config, runtime_->fast_backend_type()), config.num_codebooks, mask_value, position_value, @@ -1676,7 +2888,7 @@ class Audio8TtsARRuntime::Impl { ggml_build_forward_expand(graph_, logits_); constants.finish_graph(); constants.ensure_uploaded(); - gallocr_ = ggml_gallocr_new(ggml_backend_get_default_buffer_type(runtime_->backend())); + gallocr_ = ggml_gallocr_new(ggml_backend_get_default_buffer_type(runtime_->fast_backend())); if (gallocr_ == nullptr || !ggml_gallocr_reserve(gallocr_, graph_) || !ggml_gallocr_alloc_graph(gallocr_, graph_)) { @@ -1686,7 +2898,7 @@ class Audio8TtsARRuntime::Impl { } ~FastGraph() { - core::release_backend_graph_resources(runtime_->backend(), graph_); + core::release_backend_graph_resources(runtime_->fast_backend(), graph_); if (gallocr_ != nullptr) { ggml_gallocr_free(gallocr_); } @@ -1721,10 +2933,10 @@ class Audio8TtsARRuntime::Impl { timing_start = Clock::now(); ggml_backend_tensor_set(input_, input.data(), 0, input.size() * sizeof(float)); profile.fast_input_upload_ms += engine::debug::elapsed_ms(timing_start, Clock::now()); - core::set_backend_threads(runtime_->backend(), runtime_->threads()); + core::set_backend_threads(runtime_->fast_backend(), runtime_->threads()); timing_start = Clock::now(); - const ggml_status status = core::compute_backend_graph(runtime_->backend(), graph_, nullptr, "audio8_tts.ar.fast"); - ggml_backend_synchronize(runtime_->backend()); + const ggml_status status = core::compute_backend_graph(runtime_->fast_backend(), graph_, nullptr, "audio8_tts.ar.fast"); + ggml_backend_synchronize(runtime_->fast_backend()); profile.fast_graph_ms += engine::debug::elapsed_ms(timing_start, Clock::now()); if (status != GGML_STATUS_SUCCESS) { throw std::runtime_error("Audio8 TTS fast AR graph compute failed"); @@ -1881,6 +3093,13 @@ class Audio8TtsARRuntime::Impl { engine::debug::timing_log_scalar("audio8_tts.ar.profile.sample_main_ms", profile.sample_main_ms); engine::debug::timing_log_scalar("audio8_tts.ar.profile.sample_high_ms", profile.sample_high_ms); engine::debug::timing_log_scalar("audio8_tts.ar.profile.sample_fast_ms", profile.sample_fast_ms); + engine::debug::timing_log_scalar("audio8_tts.ar.profile.falcon_step_init_ms", profile.falcon_step_init_ms); + engine::debug::timing_log_scalar("audio8_tts.ar.profile.falcon_step_build_ms", profile.falcon_step_build_ms); + engine::debug::timing_log_scalar("audio8_tts.ar.profile.falcon_step_gallocr_ms", profile.falcon_step_gallocr_ms); + engine::debug::timing_log_scalar("audio8_tts.ar.profile.falcon_step_upload_ms", profile.falcon_step_upload_ms); + engine::debug::timing_log_scalar("audio8_tts.ar.profile.falcon_step_compute_ms", profile.falcon_step_compute_ms); + engine::debug::timing_log_scalar("audio8_tts.ar.profile.falcon_step_download_ms", profile.falcon_step_download_ms); + engine::debug::timing_log_scalar("audio8_tts.ar.profile.falcon_step_runs", profile.falcon_step_runs); engine::debug::trace_log_scalar("audio8_tts.ar.profile.prefill_runs", profile.prefill_runs); engine::debug::trace_log_scalar("audio8_tts.ar.profile.step_runs", profile.step_runs); engine::debug::trace_log_scalar("audio8_tts.ar.profile.fast_runs", profile.fast_runs); diff --git a/src/community_models/audio8_tts/codec.cpp b/src/community_models/audio8_tts/codec.cpp index fd43ed19a..03c68d396 100644 --- a/src/community_models/audio8_tts/codec.cpp +++ b/src/community_models/audio8_tts/codec.cpp @@ -1,9 +1,12 @@ +#include + #include "engine/community_models/audio8_tts/codec.h" #include "engine/framework/audio/conversion.h" #include "engine/framework/audio/resampling.h" #include "engine/framework/core/backend.h" #include "engine/framework/core/backend_weight_store.h" +#include "engine/framework/debug/profiler.h" #include "engine/framework/debug/trace.h" #include "engine/framework/core/execution_context.h" #include "engine/framework/modules/activation_modules.h" @@ -463,6 +466,131 @@ core::TensorValue build_window_transformer( return modules::TransposeModule({{0, 2, 1, 3}, 3}).build(ctx, x); } +// ---- Channel-fast decoder region (Metal only) -------------------------------------- +// The decoder's residual units are back-to-back stride-1 convolutions separated only by +// snake activations and residual adds -- all elementwise, hence layout-agnostic. Running +// a whole block (snake -> convT upsample -> 3 residual units) in channel-fast +// [channels, frames] layout avoids the two transposes per conv that the module-level +// fast path pays at every boundary, and the convT's own internal transpose. Numerics +// are unchanged: identical GEMMs, identical elementwise ops, same accumulation order. + +bool codec_channel_fast_decoder_enabled() { + static const bool enabled = [] { + const char * value = std::getenv("AUDIO8_TTS_CODEC_CHANNEL_FAST"); + return value == nullptr || value[0] != '0'; + }(); + return enabled; +} + +ggml_tensor * channel_fast_in(core::ModuleBuildContext & ctx, const core::TensorValue & x) { + const auto contiguous = core::ensure_backend_addressable_layout(ctx, x); + return ggml_cont(ctx.ggml, ggml_transpose(ctx.ggml, contiguous.tensor)); +} + +core::TensorValue channel_fast_out(core::ModuleBuildContext & ctx, ggml_tensor * x_cf, int64_t channels) { + return core::wrap_tensor( + ggml_cont(ctx.ggml, ggml_transpose(ctx.ggml, x_cf)), + core::TensorShape::from_dims({1, channels, x_cf->ne[1]}), + GGML_TYPE_F32); +} + +// snake(x)[c, t] = x + sin^2(alpha_c * x) / alpha_c, alpha broadcast along frames. +ggml_tensor * snake1d_channel_fast( + core::ModuleBuildContext & ctx, + ggml_tensor * x_cf, + const core::TensorValue & alpha) { + auto * alpha_cf = ggml_reshape_2d(ctx.ggml, alpha.tensor, alpha.tensor->ne[0], 1); + // Fused snake op: one elementwise pass instead of a 5-kernel chain (mul -> sin -> + // mul -> div -> add) over the largest decoder tensors. Per-element op order matches + // the chain; residual differences are metal sin ulp-level only (max int16 delta 12 + // over a full utterance vs the chain). Set AUDIO8_TTS_CODEC_SNAKE_FUSED=0 to fall + // back to the explicit chain. + static const bool fused_disabled = [] { + const char * e = std::getenv("AUDIO8_TTS_CODEC_SNAKE_FUSED"); + return e && e[0] == '0' && e[1] == '\0'; + }(); + if (!fused_disabled) { + ggml_tensor * alpha_f32 = alpha_cf; + if (alpha_f32->type != GGML_TYPE_F32) { + alpha_f32 = ggml_cast(ctx.ggml, alpha_f32, GGML_TYPE_F32); + } + return ggml_snake_1d(ctx.ggml, x_cf, alpha_f32); + } + auto * ax = ggml_mul(ctx.ggml, x_cf, alpha_cf); + auto * s = ggml_sin(ctx.ggml, ax); + auto * s2 = ggml_mul(ctx.ggml, s, s); + return ggml_add(ctx.ggml, x_cf, ggml_div(ctx.ggml, s2, alpha_cf)); +} + +// Causal left pad built from scaled-to-zero columns of the input itself: activations are +// finite so x * 0 is a bitwise zero, and snake(+-0) = +0 keeps the pad region exact. +ggml_tensor * channel_fast_causal_pad( + core::ModuleBuildContext & ctx, + ggml_tensor * x_cf, + int64_t channels, + int64_t left_pad) { + if (left_pad <= 0) { + return x_cf; + } + auto * head = ggml_view_2d(ctx.ggml, x_cf, channels, left_pad, x_cf->nb[1], 0); + auto * zeros = ggml_scale(ctx.ggml, head, 0.0f); + return ggml_concat(ctx.ggml, zeros, x_cf, 1); +} + +ggml_tensor * causal_conv1d_channel_fast( + core::ModuleBuildContext & ctx, + ggml_tensor * x_cf, + const modules::Conv1dWeights & weights, + int64_t channels, + int64_t kernel, + int dilation) { + const int64_t left_pad = (kernel - 1) * dilation; // stride == 1 + auto * padded = channel_fast_causal_pad(ctx, x_cf, channels, left_pad); + return modules::conv1d_pertap_channel_fast( + ctx, + weights, + padded, + modules::Conv1dConfig{channels, channels, kernel, 1, 0, dilation, true}); +} + +ggml_tensor * residual_unit_channel_fast( + core::ModuleBuildContext & ctx, + ggml_tensor * x_cf, + const ResidualUnitWeights & weights, + int64_t channels, + int dilation) { + const int64_t frames = x_cf->ne[1]; + auto * y = snake1d_channel_fast(ctx, x_cf, weights.snake1.alpha); + y = causal_conv1d_channel_fast(ctx, y, weights.conv1, channels, 7, dilation); + y = snake1d_channel_fast(ctx, y, weights.snake2.alpha); + y = causal_conv1d_channel_fast(ctx, y, weights.conv2, channels, 1, 1); + ggml_tensor * residual = x_cf; + if (y->ne[1] != frames) { + residual = ggml_view_2d(ctx.ggml, x_cf, channels, y->ne[1], x_cf->nb[1], 0); + } + return ggml_add(ctx.ggml, residual, y); +} + +// Transposed-conv upsample on channel-fast input; the col2im output is time-fast, so +// the causal trim and the transpose back happen here. +ggml_tensor * causal_conv_transpose1d_channel_fast( + core::ModuleBuildContext & ctx, + ggml_tensor * x_cf, + const modules::ConvTranspose1dWeights & weights, + int64_t in_channels, + int64_t out_channels, + int64_t kernel, + int stride) { + auto * out_tf = modules::conv_transpose1d_col2im_channel_fast( + ctx, + weights, + x_cf, + modules::ConvTranspose1dConfig{in_channels, out_channels, kernel, stride, 0, 1, true}); + const int64_t pad = kernel - stride; // padding_left == 0; drop the trailing frames + auto * trimmed = ggml_view_2d(ctx.ggml, out_tf, out_tf->ne[0] - pad, out_channels, out_tf->nb[1], 0); + return ggml_cont(ctx.ggml, ggml_transpose(ctx.ggml, trimmed)); +} + core::TensorValue build_residual_unit( core::ModuleBuildContext & ctx, const core::TensorValue & input, @@ -531,14 +659,29 @@ core::TensorValue build_decoder( auto x = causal_conv1d(ctx, input, weights.decoder_first, kCodecDim, 1536, 7, 1, 1, true); int64_t channels = 1536; const int strides[] = {8, 8, 4, 2}; - for (size_t index = 0; index < weights.decoder_blocks.size(); ++index) { - const auto & block = weights.decoder_blocks[index]; - x = modules::Snake1dModule({channels}).build(ctx, x, block.snake); - x = causal_conv_transpose1d(ctx, x, block.conv, channels, channels / 2, 2 * strides[index], strides[index], true); - channels /= 2; - x = build_residual_unit(ctx, x, block.residual1, channels, 1); - x = build_residual_unit(ctx, x, block.residual3, channels, 3); - x = build_residual_unit(ctx, x, block.residual9, channels, 9); + if (ctx.backend_type == core::BackendType::Metal && codec_channel_fast_decoder_enabled()) { + ggml_tensor * x_cf = channel_fast_in(ctx, x); + for (size_t index = 0; index < weights.decoder_blocks.size(); ++index) { + const auto & block = weights.decoder_blocks[index]; + x_cf = snake1d_channel_fast(ctx, x_cf, block.snake.alpha); + x_cf = causal_conv_transpose1d_channel_fast( + ctx, x_cf, block.conv, channels, channels / 2, 2 * strides[index], strides[index]); + channels /= 2; + x_cf = residual_unit_channel_fast(ctx, x_cf, block.residual1, channels, 1); + x_cf = residual_unit_channel_fast(ctx, x_cf, block.residual3, channels, 3); + x_cf = residual_unit_channel_fast(ctx, x_cf, block.residual9, channels, 9); + } + x = channel_fast_out(ctx, x_cf, channels); + } else { + for (size_t index = 0; index < weights.decoder_blocks.size(); ++index) { + const auto & block = weights.decoder_blocks[index]; + x = modules::Snake1dModule({channels}).build(ctx, x, block.snake); + x = causal_conv_transpose1d(ctx, x, block.conv, channels, channels / 2, 2 * strides[index], strides[index], true); + channels /= 2; + x = build_residual_unit(ctx, x, block.residual1, channels, 1); + x = build_residual_unit(ctx, x, block.residual3, channels, 3); + x = build_residual_unit(ctx, x, block.residual9, channels, 9); + } } x = modules::Snake1dModule({channels}).build(ctx, x, weights.decoder_final_snake); x = causal_conv1d(ctx, x, weights.decoder_final, channels, 1, 7, 1, 1, true); @@ -884,6 +1027,7 @@ struct DecodeGraph { throw std::runtime_error("failed to initialize Audio8 TTS codec decode graph context"); } core::ModuleBuildContext ctx{ctx_.get(), "audio8_tts.codec.decode", backend_type_}; + const auto build_start = std::chrono::steady_clock::now(); constants_.begin_graph(); for (int64_t codebook = 0; codebook < assets_->config.codec.total_codebooks; ++codebook) { auto ids = core::make_tensor(ctx, GGML_TYPE_I32, core::TensorShape::from_dims({1, frame_capacity_})); @@ -902,6 +1046,10 @@ struct DecodeGraph { if (gallocr_ == nullptr || !ggml_gallocr_alloc_graph(gallocr_.get(), graph_)) { throw std::runtime_error("failed to allocate Audio8 TTS codec decode graph"); } + engine::debug::timing_log_scalar( + "audio8_tts.codec.graph_build_ms", + engine::debug::elapsed_ms(build_start, std::chrono::steady_clock::now())); + engine::debug::trace_log_scalar("audio8_tts.codec.graph_nodes", static_cast(ggml_graph_n_nodes(graph_))); } ~DecodeGraph() { @@ -941,8 +1089,12 @@ struct DecodeGraph { core::write_tensor_i32(code_inputs_[static_cast(codebook)], padded); } core::set_backend_threads(backend_, threads_); + const auto compute_start = std::chrono::steady_clock::now(); const ggml_status status = engine::core::compute_backend_graph(backend_, graph_); ggml_backend_synchronize(backend_); + engine::debug::timing_log_scalar( + "audio8_tts.codec.graph_compute_ms", + engine::debug::elapsed_ms(compute_start, std::chrono::steady_clock::now())); if (status != GGML_STATUS_SUCCESS) { throw std::runtime_error("Audio8 TTS codec decode graph compute failed"); } diff --git a/src/community_models/audio8_tts/falcon_kv_cache.h b/src/community_models/audio8_tts/falcon_kv_cache.h new file mode 100644 index 000000000..a03090cd2 --- /dev/null +++ b/src/community_models/audio8_tts/falcon_kv_cache.h @@ -0,0 +1,39 @@ +#pragma once + +#include +#include +#include + +namespace engine::models::audio8_tts { + +// Appends one token's K (or V) vectors to a host-side KV cache stored in +// ggml's col-major [head_dim, seq, n_kv] layout: element (d, t, h) lives at +// d + head_dim*(t + seq*h), i.e. the per-head stride is the CURRENT sequence +// length. Because that stride grows with every appended token, the cached +// tokens must be re-laid into the new stride before `fresh` (the new token's +// n_kv*head_dim values, head-major) is written at t = seq. Appending without +// the re-layout makes the new token overwrite the previous head blocks and +// silently corrupts attention context from the second token on. +inline void append_falcon_kv_token( + std::vector & cache, + int64_t seq, + int64_t n_kv, + int64_t head_dim, + const float * fresh) { + const int64_t new_seq_len = seq + 1; + std::vector old; + old.swap(cache); + cache.assign(static_cast(new_seq_len * n_kv * head_dim), 0.0F); + for (int64_t h = 0; h < n_kv; ++h) { + for (int64_t t = 0; t < seq; ++t) { + std::copy_n(old.data() + head_dim * (t + seq * h), + static_cast(head_dim), + cache.data() + head_dim * (t + new_seq_len * h)); + } + std::copy_n(fresh + head_dim * h, + static_cast(head_dim), + cache.data() + head_dim * (seq + new_seq_len * h)); + } +} + +} // namespace engine::models::audio8_tts diff --git a/src/framework/modules/conv_modules.cpp b/src/framework/modules/conv_modules.cpp index e6f562b92..d3f5b3345 100644 --- a/src/framework/modules/conv_modules.cpp +++ b/src/framework/modules/conv_modules.cpp @@ -193,6 +193,99 @@ core::TensorValue depthwise_conv2d_weight( int64_t conv1d_output_frames(const Conv1dConfig & config, int64_t input_frames) { return (input_frames + 2 * config.padding - config.dilation * (config.kernel_size - 1) - 1) / config.stride + 1; } +bool is_conv1d_pertap_fast_path_eligible( + const core::ModuleBuildContext & ctx, + const Conv1dConfig & config, + const core::TensorValue & input) noexcept { + return ctx.backend_type == core::BackendType::Metal && + config.padding == 0 && + config.stride == 1 && + input.shape.dims[0] == 1 && + input.type == GGML_TYPE_F32 && + input.tensor->ne[0] == input.shape.dims[2] && + input.tensor->ne[1] == config.in_channels && + ggml_is_contiguous(input.tensor); +} + +// Per-tap GEMM accumulation on a channel-fast [in_channels, frames] F32 input: one +// contiguous GEMM per kernel tap over shifted column views, accumulated into +// [out_channels, output_frames]. No layout conversion here -- callers at region edges +// transpose; chained callers keep everything channel-fast. +ggml_tensor * conv1d_pertap_gemm_channel_fast( + core::ModuleBuildContext & ctx, + ggml_tensor * input_cf, + ggml_tensor * weight_f32, + int64_t in_channels, + int64_t out_channels, + int64_t kernel_size, + int64_t dilation, + int64_t output_frames) { + // weight logical [OC, IC, K] -> ggml ne [K, IC, OC]; regroup rows so each tap slice + // [IC, OC] is a contiguous view: row index = channel + in_channels * tap. + auto * weight_taps = ggml_reshape_2d( + ctx.ggml, + ggml_cont(ctx.ggml, ggml_permute(ctx.ggml, weight_f32, 1, 0, 2, 3)), + in_channels * kernel_size, + out_channels); + // accumulate-in-place GEMM only where the tensor-core mm kernel applies + // (mirrors the Metal supports gate); otherwise keep the mul_mat + add chain + const bool use_acc = in_channels >= 64 && output_frames > 8; + ggml_tensor * acc = nullptr; + for (int64_t tap = 0; tap < kernel_size; ++tap) { + // columns[c, j] = input[tap * dilation + j, c]: contiguous column view of input_cf. + auto * columns = ggml_view_2d( + ctx.ggml, + input_cf, + in_channels, + output_frames, + input_cf->nb[1], + static_cast(tap * dilation) * in_channels * sizeof(float)); + auto * tap_weights = ggml_view_2d( + ctx.ggml, + weight_taps, + in_channels, + out_channels, + weight_taps->nb[1], + static_cast(tap) * in_channels * sizeof(float)); + if (acc == nullptr) { + acc = ggml_mul_mat(ctx.ggml, tap_weights, columns); + } else if (use_acc) { + acc = ggml_mul_mat_acc(ctx.ggml, tap_weights, columns, acc); + } else { + acc = ggml_add(ctx.ggml, acc, ggml_mul_mat(ctx.ggml, tap_weights, columns)); + } + } + return acc; +} + +// Metal fast path for stride-1 conv1d on the time-fast [frames, channels] layout used by +// the audio codecs. ggml_conv_1d materializes an im2col matrix whose kernel taps are +// strided gathers in this layout (~200 ms per conv at [569k, 96] on M4); instead, transpose +// the input to channel-fast once, run one contiguous GEMM per kernel tap, and transpose +// the accumulator back (~3-6x faster). +core::TensorValue build_conv1d_pertap_fast_path( + core::ModuleBuildContext & ctx, + const Conv1dConfig & config, + const core::TensorValue & input, + const core::TensorValue & weight_f32, + const core::TensorShape & output_shape) { + // channel-fast copy of the input: [IC, frames]; kernel taps become contiguous columns. + auto * input_cf = ggml_cont(ctx.ggml, ggml_transpose(ctx.ggml, input.tensor)); + auto * acc = conv1d_pertap_gemm_channel_fast( + ctx, + input_cf, + weight_f32.tensor, + config.in_channels, + config.out_channels, + config.kernel_size, + config.dilation, + output_shape.dims[2]); + // mul_mat yields [OC, frames]; restore the canonical [frames, OC] orientation. + return core::wrap_tensor( + ggml_cont(ctx.ggml, ggml_transpose(ctx.ggml, acc)), + output_shape, + GGML_TYPE_F32); +} int64_t conv2d_output_dim(int64_t input, int kernel, int stride, int padding, int dilation) { return (input + 2 * padding - dilation * (kernel - 1) - 1) / stride + 1; @@ -324,6 +417,92 @@ bool is_conv_transpose1d_col2im_fast_path_eligible( config.dilation == 1; } +ggml_tensor * conv1d_pertap_channel_fast( + core::ModuleBuildContext & ctx, + const Conv1dWeights & weights, + ggml_tensor * input_cf, + const Conv1dConfig & config) { + if (ctx.ggml == nullptr || input_cf == nullptr) { + throw std::runtime_error("conv1d_pertap_channel_fast requires a ggml context and an input tensor"); + } + if (config.padding != 0 || config.stride != 1) { + throw std::runtime_error("conv1d_pertap_channel_fast requires padding=0 and stride=1"); + } + if (input_cf->type != GGML_TYPE_F32 || input_cf->ne[0] != config.in_channels || + !ggml_is_contiguous(input_cf)) { + throw std::runtime_error("conv1d_pertap_channel_fast requires contiguous F32 [in_channels, frames] input"); + } + auto weight = regular_conv_weight(ctx, weights.weight, "conv1d_pertap_channel_fast"); + if (weight.type != GGML_TYPE_F32) { + weight = core::wrap_tensor( + ggml_cast(ctx.ggml, weight.tensor, GGML_TYPE_F32), weight.shape, GGML_TYPE_F32); + } + const int64_t output_frames = input_cf->ne[1] - config.dilation * (config.kernel_size - 1); + ggml_tensor * acc = conv1d_pertap_gemm_channel_fast( + ctx, + input_cf, + weight.tensor, + config.in_channels, + config.out_channels, + config.kernel_size, + config.dilation, + output_frames); + if (config.use_bias) { + if (!weights.bias.has_value()) { + throw std::runtime_error("conv1d_pertap_channel_fast requires bias when use_bias is true"); + } + const auto bias = ensure_f32(ctx, *weights.bias); + core::validate_shape(bias, core::TensorShape::from_dims({config.out_channels}), "bias"); + acc = ggml_add(ctx.ggml, acc, ggml_reshape_2d(ctx.ggml, bias.tensor, config.out_channels, 1)); + } + return acc; +} + +ggml_tensor * conv_transpose1d_col2im_channel_fast( + core::ModuleBuildContext & ctx, + const ConvTranspose1dWeights & weights, + ggml_tensor * input_cf, + const ConvTranspose1dConfig & config) { + if (!is_conv_transpose1d_col2im_fast_path_eligible(ctx, config)) { + throw std::runtime_error("conv_transpose1d_col2im_channel_fast called with an ineligible config"); + } + if (input_cf == nullptr || input_cf->type != GGML_TYPE_F32 || + input_cf->ne[0] != config.in_channels || !ggml_is_contiguous(input_cf)) { + throw std::runtime_error( + "conv_transpose1d_col2im_channel_fast requires contiguous F32 [in_channels, frames] input"); + } + auto weight_contiguous = tensor_layout::ensure_contiguous_layout_if_needed(ctx, weights.weight); + if (weight_contiguous.type != GGML_TYPE_F32) { + weight_contiguous = core::wrap_tensor( + ggml_cast(ctx.ggml, weight_contiguous.tensor, GGML_TYPE_F32), + weight_contiguous.shape, + GGML_TYPE_F32); + } + auto * weight_perm = ggml_reshape_2d( + ctx.ggml, + ggml_cont(ctx.ggml, ggml_permute(ctx.ggml, weight_contiguous.tensor, 1, 2, 0, 3)), + config.in_channels, + config.kernel_size * config.out_channels); + auto * columns = ggml_mul_mat(ctx.ggml, weight_perm, input_cf); + auto * output = ggml_col2im_1d( + ctx.ggml, + columns, + config.stride, + static_cast(config.out_channels), + config.padding); + if (config.use_bias) { + if (!weights.bias.has_value()) { + throw std::runtime_error("conv_transpose1d_col2im_channel_fast requires bias when use_bias is true"); + } + core::validate_shape(*weights.bias, core::TensorShape::from_dims({config.out_channels}), "bias"); + output = ggml_add( + ctx.ggml, + output, + ggml_reshape_2d(ctx.ggml, weights.bias->tensor, 1, config.out_channels)); + } + return output; +} + Conv1dModule::Conv1dModule(Conv1dConfig config) : config_(config) { if (config_.in_channels <= 0 || config_.out_channels <= 0 || config_.kernel_size <= 0) { throw std::runtime_error("Conv1dConfig dimensions must be positive"); @@ -363,7 +542,10 @@ core::TensorValue Conv1dModule::build( const auto input_contiguous = ensure_f32(ctx, tensor_layout::ensure_contiguous_layout_if_needed(ctx, input)); const auto weight_contiguous = regular_conv_weight(ctx, weights.weight, "Conv1dModule"); core::TensorValue output; - if (input.shape.dims[0] == 1) { + if (is_conv1d_pertap_fast_path_eligible(ctx, config_, input) && + weight_contiguous.type == GGML_TYPE_F32) { + output = build_conv1d_pertap_fast_path(ctx, config_, input_contiguous, weight_contiguous, output_shape); + } else if (input.shape.dims[0] == 1) { output = core::wrap_tensor( ggml_conv_1d( ctx.ggml, diff --git a/tests/unittests/test_audio8_tts_falcon_kv_cache.cpp b/tests/unittests/test_audio8_tts_falcon_kv_cache.cpp new file mode 100644 index 000000000..40c3e0bbb --- /dev/null +++ b/tests/unittests/test_audio8_tts_falcon_kv_cache.cpp @@ -0,0 +1,66 @@ +// Regression test for the Falcon-H1 host KV cache append (audio8_tts 0.1B). +// +// The cache is stored in ggml's col-major [head_dim, seq, n_kv] layout where +// the per-head stride is the CURRENT sequence length. Appending a token +// without re-laying the existing entries into the new stride makes the new +// token overwrite the previous head blocks: with n_kv=2, appending token 1 +// wrote head 0 at floats [64,128) — exactly where token 0's head 1 lived — +// so from the second token on, attention read corrupted keys/values for every +// head past the first (argmax diverged from the HF reference at prompt step 3 +// and the recurrent state blew up on long sequences). +// +// This test feeds recognizable per-(token, head, dim) values through +// append_falcon_kv_token and checks the full cache contents after every +// append; the historical implementation fails from the second append on. + +#include "falcon_kv_cache.h" + +#include "test_assert.h" + +#include +#include +#include + +namespace { + +float marker(int64_t token, int64_t head, int64_t dim) { + return static_cast(token * 100000 + head * 1000 + dim); +} + +} // namespace + +int main() { + using engine::test::require; + using engine::test::require_eq; + using engine::models::audio8_tts::append_falcon_kv_token; + + constexpr int64_t n_kv = 2; + constexpr int64_t head_dim = 64; + constexpr int64_t n_tokens = 8; + + std::vector cache; + for (int64_t t = 0; t < n_tokens; ++t) { + std::vector fresh(static_cast(n_kv * head_dim)); + for (int64_t h = 0; h < n_kv; ++h) { + for (int64_t d = 0; d < head_dim; ++d) { + fresh[static_cast(d + head_dim * h)] = marker(t, h, d); + } + } + append_falcon_kv_token(cache, t, n_kv, head_dim, fresh.data()); + + const int64_t seq = t + 1; + require_eq(static_cast(cache.size()), seq * n_kv * head_dim, "cache size after append"); + for (int64_t h = 0; h < n_kv; ++h) { + for (int64_t tt = 0; tt < seq; ++tt) { + for (int64_t d = 0; d < head_dim; ++d) { + const float actual = cache[static_cast(d + head_dim * (tt + seq * h))]; + require(actual == marker(tt, h, d), + "cache entry corrupted after appending token " + std::to_string(t)); + } + } + } + } + + std::cout << "audio8_tts_falcon_kv_cache_test passed\n"; + return 0; +}