diff --git a/modules/dasLLAMA/ARCHITECTURE.md b/modules/dasLLAMA/ARCHITECTURE.md index 91b20c7562..035f2651e0 100644 --- a/modules/dasLLAMA/ARCHITECTURE.md +++ b/modules/dasLLAMA/ARCHITECTURE.md @@ -65,11 +65,13 @@ Each companion's sections are anchored by topic; a citation spells `.md#`dasllama_metal_kernels`, `dasllama_vulkan_classes` | kernel source, the threadgroup reductions every Metal body folds through (`tg_sum_all` / `tg_max_all`, `MetalTgReduceBase` over its own `partial[]`), the workgroup folds every Vulkan body whose fold is one value derives (`WgReduceBase`'s `wg_sum` / `wg_max` / `wg_rms_inv`, one value over every lane, returned to every lane) - two Vulkan bodies fold by hand by design, because their folds are split: `DnStepFused`'s q and k halves (`subgroupAdd` into per-subgroup arrays, a lane-0 loop over them, one set of results every lane reads) and `TtsPkAttn`'s `po` quarters (four partial value sums a column, added by the column's lane), the kernel-side quant-decode helpers that read a codebook table (the table-free arithmetic both homes splice is `dasllama_gpu_math`'s) and the per-word codebook accessors (`[grid_words]` bakes each from `dasllama_kqformat`'s one literal at compile time - the bytes live there, the home carries the baked copy), the derived-access/PSO census; on Vulkan the one device buffer kernel data fills (`kq_grid_dev`, the grid codebooks) and the host-side ensure/set/enc pick ladders and grid rules over its own class stamps (`gemv_*`, `q8_gemv_gu_n_*`, `q8_batch_cls_*`, `kq_batch_cls_*`, `fa_stamp_*`, `f16_gemm_*`, `da_slab_*`, `da_attn_stamp_*` with the block codecs' `kv_codec_cls_ensure`, `da_attn_b_stamp_*`, `rope_kv_b_stamp_*`, and the tile trios `cm2_cls_*` (the engine's, over every tile tail) and `khr_cls_*` (the kernel cells' KHR arm); `kq_tile_stamp` stamps the tile trios, `kq_batch_cls_*` and the GEMV ladders (its `gemv` family) over `KqFmt`, `stamp_ladder` the `fa_stamp_*` and `da_attn_stamp_*` ladders over their stamp keys); the tower row classes' mode selectors and family constants - `TowerBiasActT`'s act selector (`BIAS_ACT_NONE` / `_GELU_TANH` / `_GELU_ERF` / `_SILU` / `_RELU`, relu for canary's subsample stack) and `Q3A_TOK_PER_CHUNK`, `G4A_ATTN_PAST`, `G4A_ATTN_CAP`, `WDEC_HS`, which the drivers check against the family before they serve | device state other than `kq_grid_dev`, engine types | +| the kernel home
`dasllama_metal_kernels`, `dasllama_vulkan_classes` | kernel source, the threadgroup reductions every Metal body folds through (`tg_sum_all` / `tg_max_all`, `MetalTgReduceBase` over its own `partial[]`), the workgroup folds every Vulkan body whose fold is one value derives (`WgReduceBase`'s `wg_sum` / `wg_max` / `wg_rms_inv`, one value over every lane, returned to every lane - the value form of the shared `GkWgReduce` in `dasllama_gpu_kernels_common`, which a kernel both homes stamp derives directly) - two Vulkan bodies fold by hand by design, because their folds are split: `DnStepFused`'s q and k halves (`subgroupAdd` into per-subgroup arrays, a lane-0 loop over them, one set of results every lane reads) and `TtsPkAttn`'s `po` quarters (four partial value sums a column, added by the column's lane), the kernel-side quant-decode helpers that read a codebook table (the table-free arithmetic both homes splice is `dasllama_gpu_math`'s) and the per-word codebook accessors (`[grid_words]` bakes each from `dasllama_kqformat`'s one literal at compile time - the bytes live there, the home carries the baked copy), the derived-access/PSO census; on Vulkan the one device buffer kernel data fills (`kq_grid_dev`, the grid codebooks) and the host-side ensure/set/enc pick ladders and grid rules over its own class stamps (`gemv_*`, `q8_gemv_gu_n_*`, `q8_batch_cls_*`, `kq_batch_cls_*`, `fa_stamp_*`, `f16_gemm_*`, `da_slab_*`, `da_attn_stamp_*` with the block codecs' `kv_codec_cls_ensure`, `da_attn_b_stamp_*`, `rope_kv_b_stamp_*`, and the tile trios `cm2_cls_*` (the engine's, over every tile tail) and `khr_cls_*` (the kernel cells' KHR arm); `kq_tile_stamp` stamps the tile trios, `kq_batch_cls_*` and the GEMV ladders (its `gemv` family) over `KqFmt`, `stamp_ladder` the `fa_stamp_*` and `da_attn_stamp_*` ladders over their stamp keys); the tower row classes' mode selectors and family constants - `TowerBiasActT`'s act selector (`BIAS_ACT_NONE` / `_GELU_TANH` / `_GELU_ERF` / `_SILU` / `_RELU`, relu for canary's subsample stack) and `Q3A_TOK_PER_CHUNK`, `G4A_ATTN_PAST`, `G4A_ATTN_CAP`, `WDEC_HS`, which the drivers check against the family before they serve | device state other than `kq_grid_dev`, engine types | | `dasllama__common`
`dasllama_metal_common`, `dasllama_vulkan_common` | device state, buffer/command plumbing, hazard + capture rail, profiler, host-side quant-decode helpers (Metal's `iq4_lut`), the family's registrant of a tier seat that names a size the device state keeps (Metal's `dn_mirror_room`) | driver policy | | `dasllama__decode`
`dasllama_metal_decode`, `dasllama_vulkan_decode` | the resident token-step driver + decode-time arms; on Vulkan, for a hyper-connection MoE (`Config.hyper_conn`), the split token command and the hot pool's device slots, stacks and hit chains (`ARCHITECTURE_GPU_VULKAN_HC.md#hc-token-command`, `#hc-hot-pool`) | kernel bodies | | `dasllama__prefill`
`dasllama_metal_prefill`, `dasllama_vulkan_prefill` | the batched prefill driver + batch arms; on Vulkan, for a hyper-connection MoE, the window chain cut after each router (`ARCHITECTURE_GPU_VULKAN_HC.md#hc-window-chain`) | kernel bodies | | `dasllama__shapes`
`dasllama_metal_shapes` | PORTABLE servability gates - no GPU C++ require, so any box can bake | device calls | -| the tower driver
`dasllama_metal_tower`, `dasllama_vulkan_tower` (the vision ViT chains over the q8 image, qwen25v's over the halfword twin, and the audio chains over the q8 image - the whisper-class towers with their conv stem, the gemma4a Conformer with its chunk front and projector tail, the canary FastConformer with its front, the qwen3a front and the whisper mel seat; on a build without das_metal it fills the gemma4v and gemma3v hook slots, the blocks-only seats `register_qwen3v_gpu_blocks` and `register_qwen25v_gpu_blocks`, the Vulkan seats of those two families, and the audio seats `register_tower_blocks_gpu`, `register_tower_blocks_ln_post_gpu`, `register_tower_conv_gpu`, `register_qwen3a_front_gpu`, `register_whisper_mel_gpu`, `register_gemma4a_gpu`, `register_gemma4a_chunk_gpu`, `register_canary_gpu` and `register_canary_front_gpu`; parakeet's seat is Metal's) | one-shot embedder/encoder encodes (gemma4uv chain, the gemma4v ViT with its patch-conv stem and projector tail, gemma3v SigLIP with its patch conv and projector tail, and qwen3v block loops - qwen3v adds the vision NEOX rope, the fused-qkv weight-offset GEMMs, and the inline deepstack tap + tail merger chains - the whisper-class block loop with its post-norm or its projector tail behind it, the qwen25v window ViT, the gemma4a Conformer chain with its mel/conv front, the FastConformer chain canary and parakeet share - one block body over `CanaryLayerOffs`, parakeet's offsets mapped onto it with no GEMM biases and the tap-major depthwise stamp - the Pocket TTS codec and frame loop (`ARCHITECTURE_GPU_TOWER.md#tower-pocket-codec`, `ARCHITECTURE_GPU_TOWER.md#tower-pocket-frames`) - the whole StyleTTS2 synthesis (kitten and kokoro on both weight lanes: seven seats from PL-BERT to the inverse STFT, the front end on the f32-exact GEMM stamp, the source chain the CPU's operation for operation, `ARCHITECTURE_GPU_TOWER.md#tower-tts-chain`) - the conv frontends + the qwen3a padded-weight slab and GPU front/mel) - no session, no mirror - the Pocket frame loop's per-voice device K/V slot is the one state kept across calls; registers the gemma4uv, gemma4v (`register_gemma4v_gpu` with its `stem` flag set: the chain takes the stem's columns and runs the patch conv and the position adds itself), gemma3v, qwen3v, qwen25v, whisper-class blocks (`register_tower_blocks_gpu`, `register_tower_blocks_ln_post_gpu` with the post-norm behind them, `register_tower_blocks_tail_gpu` with the projector tail behind them), tower-conv, qwen3a-front, whisper-mel (`register_whisper_mel_gpu`, the spectrum of the qwen3a mel and of the chat towers' chunked mel), gemma4a, gemma4a-chunk, canary (`register_canary_gpu` for the blocks, `register_canary_front_gpu` for the front), parakeet (`register_parakeet_gpu` for the blocks, `register_parakeet_gpu_encode` for the whole encode with the subsample front) and StyleTTS2 (`register_styletts2_gpu`, the seven-seat record) and Pocket (`register_pocket_gpu`, the codec and frames seats, `ARCHITECTURE_GPU_TOWER.md#tower-pocket-codec` and `ARCHITECTURE_GPU_TOWER.md#tower-pocket-frames`) hooks; the whisper-class block GEMMs and its conv stem's second GEMM on a q8 tower borrow `pf_enc_q8_mm`, the projector tails borrow `enc_tw_pool2` and `enc_tw_swiglu_rows` (whisper-class), `enc_tw_pool2d` (gemma4v, gemma3v) and `enc_tw_affine_rows` (gemma4v), the canary front borrows `enc_cn_melnorm`; the FastConformer chain borrows `enc_fc_pack` / `enc_fc_softmax` / `enc_fc_unpack` around `enc_f32_mm`, `enc_cn_dw` / `enc_pk_dw`, the prefill's `pf_enc_fc_attn_dev` (every head's attention at once on the device pair, under its crown) and `enc_bias_relu_rows` with the parakeet front's `enc_pk_conv0` / `enc_pk_dw2d` / `enc_pk_xf`, and the `enc_st2_*` row, LSTM, attention and source kernels around `enc_st2_conv_mm` / `enc_st2_conv_exact_mm`, the `enc_pk_*` row, add-and-norm, attention-row and prologue/epilogue GEMV kernels for the Pocket codec and frame loop beside the decode rail's `enc_gemv` / `enc_rpst32_c` and the prefill's `enc_rope`; on Vulkan the whisper-class encoder-output handoff - the served post-norm chain's `xb` plane, kept for the encode in flight alone, which the ASR-decoder driver's cross-KV chain reads (`vulkan_tower_enc_out`) - is the second state kept across calls beside the Pocket slot | LLM or ASR decoder session state | +| the tower driver
`dasllama_metal_tower`, `dasllama_vulkan_tower` (the vision ViT chains over the q8 image, qwen25v's over the halfword twin, and the audio chains over the q8 image - the whisper-class towers with their conv stem, the gemma4a Conformer with its chunk front and projector tail, the canary FastConformer with its front, the qwen3a front and the whisper mel seat; on a build without das_metal it fills the gemma4v and gemma3v hook slots, the blocks-only seats `register_qwen3v_gpu_blocks` and `register_qwen25v_gpu_blocks`, the Vulkan seats of those two families, and the audio seats `register_tower_blocks_gpu`, `register_tower_blocks_ln_post_gpu`, `register_tower_conv_gpu`, `register_qwen3a_front_gpu`, `register_whisper_mel_gpu`, `register_gemma4a_gpu`, `register_gemma4a_chunk_gpu`, `register_canary_gpu` and `register_canary_front_gpu`; parakeet's seat is Metal's) | one-shot embedder/encoder encodes (gemma4uv chain, the gemma4v ViT with its patch-conv stem and projector tail, gemma3v SigLIP with its patch conv and projector tail, and qwen3v block loops - qwen3v adds the vision NEOX rope, the fused-qkv weight-offset GEMMs, and the inline deepstack tap + tail merger chains - the whisper-class block loop with its post-norm or its projector tail behind it, the qwen25v window ViT, the gemma4a Conformer chain with its mel/conv front, the FastConformer chain canary and parakeet share - one block body over `CanaryLayerOffs`, parakeet's offsets mapped onto it with no GEMM biases and the tap-major depthwise stamp - the Pocket TTS codec, frame loop and text prompt (`ARCHITECTURE_GPU_TOWER_POCKET.md#tower-pocket-codec`, `ARCHITECTURE_GPU_TOWER_POCKET.md#tower-pocket-frames`) - the whole StyleTTS2 synthesis (kitten and kokoro on both weight lanes: seven seats from PL-BERT to the inverse STFT, the front end on the f32-exact GEMM stamp, the source chain the CPU's operation for operation, `ARCHITECTURE_GPU_TOWER.md#tower-tts-chain`) - the conv frontends + the qwen3a padded-weight slab and GPU front/mel) - no session, no mirror - the Pocket frame loop's per-voice device K/V slot is the one state kept across calls; registers the gemma4uv, gemma4v (`register_gemma4v_gpu` with its `stem` flag set: the chain takes the stem's columns and runs the patch conv and the position adds itself), gemma3v, qwen3v, qwen25v, whisper-class blocks (`register_tower_blocks_gpu`, `register_tower_blocks_ln_post_gpu` with the post-norm behind them, `register_tower_blocks_tail_gpu` with the projector tail behind them), tower-conv, qwen3a-front, whisper-mel (`register_whisper_mel_gpu`, the spectrum of the qwen3a mel and of the chat towers' chunked mel), gemma4a, gemma4a-chunk, canary (`register_canary_gpu` for the blocks, `register_canary_front_gpu` for the front), parakeet (`register_parakeet_gpu` for the blocks, `register_parakeet_gpu_encode` for the whole encode with the subsample front) and StyleTTS2 (`register_styletts2_gpu`, the seven-seat record) and Pocket (`register_pocket_gpu`, the codec, frames and prompt seats, `ARCHITECTURE_GPU_TOWER_POCKET.md#tower-pocket-codec` and `ARCHITECTURE_GPU_TOWER_POCKET.md#tower-pocket-frames`) hooks; the whisper-class block GEMMs and its conv stem's second GEMM on a q8 tower borrow `pf_enc_q8_mm`, the projector tails borrow `enc_tw_pool2` and `enc_tw_swiglu_rows` (whisper-class), `enc_tw_pool2d` (gemma4v, gemma3v) and `enc_tw_affine_rows` (gemma4v), the canary front borrows `enc_cn_melnorm`; the FastConformer chain borrows `enc_fc_pack` / `enc_fc_softmax` / `enc_fc_unpack` around `enc_f32_mm`, `enc_cn_dw` / `enc_pk_dw`, the prefill's `pf_enc_fc_attn_dev` (every head's attention at once on the device pair, under its crown) and `enc_bias_relu_rows` with the parakeet front's `enc_pk_conv0` / `enc_pk_dw2d` / `enc_pk_xf`, and the `enc_st2_*` row, LSTM, attention and source kernels around `enc_st2_conv_mm` / `enc_st2_conv_exact_mm`, the `enc_pk_*` row, add-and-norm, attention-row and prologue/epilogue GEMV kernels for the Pocket codec and frame loop beside the decode rail's `enc_gemv` / `enc_rpst32_c` and the prefill's `enc_rope`; on Vulkan the whisper-class encoder-output handoff - the served post-norm chain's `xb` plane, kept for the encode in flight alone, which the ASR-decoder driver's cross-KV chain reads (`vulkan_tower_enc_out`) - is the second state kept across calls beside the Pocket slot | LLM or ASR decoder session state | | the ASR-decoder driver
`dasllama_metal_asr_dec`, `dasllama_vulkan_asr_dec` | the whisper decoder on the GPU: Metal's 34B weight blob or Vulkan's row-major q8 gather, the f16 resident cross/self K/V, window-granular cross-KV + decode-step serves; registers the whisper cross-KV and decode hooks (`register_whisper_cross_kv_gpu`, `register_whisper_decode_gpu`; family registries in `dasllama_whisper`); on a build without das_metal the Vulkan driver fills them (`ARCHITECTURE_GPU_TOWER_VULKAN.md#vk-asr-decoder`) | kernel bodies, LLM session state | -| the TTS driver
`dasllama_vulkan_tts` (Vulkan; the Metal tower driver's own row carries the TTS seats on an Apple build) | the StyleTTS2 synthesis seats and the Pocket codec and frames seats on the tower's knob - the f32 weight slabs the shared host writer (`dasllama_tts_slab`) lays out, the front end on the f32-exact tile GEMM, the seats registered through `register_styletts2_gpu` and `register_pocket_gpu` on a build without das_metal; the seats serve only while the tier's want arms the device (`DASLLAMA_GPU=1`, off by default), so a run with no flags and no environment overrides synthesizes on the CPU chain on a Vulkan build, where on an Apple build the Metal tower driver's TTS seats serve by default (`DASLLAMA_METAL_TOWER` on) | kernel bodies, the CPU chain's arithmetic (the CPU chain is the specification), LLM or ASR session state | +| the TTS driver
`dasllama_vulkan_tts` (Vulkan; the Metal tower driver's own row carries the TTS seats on an Apple build) | the StyleTTS2 synthesis seats and the Pocket codec, frames and prompt seats on the tower's knob - the f32 weight slabs the shared host writer (`dasllama_tts_slab`) lays out, the front end on the f32-exact tile GEMM, the seats registered through `register_styletts2_gpu` and `register_pocket_gpu` on a build without das_metal; the seats serve only while the tier's want arms the device (`DASLLAMA_GPU=1`, off by default), so a run with no flags and no environment overrides synthesizes on the CPU chain on a Vulkan build, where on an Apple build the Metal tower driver's TTS seats serve by default (`DASLLAMA_METAL_TOWER` on) | kernel bodies, the CPU chain's arithmetic (the CPU chain is the specification), LLM or ASR session state | | the assistant-drafter driver
`dasllama_metal_mtp_gemma` | the gemma-4 assistant drafter on Metal: the sidecar blob upload, the Q-only layer chain reading the TARGET mirror at the two capture layers with the decode's own attention kernels, the speculative round over the batch driver's same-slab verify; registers the `metal` round override and delegates head-less-drafter-less models to `metal_mtp_spec_round` | kernel bodies, mirror ownership | | the kernel-access lens
`dasllama_metal_lens` (Metal), `dasllama_vulkan_dispatch` (Vulkan - the `[vk_dispatch]` macro derives access per class) | the kernel-access macro and its dispatch-support macros (`compile_stamp`, `release_handles`) | anything else | @@ -273,7 +273,15 @@ key it has on that home. in the arguments is one set. A kernel every site reads from element zero carries none. - **The two emitters place the body differently.** The MSL emitter splices a called method flat into the kernel; the SPIR-V emitter keeps it a function the entry point calls. A kernel moved here therefore - reads, on Vulkan, as its old words plus one call. + reads, on Vulkan, as its old words plus one call. Because the MSL emitter splices, a template method of + more than one statement cannot sit in value position: a helper that returns a value - a workgroup + reduction, a partial sum - is written as a statement method that lands its result in a `var` reference + (`wg_max_into`, `wg_sum_into`, `value_sum_into` in `GkPkAttn`). +- **A home's lane primitive is reached by one name.** A template method resolves in the module that + stamps it, so a shared reduction calls `gk_subgroup_add` / `gk_subgroup_max`, and each kernel home + defines the pair over its own primitive (`subgroupAdd` on Vulkan, `simd_sum` on Metal). A free function + in the common module resolves in the common module and sees neither home, so a body that needs a + home's primitive keeps it in a method. - **What stays per home:** the GEMM and attention tiles, which each home builds on its own matrix primitives. A pair whose two forms differ in shape (the elementwise maps run one element a thread on Metal and four an invocation on Vulkan) folds onto one body when both homes profile the same on it, and diff --git a/modules/dasLLAMA/ARCHITECTURE_GPU_TOWER.md b/modules/dasLLAMA/ARCHITECTURE_GPU_TOWER.md index 95371f2fb7..4acee32cd1 100644 --- a/modules/dasLLAMA/ARCHITECTURE_GPU_TOWER.md +++ b/modules/dasLLAMA/ARCHITECTURE_GPU_TOWER.md @@ -4,9 +4,11 @@ Companion to `ARCHITECTURE_GPU.md`; a section is cited by its anchor. This docum carries what a tower driver is and the registrations its hooks come through, the three routes that serve tower attention on Metal, the one-command-buffer encode chain each family the Metal tower driver serves gets, and the StyleTTS2 -synthesis chain, and the Pocket TTS codec and frames seats. The GPU backend role table these -sections build on - the tower driver's role row included - stays in `ARCHITECTURE_GPU.md#gpu-backends`. +synthesis chain. The GPU backend role table these sections build on - the tower driver's role row +included - stays in `ARCHITECTURE_GPU.md#gpu-backends`. +- `ARCHITECTURE_GPU_TOWER_POCKET.md` - the Metal tower driver's three Pocket TTS seats: the + codec, the frame loop and the text prompt. - `ARCHITECTURE_GPU_TOWER_VULKAN.md` - the Vulkan tower driver's row classes and attention routes, its encode chains, and the Vulkan ASR-decoder driver. @@ -233,67 +235,3 @@ the slab allocation), `gpu_error`. Engage is each seat's `served` counter beside `metal_tower_stats` (one encode per seat per chunk); the parity instruments are the per-stage cells on identical inputs on the f32 lane (`tests/CLAUDE.md`, the kitten file) and the served synthesis across the knob. - -### The tower driver's Pocket codec seat {#tower-pocket-codec} - -The Pocket TTS codec decoder rides the tower driver as the first seat of the family's hook record -(`register_pocket_gpu`, `ARCHITECTURE_MEDIA.md#tower-gpu-hook`): the latents of a chunk go up, the -samples come back, one command buffer. The chain is the CPU's `pocket_decode_latents` run as its -first window over the whole chunk - the stream's carries are the first window's zero rows, and the -CPU's windows equal that one shot to float noise (`ARCHITECTURE_POCKET.md#pocket-codec-stream`) - so no carry -crosses a dispatch: the 1x1 latent projection and every dense conv on the conv-gathering GEMM stamp -(a forward conv behind `k - stride` zero rows, a transposed conv's first `t * stride` output rows; -the input row stride `xs` of `St2ConvArgs` lets a conv read the padded width its producer wrote), -the depthwise upsample on the residual pool kernel, the two codec transformer layers on the dense -GEMM stamp with the layer scale, the prefill rows rope over the layer's row stride with the CPU's own -cos and sin tables (no device trigonometry) and a windowed causal attention kernel over device K/V -rows, ELU and the row copies on one row kernel, the sample column picked out of the last conv's -padded rows. Every GEMM is the -f32-exact stamp: the codec reads within 1e-6 of the CPU chain on the f32 lane that way on the M5 -Max (1e-2 on the served q8 and K-quant lanes, whose CPU chains quantize their activations), and the f16-staged twins -bought no time on this chain (the row kernels and the attention, not the GEMMs, carry its cost) -at three orders of magnitude of agreement. The slab holds every codec weight, keyed on the -addresses and the lane as the StyleTTS2 slabs are, and drops with them. A chunk past -`PK_CODEC_MAX_FRAMES` latent frames declines on its row budget (the one-shot rows scale with the -chunk) and the CPU's windows serve it; the other declines are `knob`, `shape` (a channel count off -the 8-run lattice, a head over 128 wide, a transformer width off the 32 lattice), `device` and -`gpu_error`. The frame loop is the family's second seat (`ARCHITECTURE_GPU_TOWER.md#tower-pocket-frames`). - -### The tower driver's Pocket frames seat {#tower-pocket-frames} - -The Pocket frame loop - the backbone step and the flow head for every frame of a chunk - rides -the tower driver as the family's second seat once the prompt's rows sit in the voice's caches: -the CPU's `pocket_synthesize` loop, dispatch for dispatch, in batches of `pocket_frame_batch()` frames -(eight) per command buffer, the EOS logits read back and the stop rule walked on the host between -batches (`pocket_frames_batched`, the host loop both GPU drivers run), the latents and the -conditioning rows read back once at the end. A voice slot holds the device K/V rows `[cap][d]` per -backbone layer, keyed (`tts_pk_voice_current`, the residency both drivers read) on the caches' -addresses, their fill and capacity and a sample of the rows they hold: the voice's prompt rows are transposed in once at -attach, a chunk's text rows behind them every chunk, and the frames append after those - nothing -comes back to the host, since the CPU chain forgets a chunk's rows by resetting the fill. The -key samples the rows because an address alone outlives the voice that held it: a later voice's -caches can land at the freed address with the same fill and capacity. The -backbone's q8 linears (a K-quant linear requantized to q8 from its dequantized rows, so the small -form serves) ride the decode GEMV over a 34B-block blob with their bias rows in the slab, the bias -a row add after the GEMV (the row GEMV's fused-bias q8 form as it stood before the lane-map fold, -one lane a block, ran a frame 20% slower than the decode GEMV's split-K walk - `PERF_LEDGER.md`'s -Pocket frame loop entry; the folded row stamp keeps the epilogue kinds only); every -f32 linear rides the row GEMV over the slab's rows - a simdgroup a row, x staged in threadgroup -memory - so the parity lane runs exact. Per layer: the first norm (the layer before's residual -joined in the same dispatch - the tower LayerNorm's ADD stamp), the fused q/k/v projection, the -decode rail's rope-and-store kernel rotating q and k and storing k and v into the voice slot's row -(its f32 stamp, the rope tables and the caches bound at the position's row), the attention row (a -threadgroup per head, the scores staged, one softmax, the value sum in parts), the out projection, -the residual join with the second norm, the two ffn projections around the tanh GELU. The out norm writes the conditioning row straight -into the readback rows, and the head's GEMVs carry their norm, modulation and activations as -prologue and epilogue stamps (the LayerNorm and adaLN modulation in, the SiLU, the gated -residual, the SiLU over the time constant, the tail that adds the noise and denormalizes the -latent). The noise is drawn a batch ahead in the frame order the CPU draws it, so a teacher-forced -run reads the same stream, and the last batch's draws past the frames made are rewound so the -generator ends where the CPU loop's does; a command buffer that fails declines the loop as -`gpu_error`, the generator is put back to where the chunk found it, and the CPU loop reruns the -chunk from the caches and the stream as they were. The other declines are `knob`, `shape` (the -backbone and head widths off the 32 lattice, a head size other than the 64 the attention row is -stamped for, a cache past 2048 rows, a width past the GEMV's 4096-wide stage) and `device`. On the served lanes the seat reads within the CPU chain's own -distance from the reference: the CPU quantizes the activations it feeds a q8 or K-quant plane, the -tower feeds them f32. diff --git a/modules/dasLLAMA/ARCHITECTURE_GPU_TOWER_POCKET.md b/modules/dasLLAMA/ARCHITECTURE_GPU_TOWER_POCKET.md new file mode 100644 index 0000000000..78e295d560 --- /dev/null +++ b/modules/dasLLAMA/ARCHITECTURE_GPU_TOWER_POCKET.md @@ -0,0 +1,89 @@ +# dasLLAMA Architecture - the Metal tower driver's Pocket TTS seats + +Companion to `ARCHITECTURE_GPU_TOWER.md`; a section is cited by its anchor. This document carries +the three seats the Metal tower driver fills for the Pocket TTS family - the codec, the frame loop +and the text prompt. The Vulkan TTS driver's Pocket seats are +`ARCHITECTURE_GPU_TOWER_VULKAN_TTS.md#vk-pocket-chain`; the family itself is `ARCHITECTURE_POCKET.md`. + +### The tower driver's Pocket codec seat {#tower-pocket-codec} + +The Pocket TTS codec decoder rides the tower driver as the first seat of the family's hook record +(`register_pocket_gpu`, `ARCHITECTURE_MEDIA.md#tower-gpu-hook`): the latents of a chunk go up, the +samples come back, one command buffer. The chain is the CPU's `pocket_decode_latents` run as its +first window over the whole chunk - the stream's carries are the first window's zero rows, and the +CPU's windows equal that one shot to float noise (`ARCHITECTURE_POCKET.md#pocket-codec-stream`) - so no carry +crosses a dispatch: the 1x1 latent projection and every dense conv on the conv-gathering GEMM stamp +(a forward conv behind `k - stride` zero rows, a transposed conv's first `t * stride` output rows; +the input row stride `xs` of `St2ConvArgs` lets a conv read the padded width its producer wrote), +the depthwise upsample on the residual pool kernel, the two codec transformer layers on the dense +GEMM stamp with the layer scale, the prefill rows rope over the layer's row stride with the CPU's own +cos and sin tables (no device trigonometry) and a windowed causal attention kernel over device K/V +rows, ELU and the row copies on one row kernel, the sample column picked out of the last conv's +padded rows. Every GEMM is the +f32-exact stamp: the codec reads within 1e-6 of the CPU chain on the f32 lane that way on the M5 +Max (1e-2 on the served q8 and K-quant lanes, whose CPU chains quantize their activations), and the f16-staged twins +bought no time on this chain (the row kernels and the attention, not the GEMMs, carry its cost) +at three orders of magnitude of agreement. The slab holds every codec weight, keyed on the +addresses and the lane as the StyleTTS2 slabs are, and drops with them. A chunk past +`PK_CODEC_MAX_FRAMES` latent frames declines on its row budget (the one-shot rows scale with the +chunk) and the CPU's windows serve it; the other declines are `knob`, `shape` (a channel count off +the 8-run lattice, a head over 128 wide, a transformer width off the 32 lattice), `device` and +`gpu_error`. The frame loop is the family's second seat and the text prompt its third +(`ARCHITECTURE_GPU_TOWER_POCKET.md#tower-pocket-frames`). + +### The tower driver's Pocket frames seat {#tower-pocket-frames} + +The Pocket frame loop - the backbone step and the flow head for every frame of a chunk - rides +the tower driver as the family's second seat once the prompt's rows sit in the voice's caches: +the CPU's `pocket_synthesize` loop, dispatch for dispatch, in batches of `pocket_frame_batch()` frames +(eight) per command buffer, the EOS logits read back and the stop rule walked on the host between +batches (`pocket_frames_batched`, the host loop both GPU drivers run), the latents and the +conditioning rows read back once at the end. A voice slot holds the device K/V rows `[cap][d]` per +backbone layer, keyed (`tts_pk_voice_current`, the residency both drivers read) on the caches' +addresses, their fill and capacity and a sample of the rows they hold: the voice's prompt rows are transposed in once at +attach, a chunk's text rows behind them every chunk - uploaded from the host caches, or left by the +prompt seat below - and the frames append after those: the frames' rows never come back to the host, +since the CPU chain forgets a chunk's rows by resetting the fill. The +key samples the rows because an address alone outlives the voice that held it: a later voice's +caches can land at the freed address with the same fill and capacity. The +backbone's q8 linears (a K-quant linear requantized to q8 from its dequantized rows, so the small +form serves) ride the decode GEMV over a 34B-block blob with their bias rows in the slab, the bias +a row add after the GEMV (the row GEMV's fused-bias q8 form as it stood before the lane-map fold, +one lane a block, ran a frame 20% slower than the decode GEMV's split-K walk - `PERF_LEDGER.md`'s +Pocket frame loop entry; the folded row stamp keeps the epilogue kinds only); every +f32 linear rides the row GEMV over the slab's rows - a simdgroup a row, x staged in threadgroup +memory - so the parity lane runs exact. Per layer: the first norm (the layer before's residual +joined in the same dispatch - the tower LayerNorm's ADD stamp), the fused q/k/v projection, the +decode rail's rope-and-store kernel rotating q and k and storing k and v into the voice slot's row +(its f32 stamp, the rope tables and the caches bound at the position's row), the attention row (a +threadgroup per head, the scores staged, one softmax, the value sum in parts), the out projection, +the residual join with the second norm, the two ffn projections around the tanh GELU. The out norm writes the conditioning row straight +into the readback rows, and the head's GEMVs carry their norm, modulation and activations as +prologue and epilogue stamps (the LayerNorm and adaLN modulation in, the SiLU, the gated +residual, the SiLU over the time constant, the tail that adds the noise and denormalizes the +latent). The noise is drawn a batch ahead in the frame order the CPU draws it, so a teacher-forced +run reads the same stream, and the last batch's draws past the frames made are rewound so the +generator ends where the CPU loop's does; a command buffer that fails declines the loop as +`gpu_error`, the generator is put back to where the chunk found it, and the CPU loop reruns the +chunk from the caches and the stream as they were. The other declines are `knob`, `shape` (the +backbone and head widths off the 32 lattice, a head size other than the 64 the attention row is +stamped for, a cache past 2048 rows, a width past the GEMV's 4096-wide stage) and `device`. On the served lanes the seat reads within the CPU chain's own +distance from the reference: the CPU quantizes the activations it feeds a q8 or K-quant plane, the +tower feeds them f32. + +The text prompt is the family's third seat, the rows form of the same layers: a chunk's text rows +(the host's embedding gather, `pocket_prompt_embed`) through the frames slab's layers at the +positions after the voice's rows, each f32 linear the exact f32 GEMM and each q8 linear the prefill +ladder's q8-blob GEMM (`pf_enc_q8_mm` over the frames blob's regions, the bias row after it), the +rope from the voice slot's tables at the first row's position (`MetalPkRopeTab`, the position-based +rope both homes stamp from `GkRopeTab`), the keys and values into the voice slot's rows at those +positions and the attention over the slot's rows below them. The last layer ends at its keys and +values: a prompt's residual is read by nothing, so its attention, out projection and FFN are not run - +the CPU chain's `transformer_rows` makes the same cut under `kv_only_last` for the prompt and for the voice state. The rows +come back to the host caches after the command buffer (`kv_cache_append_rows`), so the CPU frame loop +and the frames seat's upload read the same cache either way; a served prompt also leaves +`TtsPkPromptDev` naming the voice, the fill and the row count, which the frames call after it spends +on entry - served or declined, so no decline leaves it for a later chunk - to skip uploading rows the +device already holds; a frames call over another voice, fill or prompt finds the record stale and +uploads every row. The prompt declines as the frames seat does (`knob`, +`shape` - the frames admission plus the GEMM widths on the 64 lattice - `device`, `gpu_error`). diff --git a/modules/dasLLAMA/ARCHITECTURE_GPU_TOWER_VULKAN_TTS.md b/modules/dasLLAMA/ARCHITECTURE_GPU_TOWER_VULKAN_TTS.md index 7db752c51e..5850791a80 100644 --- a/modules/dasLLAMA/ARCHITECTURE_GPU_TOWER_VULKAN_TTS.md +++ b/modules/dasLLAMA/ARCHITECTURE_GPU_TOWER_VULKAN_TTS.md @@ -2,9 +2,9 @@ Companion to `ARCHITECTURE_GPU_TOWER_VULKAN.md`; a section is cited by its anchor. This document carries the seats the Vulkan TTS driver serves: the StyleTTS2 synthesis seats of the kitten and -kokoro families, and the Pocket TTS codec and frames seats. The Metal twin of every StyleTTS2 seat +kokoro families, and the Pocket TTS codec, frames and prompt seats. The Metal twin of every StyleTTS2 seat is `ARCHITECTURE_GPU_TOWER.md#tower-tts-chain`, of the Pocket seats -`ARCHITECTURE_GPU_TOWER.md#tower-pocket-codec` and `ARCHITECTURE_GPU_TOWER.md#tower-pocket-frames`, +`ARCHITECTURE_GPU_TOWER_POCKET.md#tower-pocket-codec` and `ARCHITECTURE_GPU_TOWER_POCKET.md#tower-pocket-frames`, and the CPU chain is the specification, dispatch for dispatch. The GPU backend role table these sections build on stays in `ARCHITECTURE_GPU.md#gpu-backends`. @@ -130,8 +130,8 @@ host inverse STFT, as the Metal twin does. ### The Pocket seats on Vulkan {#vk-pocket-chain} -The Pocket family's two seats ride the same driver and knob, registered through -`register_pocket_gpu`. Every Pocket linear is f32 rows in the slab - a q8 or K-quant file +The Pocket family's three seats - the codec, the frame loop and the text prompt - ride the same +driver and knob, registered through `register_pocket_gpu`. Every Pocket linear is f32 rows in the slab - a q8 or K-quant file dequantized through the active repack at slab time; the driver carries no q8 blob route, so the served lanes read the weights at f32 where the CPU chain reads quants. The codec seat is one submit over the whole run (at most 512 frames; a longer chunk declines by shape and the CPU's @@ -149,7 +149,8 @@ ratio the ELU (`TtsPkRowsElu`), the transposed upsample and the ELU-conv-ELU-con block, the last ELU, dec_out and the sample column copied out. The frames seat: the voice's K/V rows live on the device per backbone layer as [cap][d] rows under a key over the host caches, with the rope tables for every position they can hold; the chunk's text rows come up -before the loop; each frame is the CPU's `frame_step` and `head_step` at t = 1 - the input +before the loop, but for the rows a served prompt seat left there (`TtsPkPromptDev`, spent on entry +by the frames call after it, served or declined); each frame is the CPU's `frame_step` and `head_step` at t = 1 - the input linear, per layer five dispatches (the qkv row off the residual's norm with its q span roped in place and its k and v spans roped and stored into the caches' row at the frame's position, the attention over the cache, the out projection added into the residual under its layer scale, the @@ -165,7 +166,15 @@ the four lattice - and the tail that adds the noise row and denormalizes the lat batches of eight a submit (`set_pocket_frame_batch`, the one knob both drivers read), the EOS rule walked on the host between batches from the logits read back, the generator rewound past the frames made - the one host loop both drivers run (`pocket_frames_batched`, `dasllama_pocket.das`), each -handing it its noise sink, its submit and its EOS source. +handing it its noise sink, its submit and its EOS source. The prompt seat is the codec +transformer's layer loop (`ts_pk_rows_tf`, the one rows form both chains dispatch) over the frames +slab's layers and the chunk's embedding rows (the host's gather, `pocket_prompt_embed`) at the +positions after the voice's rows: the rope from the voice slot's tables at the first row's +position (`TtsPkRope`, stamped from the `GkRopeTab` template both homes share), the keys and +values into the voice slot's rows at those positions, the attention over the slot's rows below +them, the last layer ending at its keys and values (only the caches are read after a prompt - the +CPU chain's `transformer_rows` makes the same cut under `kv_only_last`); the rows come back to the host caches after the +submit, so the CPU frame loop reads them as its own. The declines: `knob`, `shape` (a width off the 64 lattice, a head width other than 64 or 128, more than 512 tokens for the attention stage, an LSTM direction over 256 hidden), `device` (the diff --git a/modules/dasLLAMA/ARCHITECTURE_POCKET.md b/modules/dasLLAMA/ARCHITECTURE_POCKET.md index 6a1b33872d..afef42619e 100644 --- a/modules/dasLLAMA/ARCHITECTURE_POCKET.md +++ b/modules/dasLLAMA/ARCHITECTURE_POCKET.md @@ -28,12 +28,15 @@ The TTS block home, facade and phoneme families are `ARCHITECTURE_TTS.md`. block home's but the residual add, the towers' `add_inplace_rows`: the transformer layer runs on `linear_rows`, `layernorm_rows_into`, `rope_rows`, `attention_causal_rows` over a `TtsKvCache`, `gelu`, `layer_scale_rows`; the codec on `conv1d_rows`, `conv1d_rows_transposed_depthwise` and - `elu_rows`. The codec decoder and the - frame loop are the two seats of the family's hook record (`PocketGpuDriver`, - `register_pocket_gpu`, `pocket_gpu_stats`; `ARCHITECTURE_GPU_TOWER.md#tower-pocket-codec` and - `ARCHITECTURE_GPU_TOWER.md#tower-pocket-frames`): a registered driver gets the first refusal of `pocket_decode_latents` and of the - loop inside `pocket_synthesize` (once the prompt's rows sit in the caches), and the CPU form - serves a decline; a served loop's wall reads as the backbone's timing, its head timing zero. + `elu_rows`. The codec decoder, the + frame loop and the text prompt are the three seats of the family's hook record (`PocketGpuDriver`, + `register_pocket_gpu`, `pocket_gpu_stats`; `ARCHITECTURE_GPU_TOWER_POCKET.md#tower-pocket-codec` and + `ARCHITECTURE_GPU_TOWER_POCKET.md#tower-pocket-frames`): a registered driver gets the first refusal of `pocket_decode_latents`, of the + prompt's rows into the caches (`prompt_rows`, over the embedding rows `pocket_prompt_embed` gathers on + either rail) and of the loop inside `pocket_synthesize` once those rows sit there, and the CPU form + serves a decline; a served loop's wall reads as the backbone's timing, its head timing zero. A prompt's + residual is read by nothing after its last layer's keys and values, so the CPU chain runs the prompt and + the voice state through `transformer_rows` under `kv_only_last` - every layer whole but the last, which ends at its cache append. Both drivers run the served loop's host side through `pocket_frames_batched` (the batch of frames a submit carries, the noise draws, the end-of-speech check) at the one batch knob `set_pocket_frame_batch` / `pocket_frame_batch()`. diff --git a/modules/dasLLAMA/ARCHITECTURE_TTS.md b/modules/dasLLAMA/ARCHITECTURE_TTS.md index 7f9e51ded3..6458fa0abb 100644 --- a/modules/dasLLAMA/ARCHITECTURE_TTS.md +++ b/modules/dasLLAMA/ARCHITECTURE_TTS.md @@ -85,8 +85,8 @@ buffers, the chunk cap and the idle release, the streamed source - is `ARCHITECT - **`dasllama_tts_slab.das`** - the shared slab writer both GPU tower drivers build their TTS slabs through: the writer struct and allocator, the q8 and K-quant dequant into f32 rows over the job pool, the conv, linear and norm row writers, the block slot writers (an LSTM transposed or not by - `lstm_transposed`), the two-pass build, every slab key, `tts_style_rows`, the Pocket frames admission - and voice record, and the decoder shape walks, generic over the driver's record and `tts_note_*` + `lstm_transposed`), the two-pass build, every slab key, `tts_style_rows`, the Pocket cache admissions + (frames, prompt), the voice record and the prompt's residency record, and the decoder shape walks, generic over the driver's record and `tts_note_*` hooks. Element offsets only - a driver turns them into its own binding offsets - and no device call. - **`dasllama_styletts2.das`** - the StyleTTS2-lineage model both families share: the weight map of the converted GGUF (conv geometry rides as `styletts2.conv.` metadata, so the diff --git a/modules/dasLLAMA/PERF_LEDGER.md b/modules/dasLLAMA/PERF_LEDGER.md index 4759236a90..d46d18e2f0 100644 --- a/modules/dasLLAMA/PERF_LEDGER.md +++ b/modules/dasLLAMA/PERF_LEDGER.md @@ -11,6 +11,33 @@ what it costs today and what the fix would change. ## Entries +- **MEASURED (2026-10-06, `direction-grade`) - Pocket's text prompt on the device takes a third of its request on + both boxes; the thirteen TTS kernels moved to templates both homes stamp read as before.** The prompt seat runs the + chunk's text rows through the frames slab's layers at the voice's positions (`ARCHITECTURE_GPU_TOWER_POCKET.md#tower-pocket-frames`, + `ARCHITECTURE_GPU_TOWER_VULKAN_TTS.md#vk-pocket-chain`), the keys and values back to the host; the last layer ends at + its K/V on every rail, including the CPU's. `dasllama-server` with the TTS model alone, `harness/served_bench.das + --no-chat --tts-url --reps 8`, one sentence, master and the branch alternated twice in one sitting on each + box (the master lives first); the stages are the engine's own stage clock. The M5's tune sidecar predates the binary. + - M5 Max, Metal, ms a request, master / branch: Pocket q8 67.5 / 49.5 and 67.5 / 49.4 (4.56 s of speech on master, + 4.64 on the branch - the device prompt moves the EOS frame by one), its stages prompt 15.5 / 2.1, backbone + 38.4 / 34.2 (the attention's unrolled loads and the frames call's upload of rows the device already holds), + codec 8.8 / 8.6; Kokoro-82M 73.5 / 73.5 (a third branch life 73.5, two branch lives void at cv 3.3% and 3.4%), + its stages the same to the tenth (decoder 41.22 / 41.29); Kitten mini 64.5 / 64.5 (one branch life void at cv + 4.1%), decoder 29.37 / 29.36. + - RTX PRO 4500, Vulkan (driver 580.159), ms a request, master / branch: Pocket q8 114.7 / 68.1 and 114.2 / 67.9, + its stages prompt 50.0 / 3.8, backbone 50.2 / 49.2, codec 8.7 / 8.8; Kokoro-82M 48.3 / 48.1 and 48.3 / 47.9 (the + first master row void at cv 9.3%), its stages the same to the tenth (bert 4.29 / 4.29, decoder 13.9 / 13.9); + Kitten mini 48.6 / 46.4 and 52.5 / 47.8, every row void on spread (that box's round trip jitters), its stages the + same to the tenth (decoder 10.5 / 10.4). The shared reduce base the row kernels now derive puts one call under + every `WgReduceBase` reduce (eight audio-tower stamps and five TTS stamps moved by it, `harness/vk_spv_diff.das`): + the backbone and bert stages above ride it, and whisper large-v3-turbo q8 served (`main.das -- --asr `, + `served_bench.das --no-chat --asr-url --clip jfk_ask_not.wav --reps 8`) reads 48.8 / 48.5 on the one + round both sides held under 3% (54.2 / 54.7 on the other, both void). + - Quality, `harness/tts_rig.py` on the 200-sentence corpus, flag-free on the M5 (the Metal seats serve), WER % / + UTMOS before -> after: Pocket f32 4.18 / 4.373 -> 3.95 / 4.368, q8 4.41 / 4.363 -> 4.23 / 4.354, kq 3.91 / 4.325 -> + 3.77 / 4.314, stuart-kq 3.32 / 4.137 -> 3.27 / 4.139 (the device prompt moves an EOS frame here and there); + kitten-nano 3.09 / 3.980 -> 3.09 / 3.976, kitten-mini 2.77 / 4.331 -> 2.77 / 4.331, kokoro 2.73 / 4.501 -> + 2.73 / 4.501. - **MEASURED (2026-10-05, `direction-grade`, `debug-jit`) - the depth-1 NextN round on the M1 Max: the draft at two rows, then chained into the verify's command buffer.** M1 Max (MacBookPro18,2, 64 GB), Metal, Qwen3.6-35B-A3B-MTP UD-Q4_K_M, `benchmarks/lcpp_bench.das` as the `-jit` script (`-no-module-cache`): @@ -1234,7 +1261,7 @@ what it costs today and what the fix would change. pays the device bring-up and the slab builds; of the spread only the min named below is in the record, the rest is the run's log). Pocket q8 (`pocket-tts-en-q8.gguf`, alba): 132 ms a sentence, the first 641, a steady sentence about 99 - the prompt 27 ms of it on the CPU chain - (`followup_vulkan.md` 107), the backbone 62, the codec 10. kokoro (`kokoro-82m.gguf`): 98 ms, + (the prompt seat since serves it on the device, the entry above), the backbone 62, the codec 10. kokoro (`kokoro-82m.gguf`): 98 ms, the first 785, min 36, the steady sentences about 62, against the torch CUDA reference row of 52 ms a sentence (`external`: `harness/tts_ref_bench.py --device cuda --models kokoro-82m:af_heart --limit 20` on the pod under @@ -1290,7 +1317,7 @@ what it costs today and what the fix would change. before the change: sampled within 5% of greedy. `lcpp_bench` carries no sampler flag (its `--mtp-temp` sets a temperature alone), so no board row holds the sampled rate. - **LANDED (2026-09-25) - the Pocket TTS frame loop rides the Metal tower as the family's second - seat (`ARCHITECTURE_GPU_TOWER.md#tower-pocket-frames`): the backbone step and the flow head for every + seat (`ARCHITECTURE_GPU_TOWER_POCKET.md#tower-pocket-frames`): the backbone step and the flow head for every frame, eight frames a command buffer over a per-voice device K/V slot, the EOS rule on the host between batches, the q8 backbone on the decode GEMV, the head's GEMVs carrying their norm and activations, the attention row a threadgroup a head with its scores staged.** The bare q8 @@ -1338,7 +1365,7 @@ what it costs today and what the fix would change. `harness/tts_rig.py` on the 200-sentence corpus against the codec-seat rows: q8 4.23 / 4.364 -> 4.18 / 4.364, kq 4.04 / 4.333 -> 3.91 / 4.325, stuart-kq 3.23 / 4.117 -> 3.36 / 4.136 (WER / UTMOS, a word or two of the 2201 either way, the UTMOS within a hundredth). The levers left are `followup_metal.md` sec.27. -- **LANDED (2026-09-25) - the Pocket TTS codec rides the Metal tower (`ARCHITECTURE_GPU_TOWER.md#tower-pocket-codec`): a chunk's latents up, its samples back, one command buffer, the CPU's windowed codec +- **LANDED (2026-09-25) - the Pocket TTS codec rides the Metal tower (`ARCHITECTURE_GPU_TOWER_POCKET.md#tower-pocket-codec`): a chunk's latents up, its samples back, one command buffer, the CPU's windowed codec run as one shot over the chunk on the f32-exact GEMM stamps, the K-quant transformer linears dequantized into the slab so the small form serves too.** Box: the M5 Max, every das figure `-jit` on this tree with `DAS_TUNE_MANIFEST=performance/m5.tune.json` (its runtime section diff --git a/modules/dasLLAMA/REVIEW.md b/modules/dasLLAMA/REVIEW.md index 67e86dc4b0..a6553b1c9e 100644 --- a/modules/dasLLAMA/REVIEW.md +++ b/modules/dasLLAMA/REVIEW.md @@ -273,8 +273,9 @@ a model card (the provenance-and-licence page beside a released model or pack), adopt or reject a model, a dataset, or a dependency; anywhere else in prose it is a defect.** **A diff that moves a family encode stage - a `dasllama/dasllama_.das` stage that turns -input into embeddings - onto a GPU hook leaves the CPU form in place and changes none of its -arithmetic.** The CPU form serves every box with no driver. +input into embeddings or fills a cache the next stage reads - onto a GPU hook leaves the CPU form +in place and changes no value that code after the stage reads.** The CPU form serves every box +with no driver. **A call to a `set_*_q8` lane setter - one that picks whether a model family's weights run the q8 or the float path - in a file under `dasllama/` or `harness/`, outside the body of another diff --git a/modules/dasLLAMA/REVIEW_GPU_KERNEL_BODY.md b/modules/dasLLAMA/REVIEW_GPU_KERNEL_BODY.md index 444eddda0a..c3e2d28063 100644 --- a/modules/dasLLAMA/REVIEW_GPU_KERNEL_BODY.md +++ b/modules/dasLLAMA/REVIEW_GPU_KERNEL_BODY.md @@ -102,7 +102,13 @@ declares), a hand-written fold of one plain float sum or max (no compensation te carried alongside) over every lane of a Vulkan kernel's workgroup - a subgroup shuffle loop, a lane-0 loop over a `@workgroup` array, one `@workgroup` value that one lane writes and every lane reads - is a defect: derive `WgReduceBase` and call its `wg_sum`, `wg_max` or `wg_rms_inv` -instead.** A fold into more than one result - separate sums over parts of the +instead.** + +**A diff that adds or changes, in a class template in `dasllama/dasllama_gpu_kernels_common.das`, a +hand-written fold of one plain float sum or max (no compensation term, no index carried +alongside) over every lane of the workgroup - a subgroup or simd shuffle loop, a lane-0 loop over +a `@workgroup` array, one `@workgroup` value that one lane writes and every lane reads - is a +defect: derive `GkWgReduce` and call its `wg_sum_into` / `wg_max_into` instead.** A fold into more than one result - separate sums over parts of the workgroup - is not one value; `ARCHITECTURE_GPU.md#gpu-backends` names the bodies that fold that way. diff --git a/modules/dasLLAMA/REVIEW_GPU_KERNEL_CLASSES.md b/modules/dasLLAMA/REVIEW_GPU_KERNEL_CLASSES.md index 3cf34e7d89..8abf83ab5c 100644 --- a/modules/dasLLAMA/REVIEW_GPU_KERNEL_CLASSES.md +++ b/modules/dasLLAMA/REVIEW_GPU_KERNEL_CLASSES.md @@ -63,7 +63,8 @@ writes any field declared on it - fields in the stamp or in the shell may share **A diff that moves a kernel class out of the class template its siblings stamp, or off the base shell they derive from, gives the class a `//!` line above its `[metal_dispatch]` / `[vk_dispatch]` declaration naming the body difference that keeps it out of that template or -shell.** +shell - or, for a move onto a template both homes stamp (`dasllama/dasllama_gpu_kernels_common.das`), +the template it joined.** **A diff that adds or changes a `[metal_dispatch]` / `[vk_dispatch]` binding that no site writes after arming - a binding filled before the first encode and never written again - puts diff --git a/modules/dasLLAMA/REVIEW_GPU_VULKAN.md b/modules/dasLLAMA/REVIEW_GPU_VULKAN.md index 116ea330ab..540d878a49 100644 --- a/modules/dasLLAMA/REVIEW_GPU_VULKAN.md +++ b/modules/dasLLAMA/REVIEW_GPU_VULKAN.md @@ -194,7 +194,8 @@ whatever their number (`ARCHITECTURE_GPU_VULKAN.md#kq-gemv-grid-buffer`). **Never let a reduce write a workgroup slot while the previous reduce's partials still occupy it - pass the other slot, or put a `barrier()` between the two reduces.** A reduce sums a value across -the workgroup through a `@workgroup` staging array - every `WgReduceBase` reduce that takes a `slot` +the workgroup through a `@workgroup` staging array - every `WgReduceBase` reduce +(`dasllama/dasllama_vulkan_classes.das`) or `GkWgReduce` reduce (`dasllama/dasllama_gpu_kernels_common.das`) that takes a `slot` argument, 0 or 1 (`dasllama/dasllama_vulkan_classes.das`); a slot is the run of partials one reduce writes into that array. A reduce carries one barrier, so a thread still summing the first reduce's partials would read the second's writes out of the same slot. @@ -252,14 +253,15 @@ profile to another form's role names. returned before it submits any command that writes the buffer that copy reads.** The host's wait is the only order between the copy's read and that write. -**A diff that adds an integer division or modulo to `dasllama/dasllama_vulkan_classes.das` kernel -code - a kernel body or any method it reaches through calls, `dasllama/dasllama_gpu_math.das`'s -helpers included - or changes one's divisor, the values that feed it, or the clamp or `if` that -guards it, where the divisor is not a literal or a template constant, clamps the divisor to at +**A diff that adds an integer division or modulo to kernel code in `dasllama/dasllama_vulkan_classes.das` +or `dasllama/dasllama_gpu_kernels_common.das` - a kernel body or any method it calls, +`dasllama/dasllama_gpu_math.das`'s helpers included - or changes its divisor, the values that feed it, +or its guard, where the divisor is not a literal or a template constant, clamps the divisor to at least one (`max(1u, ...)`) before the division, or places the division inside an `if` whose -condition tests that same divisor expression and is false when it is zero.** An integer division -by zero is undefined in SPIR-V, and some drivers evaluate both arms of a `?:` select, so neither -a select nor a test on a field the divisor is computed from guards it. +condition compares that divisor expression itself with zero.** An integer division by zero is +undefined in SPIR-V, and some drivers evaluate both arms of a `?:` select, so neither a select nor a +condition on another expression - a field the divisor is computed from, a product it is a factor +of - guards it. **A diff that adds or changes a path under `dasllama/` that re-records the split form of the one-row token command in `RDec.cmd[region]` - the chain recorded there with the attention at diff --git a/modules/dasLLAMA/dasllama/dasllama_gpu_kernels_common.das b/modules/dasLLAMA/dasllama/dasllama_gpu_kernels_common.das index 4ef7525a62..6a6de68b57 100644 --- a/modules/dasLLAMA/dasllama/dasllama_gpu_kernels_common.das +++ b/modules/dasLLAMA/dasllama/dasllama_gpu_kernels_common.das @@ -27,6 +27,64 @@ class template GkAxpy { } } +struct GkRopeTabArgs { + rows : uint //! the rows roped, row r at the tables' row r past their element bases + qs : uint //! x's row stride + col : uint //! the roped span's first column (a fused [q | k | v] row's q or k slot) + d : uint //! the span's width, its heads side by side + dh : uint //! the head size + coff : uint //! the cos table [pos][dh / 2]'s element base in its binding + soff : uint //! the sin table's element base in its binding, laid out as the cos table +} + +//! the rotary embedding over a span of rows in place from per-position tables: row r turns each pair of each head by +//! the tables' angle j at row r - NEOX the pairs (j, j + dh / 2), the CPU `rope_neox_tab_rows` (the vision mrope tables), +//! else the adjacent pairs (2j, 2j + 1), the CPU `rope_rows`; one pair an invocation over rows * d / 2 +[ |> template_struct_instance] +class template GkRopeTab { + @template_constant NEOX : bool = false + @ssbo @binding = 0 @role = "readwrite" x : array + @ssbo @binding = 1 @role = "read" cos_t : array + @ssbo @binding = 2 @role = "read" sin_t : array + @kargs pa : GkRopeTabArgs + + def pair_first(hb, j : uint) : uint { + static_if (NEOX) { + return hb + j + } else { + return hb + 2u * j + } + } + + def pair_partner_dist(hh : uint) : uint { + static_if (NEOX) { + return hh + } else { + return 1u + } + } + + def rope_tab_body { + let gid = gl_GlobalInvocationID.x + let half = pa.d / 2u + if (gid < pa.rows * half) { + let r = gid / max(half, 1u) + let pi = gid % max(half, 1u) + let hh = max(pa.dh / 2u, 1u) + let j = pi % hh + let i0 = pair_first(r * pa.qs + pa.col + (pi / hh) * pa.dh, j) + let i1 = i0 + pair_partner_dist(hh) + let tb = r * hh + j + let cv = cos_t[pa.coff + tb] + let sv = sin_t[pa.soff + tb] + let v0 = x[i0] + let v1 = x[i1] + x[i0] = v0 * cv - v1 * sv + x[i1] = v0 * sv + v1 * cv + } + } +} + struct GkConcatArgs { c : uint //! x's row width rc : uint //! the residual columns, read only under `with_res` @@ -70,3 +128,630 @@ class template GkConcat { } } } + +struct GkPkRowsArgs { + rows : uint + width : uint //! the floats a row copies + ss : uint //! src's row stride + sc : uint //! src's first column + sr0 : uint //! src's first row + ds : uint //! dst's row stride + dc : uint //! dst's first column + dr0 : uint //! dst's first row +} + +//! `rows` rows of `width` floats from src [.][ss] at column sc, rows sr0.., into dst [.][ds] at column dc, rows dr0.. - the +//! Pocket chains' row copy (a conv's window behind its carry, a carry's tail to its head, the K and V columns into their +//! caches, the sample column out); the ELU stamp maps each value on the way (the CPU `elu_rows` before a conv); one element an invocation +[ |> template_struct_instance] +class template GkPkRows { + @template_constant ELU : bool = false + @ssbo @binding = 0 @role = "read" src : array + @ssbo @binding = 1 @role = "write" dst : array + @kargs pa : GkPkRowsArgs + + def pk_rows_body { + let gid = gl_GlobalInvocationID.x + if (gid < pa.rows * pa.width) { + let w = max(pa.width, 1u) + let r = gid / w + let j = gid % w + var v = src[(pa.sr0 + r) * pa.ss + pa.sc + j] + static_if (ELU) { + v = max(v, 0.0) + min(exp(v) - 1.0, 0.0) + } + dst[(pa.dr0 + r) * pa.ds + pa.dc + j] = v + } + } +} + +struct GkSigSumArgs { + t : uint //! the rows + nb : uint //! the logits summed a row + ls : uint //! the logit rows' stride + speed : float //! the divisor of each row's sum +} + +//! raw[r] = (the sum of sigmoid(logits[r][j]) over j < nb, in j order) / speed - the CPU `styletts2_durations`' raw +//! durations (the CPU divides; a reciprocal multiply rounds a raw duration apart from it); one row an invocation +[ |> template_struct_instance] +class template GkSigSum { + @ssbo @binding = 0 @role = "read" logits : array + @ssbo @binding = 1 @role = "write" raw : array + @kargs pa : GkSigSumArgs + + def sigsum_body { + let r = gl_GlobalInvocationID.x + if (r < pa.t) { + var acc = 0.0 + var j = 0u + while (j < pa.nb) { + acc += sigmoid_f32(logits[r * pa.ls + j]) + j++ + } + raw[r] = acc / pa.speed + } + } +} + +struct GkSourceArgs { + n : uint //! the samples, frames * up + nlow : uint //! the phase frames, `resize_len(n, 1 / up)` + nh : uint //! the harmonics + up : uint + sr : float + sine_amp : float + noise_std : float + voiced_thr : float + seed : uint //! the own noise draw's key + woff : uint //! the source linear's weight row [nh] in the slab + boff : uint //! its bias [1] in the slab +} + +//! the phase per frame lowc [nh][nlow] as the running sum of the increments - TORCH: the double accumulator as a two-float +//! sum carrying its rounding error, else the plain float sum; one workgroup, invocation h < nh walks its harmonic +[ |> template_struct_instance] +class template GkSrcCumsum { + @template_constant TORCH : bool = false + @ssbo @binding = 0 @role = "read" low : array + @ssbo @binding = 1 @role = "write" lowc : array + @kargs pa : GkSourceArgs + + def src_cumsum_body { + let h = gl_LocalInvocationID.x + if (h < pa.nh) { + var acc = 0.0 + var lo = 0.0 + var k = 0u + while (k < pa.nlow) { + let x = low[h * pa.nlow + k] + static_if (TORCH) { + tts_two_sum_add(acc, lo, x) + } else { + acc += x + } + lowc[h * pa.nlow + k] = tts_phase_rad(acc, pa.up) + k++ + } + } + } +} + +//! the source's own noise draw noise_n [n][nh], `tts_hash_normal` at the seed - the captured stream's seat when no stream +//! was captured; one element an invocation +[ |> template_struct_instance] +class template GkSrcNoise { + @ssbo @binding = 0 @role = "write" noise_n : array + @kargs pa : GkSourceArgs + + def src_noise_body { + let gid = gl_GlobalInvocationID.x + if (gid < pa.n * pa.nh) { + noise_n[gid] = tts_hash_normal(pa.seed, gid) + } + } +} + +struct GkStftArgs { + n : uint //! the signal's samples + frames : uint + bins : uint + k : uint //! the frame length, each basis row's taps + hop : uint + pad : uint //! the samples padded on each side + eps : float + reoff : uint //! the real basis wre [bins][k] in wn + imoff : uint //! the imaginary basis, the same layout +} + +//! the STFT as two convs over the padded signal - REFLECT mirrors the signal past its ends, else the edge sample repeats - +//! into spec [frames][2 bins], the magnitudes sqrt(re^2 + im^2 + eps) then the phases, a zero imaginary part on the +//! negative real axis reading +pi (the CPU `magnitude_phase`); one (frame, bin) an invocation +[ |> template_struct_instance] +class template GkStft { + @template_constant REFLECT : bool = false + @ssbo @binding = 0 @role = "read" har : array + @ssbo @binding = 1 @role = "weight" wn : array + @ssbo @binding = 2 @role = "write" spec : array + @kargs pa : GkStftArgs + + def stft_body { + let gid = gl_GlobalInvocationID.x + let bins = max(pa.bins, 1u) + if (gid < pa.frames * bins) { + let f = gid / bins + let b = gid % bins + var re = 0.0 + var im = 0.0 + var tap = 0u + while (tap < pa.k) { + let p = int(f * pa.hop + tap) - int(pa.pad) + var src = 0 + static_if (REFLECT) { + src = tts_pad_reflect(p, int(pa.n)) + } else { + src = clamp(p, 0, int(pa.n) - 1) + } + let v = har[uint(src)] + re += wn[pa.reoff + b * pa.k + tap] * v + im += wn[pa.imoff + b * pa.k + tap] * v + tap++ + } + let row = f * 2u * pa.bins + spec[row + b] = tts_stft_mag(re, im, pa.eps) + spec[row + pa.bins + b] = tts_stft_phase(re, im) + } + } +} + +struct GkIstftArgs { + tt : uint //! the frames + stride : uint //! y's row width + bins : uint + k : uint //! the frame length + hop : uint + pad : uint //! the samples trimmed from each end + envelope : uint //! 1 = the summed squared window divided out where it is not negligible + n : uint //! the samples out + reoff : uint //! the real basis wre [k][bins] in wn - the transposed conv's operand layout + imoff : uint //! the imaginary basis, the same layout + wnoff : uint //! the window [k] in wn +} + +//! the inverse STFT from conv_post's rows y [tt][stride] - a bin's log magnitude, then its phase - as the overlap-added +//! transposed conv, a non-negligible envelope divided out under `envelope`; one sample an invocation +[ |> template_struct_instance] +class template GkIstft { + @ssbo @binding = 0 @role = "read" y : array + @ssbo @binding = 1 @role = "weight" wn : array + @ssbo @binding = 2 @role = "write" wave : array + @kargs pa : GkIstftArgs + + def istft_body { + let i = gl_GlobalInvocationID.x + if (i < pa.n) { + let m = i + pa.pad + let hop = max(pa.hop, 1u) + var acc = 0.0 + var env = 0.0 + var tap = 0u + while (tap < pa.k) { + if (m >= tap && (m - tap) % hop == 0u) { + let q = (m - tap) / hop + if (q < pa.tt) { + var b = 0u + while (b < pa.bins) { + let wa = tap * pa.bins + b + acc += tts_istft_term(y[q * pa.stride + b], y[q * pa.stride + pa.bins + b], wn[pa.reoff + wa], wn[pa.imoff + wa]) + b++ + } + let wv = wn[pa.wnoff + tap] + env += wv * wv + } + } + tap++ + } + wave[i] = (pa.envelope != 0u && env > 1e-11) ? acc / env : acc + } + } +} + +struct GkAddScaleArgs { + nelem : uint + scale : float +} + +//! o = (a + b) * scale over nelem elements - the residual join's 1 / sqrt(2) and the averaged stage sums; o may be a; +//! one element an invocation +[ |> template_struct_instance] +class template GkAddScale { + @ssbo @binding = 0 @role = "read" a : array + @ssbo @binding = 1 @role = "read" b : array + @ssbo @binding = 2 @role = "write" o : array + @kargs pa : GkAddScaleArgs + + def add_scale_body { + let i = gl_GlobalInvocationID.x + if (i < pa.nelem) { + o[i] = (a[i] + b[i]) * pa.scale + } + } +} + +//! the source's resample law - TORCH: torch's taps and mix, else onnxruntime's - the base of the source stamps; the taps' +//! quotients divide through `tts_div` +[ |> template_struct_instance] +class template GkSrcLaw { + @template_constant TORCH : bool = false + + def law_taps(n, to : uint; scale : float; var i0, i1 : uint&; var l0, l1 : float&) { + static_if (TORCH) { + tts_resize_taps_torch(n, to, tts_div(1.0, scale), i0, i1, l0, l1) + } else { + tts_resize_taps_onnx(n, tts_div(float(to) + 0.5, scale), i0, i1, l0, l1) + } + } + + def law_mix(x0, x1, l0, l1 : float; last : bool) : float { + static_if (TORCH) { + return tts_resize_mix_torch(x0, x1, l0, l1) + } else { + return tts_resize_mix_onnx(x0, x1, l0, l1, last) + } + } +} + +//! the low-rate phase increments low [nh][nlow] on the stamp's resample law, the first sample's initial phases folded in; +//! one element an invocation +[ |> template_struct_instance] +class template GkSrcLow : GkSrcLaw { + @ssbo @binding = 0 @role = "read" f0 : array + @ssbo @binding = 1 @role = "read" noise_u : array + @ssbo @binding = 2 @role = "write" low : array + @kargs pa : GkSourceArgs + + def src_low_body { + let gid = gl_GlobalInvocationID.x + if (gid < pa.nh * pa.nlow) { + let nlow = max(pa.nlow, 1u) + let up = max(pa.up, 1u) + let h = gid / nlow + var i0 = 0u + var i1 = 0u + var l0 = 0.0 + var l1 = 0.0 + law_taps(pa.n, gid % nlow, tts_div(1.0, float(pa.up)), i0, i1, l0, l1) + let r0 = tts_src_rad(tts_div(f0[i0 / up] * float(h + 1u), pa.sr), noise_u[h], i0, h) + let r1 = tts_src_rad(tts_div(f0[i1 / up] * float(h + 1u), pa.sr), noise_u[h], i1, h) + low[gid] = law_mix(r0, r1, l0, l1, h + 1u == pa.nh) + } + } +} + +//! the mixed source signal har [n] on the stamp's resample law: the harmonics' sines at the interpolated phases, the noise +//! rows scaled by voicing, the source linear (its weight row [nh] at woff, its bias [1] at boff in wn) and its tanh; one +//! sample an invocation +[ |> template_struct_instance] +class template GkSrcSines : GkSrcLaw { + @ssbo @binding = 0 @role = "read" f0 : array + @ssbo @binding = 1 @role = "read" lowc : array + @ssbo @binding = 2 @role = "read" noise_n : array + @ssbo @binding = 3 @role = "weight" wn : array + @ssbo @binding = 4 @role = "write" har : array + @kargs pa : GkSourceArgs + + def src_sines_body { + let i = gl_GlobalInvocationID.x + if (i < pa.n) { + var i0 = 0u + var i1 = 0u + var l0 = 0.0 + var l1 = 0.0 + law_taps(pa.nlow, i, float(pa.up), i0, i1, l0, l1) + let uv = f0[i / max(pa.up, 1u)] > pa.voiced_thr ? 1.0 : 0.0 + let noise_amp = tts_noise_amp(uv, pa.noise_std, pa.sine_amp) + var acc = wn[pa.boff] + var h = 0u + while (h < pa.nh) { + let phase = law_mix(lowc[h * pa.nlow + i0], lowc[h * pa.nlow + i1], l0, l1, h + 1u == pa.nh) + acc += tts_src_term(wn[pa.woff + h], tts_sin_big(phase) * pa.sine_amp, uv, noise_amp, noise_n[i * pa.nh + h]) + h++ + } + har[i] = tts_tanh(acc) + } + } +} + +let GK_ADAIN_EPS = 1e-5 //! the instance norms' eps - the body spells the literal (a kernel reads no module global); a driver holds it to its model's + +struct GkAdainArgs { + c : uint //! the channels + t : uint //! the rows the statistics ran over + hoff : uint //! the style fc rows' base in h: gamma [c], then beta [c] + gwoff : uint //! the norm's scale row in wn + gboff : uint //! the norm's shift row in wn + aoff : uint //! Snake's alpha row in wn, read by the Snake stamp alone +} + +//! AdaIN as the CPU `adain_affine` folds it - the column statistics, the norm's own scale and shift, then (1 + gamma) and +//! beta - followed by Snake or LeakyReLU(0.2); y may be x itself; one element an invocation over t * c +[ |> template_struct_instance] +class template GkAdain { + @template_constant SNAKE : bool = false + @ssbo @binding = 0 @role = "read" x : array + @ssbo @binding = 1 @role = "write" y : array + @ssbo @binding = 2 @role = "read" stats : array + @ssbo @binding = 3 @role = "read" h : array + @ssbo @binding = 4 @role = "weight" wn : array + @kargs pa : GkAdainArgs + + def adain_body { + let gid = gl_GlobalInvocationID.x + if (gid < pa.t * pa.c) { + let ci = gid % max(pa.c, 1u) + var scale = 0.0 + var shift = 0.0 + tts_adain_affine(stats[ci], stats[pa.c + ci], float(pa.t), 1e-5, wn[pa.gwoff + ci], wn[pa.gboff + ci], h[pa.hoff + ci], + h[pa.hoff + pa.c + ci], scale, shift) + let v = mad(x[gid], scale, shift) + static_if (SNAKE) { + let a = wn[pa.aoff + ci] + let sn = sin(a * v) + y[gid] = mad(1.0 / a, sn * sn, v) + } else { + y[gid] = v > 0.0 ? v : v * 0.2 + } + } + } +} + +struct GkDwConvRowsArgs { + c : uint //! the channels, one tap row each + k : uint //! the taps, where the stamp does not fix them + stride : uint //! the transposed stamp's stride + pad_l : uint //! the transposed stamp's cropped left edge + dil : uint //! the transposed stamp's dilation + t_in : uint //! the input rows, and the output rows of the causal and the centered stamps + t_out : uint //! the transposed stamp's output rows + woff : uint //! w [c][k]'s element base in wn + boff : uint //! the transposed stamp's bias row in wn + bnwoff : uint //! the centered stamp's folded BatchNorm scale row in wn + bnboff : uint //! the centered stamp's folded BatchNorm shift row in wn +} + +//! a depthwise conv on rows, y [t_out][c] from x [t_in][c] and w [c][k], one element an invocation; FORM is the source rule +//! and the epilogue: 0 the transposed conv gathered, bias first and one mad a tap (the CPU `conv1d_rows_transposed_depthwise`); +//! 1 the causal conv, a row reading the k - 1 rows before it (gemma4a); 2 the centered conv, folded BatchNorm and silu after it (canary) +[ |> template_struct_instance] +class template GkDwConvRows { + @template_constant FORM : int = 0 //!< spelled as its literal: a template constant folds only a literal + @template_constant K : uint = 0u //!< the taps where the stamp fixes them, else 0 and the taps are `pa.k` + @ssbo @binding = 0 @role = "read" x : array + @ssbo @binding = 1 @role = "write" y : array + @ssbo @binding = 2 @role = "weight" wn : array + @kargs pa : GkDwConvRowsArgs + + def rows_out : uint { + static_if (FORM == 0) { + return pa.t_out + } else { + return pa.t_in + } + } + + def taps : uint { + static_if (K != 0u) { + return K + } else { + return pa.k + } + } + + def src_row(t, kk : uint) : int { + static_if (FORM == 0) { + return conv_src_tr_row(t, kk, 0u, pa.pad_l, pa.dil, max(pa.stride, 1u), pa.t_in) + } else { + static_if (FORM == 1) { + return conv_src_fwd_row(t, kk, 0u, taps() - 1u, 1u, 1u, pa.t_in) + } else { + return conv_src_fwd_row(t, kk, 0u, (taps() - 1u) / 2u, 1u, 1u, pa.t_in) + } + } + } + + def dw_conv_rows_body { + let e = gl_GlobalInvocationID.x + if (e < rows_out() * pa.c) { + let c = max(pa.c, 1u) + let t = e / c + let j = e % c + let nk = taps() + var acc = 0.0 + static_if (FORM == 0) { + acc = wn[pa.boff + j] + } + var kk = 0u + while (kk < nk) { + let src = src_row(t, kk) + if (src >= 0) { + static_if (FORM == 0) { + acc = mad(wn[pa.woff + j * nk + kk], x[uint(src) * pa.c + j], acc) + } else { + acc += wn[pa.woff + j * nk + kk] * x[uint(src) * pa.c + j] + } + } + kk++ + } + static_if (FORM == 2) { + y[e] = silu_f32(acc * wn[pa.bnwoff + j] + wn[pa.bnboff + j]) + } else { + y[e] = acc + } + } + } +} + +//! the workgroup reductions every row kernel of either home shares: one barrier a reduce, every thread summing the subgroup +//! partials in the same order, so every thread holds the same sum or max; the lane primitives are each home's +//! `gk_subgroup_add` / `gk_subgroup_max`, and the result lands by reference, the statement form the MSL emitter splices +[ |> template_struct_instance] +class template GkWgReduce { + @workgroup part : float[128] //! two slots of 64 subgroup partials, one slot a reduce + + //! the workgroup sum every thread holds alike; `slot` alternates 0 / 1 between consecutive reduces; a statement with its + //! result by reference, the form the MSL emitter inlines + def wg_sum_into(v0 : float; slot : uint; var tot : float&) { + let ss = gk_subgroup_add(v0) + if (gl_SubgroupInvocationID == 0u) { + part[slot * 64u + gl_SubgroupID] = ss + } + barrier() + tot = 0.0 + var sgi = 0u + while (sgi < gl_NumSubgroups) { + tot += part[slot * 64u + sgi] + sgi++ + } + } + + //! the workgroup max every thread holds alike + def wg_max_into(v0 : float; slot : uint; var m : float&) { + let sm = gk_subgroup_max(v0) + if (gl_SubgroupInvocationID == 0u) { + part[slot * 64u + gl_SubgroupID] = sm + } + barrier() + m = -3.0e38 + var sgi = 0u + while (sgi < gl_NumSubgroups) { + m = max(m, part[slot * 64u + sgi]) + sgi++ + } + } +} + +let GK_PK_ATTN_HS = 64u //! the head size the attention row is stamped for, the body's literal: another head size declines at the driver +let GK_PK_ATTN_MAX_KEYS = 2048u //! the scores a row stages, the body's literal: a longer key range declines at the driver + +struct GkPkAttnArgs { + t : uint //! the query rows + qs : uint //! q's row stride + qcol : uint //! the query span's first column in q's rows + d : uint //! the heads' width, the cache rows' and o's row stride + pos0 : uint //! the first query's position + ctx : uint //! the keys a query sees at most, 0 = every key up to itself +} + +//! the causal cached attention of one query row over the caches' rows (the CPU `attention_causal_rows`): query i at position +//! pos0 + i sees the keys up to itself, the last `ctx` of them when ctx is nonzero; the scaled query staged, the scores staged +//! and softmaxed in two passes, the value sum in four parts of 64 columns; a workgroup of 256 a (head, row), reductions as statement methods +[ |> template_struct_instance] +class template GkPkAttn : GkWgReduce { + @ssbo @binding = 0 @role = "read" q : array + @ssbo @binding = 1 @role = "read" k : array + @ssbo @binding = 2 @role = "read" v : array + @ssbo @binding = 3 @role = "write" o : array + @kargs pa : GkPkAttnArgs + @workgroup qh : float[64] + @workgroup sc : float[2048] //! == GK_PK_ATTN_MAX_KEYS, the array's literal + @workgroup po : float[256] + //! a column's quarter of the value sum: the keys `quarter`, `quarter` + 4, ... in order, eight a step with their loads + //! in flight together, the column at `vb` in the first key's row + def value_sum_into(vb, quarter, nk : uint; var acc : float&) { + acc = 0.0 + var jj = quarter + while (jj + 28u < nk) { + let v0 = v[vb + jj * pa.d] + let v1 = v[vb + (jj + 4u) * pa.d] + let v2 = v[vb + (jj + 8u) * pa.d] + let v3 = v[vb + (jj + 12u) * pa.d] + let v4 = v[vb + (jj + 16u) * pa.d] + let v5 = v[vb + (jj + 20u) * pa.d] + let v6 = v[vb + (jj + 24u) * pa.d] + let v7 = v[vb + (jj + 28u) * pa.d] + acc += sc[jj] * v0 + acc += sc[jj + 4u] * v1 + acc += sc[jj + 8u] * v2 + acc += sc[jj + 12u] * v3 + acc += sc[jj + 16u] * v4 + acc += sc[jj + 20u] * v5 + acc += sc[jj + 24u] * v6 + acc += sc[jj + 28u] * v7 + jj += 32u + } + while (jj < nk) { + acc += sc[jj] * v[vb + jj * pa.d] + jj += 4u + } + } + + def pk_attn_body { + let wg = gl_WorkGroupID.x + let t = max(pa.t, 1u) + let h = wg / t + let i = wg % t + let lid = gl_LocalInvocationID.x + let p = pa.pos0 + i + let j0 = tts_window_start(p, pa.ctx) + let nk = p + 1u - j0 + if (lid < 64u) { + qh[lid] = q[i * pa.qs + pa.qcol + h * 64u + lid] / sqrt(float(64u)) + } + barrier() + var m = -1.0e30 + var j = lid + while (j < nk) { + let krow = (j0 + j) * pa.d + h * 64u + var sdot = 0.0 + var e = 0u + while (e < 64u) { + let k0 = k[krow + e] + let k1 = k[krow + e + 1u] + let k2 = k[krow + e + 2u] + let k3 = k[krow + e + 3u] + sdot += qh[e] * k0 + sdot += qh[e + 1u] * k1 + sdot += qh[e + 2u] * k2 + sdot += qh[e + 3u] * k3 + e += 4u + } + sc[j] = sdot + m = max(m, sdot) + j += 256u + } + var mall = 0.0 + wg_max_into(m, 0u, mall) + var l = 0.0 + j = lid + while (j < nk) { + let pw = exp(sc[j] - mall) + sc[j] = pw + l += pw + j += 256u + } + var lsum = 0.0 + wg_sum_into(l, 1u, lsum) + let inv = 1.0 / lsum + j = lid + while (j < nk) { + sc[j] = sc[j] * inv + j += 256u + } + barrier() + let col = lid % 64u + let quarter = lid / 64u + var vs = 0.0 + value_sum_into(j0 * pa.d + h * 64u + col, quarter, nk, vs) + po[quarter * 64u + col] = vs + barrier() + if (lid < 64u) { + var tot = 0.0 + var pp = 0u + while (pp < 4u) { + tot += po[pp * 64u + lid] + pp++ + } + o[i * pa.d + h * 64u + lid] = tot + } + } +} diff --git a/modules/dasLLAMA/dasllama/dasllama_gpu_math.das b/modules/dasLLAMA/dasllama/dasllama_gpu_math.das index b3c2b8f578..0bafc31f2d 100644 --- a/modules/dasLLAMA/dasllama/dasllama_gpu_math.das +++ b/modules/dasLLAMA/dasllama/dasllama_gpu_math.das @@ -187,3 +187,11 @@ def tts_hash_normal(seed, idx : uint) : float { let u2 = float(b >> 8u) * (1.0 / 16777216.0) return sqrt(-2.0 * log(u1)) * cos(6.28318530717959 * u2) } + +//! a / b rounded as the CPU's IEEE divide: a SPIR-V divide sits up to 2.5 ulp off, which the phase sums carry; the correction lands the same quotient over Metal's IEEE divide +def tts_div(a, b : float) : float { + var r = 1.0 / b + r = mad(mad(-b, r, 1.0), r, r) + let q = a * r + return mad(mad(-q, b, a), r, q) +} diff --git a/modules/dasLLAMA/dasllama/dasllama_metal_common.das b/modules/dasLLAMA/dasllama/dasllama_metal_common.das index cfd43b8907..df186f17ea 100644 --- a/modules/dasLLAMA/dasllama/dasllama_metal_common.das +++ b/modules/dasLLAMA/dasllama/dasllama_metal_common.das @@ -274,6 +274,7 @@ var g_pso_pk_rows : MetalComputePipeline? var g_pso_pk_rows_elu : MetalComputePipeline? var g_pso_pk_row_scale : MetalComputePipeline? var g_pso_pk_attn : MetalComputePipeline? +var g_pso_pk_rope_tab : MetalComputePipeline? var g_pso_pk_add_ln : MetalComputePipeline? var g_pso_pk_gemv_f32 : MetalComputePipeline? var g_pso_pk_gemv_f32_ln_silu : MetalComputePipeline? diff --git a/modules/dasLLAMA/dasllama/dasllama_metal_kernels.das b/modules/dasLLAMA/dasllama/dasllama_metal_kernels.das index 1d9ad77569..e41f516be7 100644 --- a/modules/dasLLAMA/dasllama/dasllama_metal_kernels.das +++ b/modules/dasLLAMA/dasllama/dasllama_metal_kernels.das @@ -3724,6 +3724,7 @@ def metal_decode_init : bool { // nolint:STYLE038 — a flat compile_pso per k g_pso_pk_rows_elu = compile_stamp(metal_pk_rows_elu_msl, ok) g_pso_pk_row_scale = compile_stamp(metal_pk_row_scale_msl, ok) g_pso_pk_attn = compile_stamp(metal_pk_attn_msl, ok) + g_pso_pk_rope_tab = compile_stamp(metal_pk_rope_tab_msl, ok) g_pso_pk_add_ln = compile_stamp(metal_pk_add_ln_msl, ok) g_pso_pk_gemv_f32 = compile_stamp(metal_pk_gemv_f32_msl, ok) g_pso_pk_gemv_f32_ln_silu = compile_stamp(metal_pk_gemv_f32_ln_silu_msl, ok) @@ -3738,7 +3739,7 @@ def metal_decode_init : bool { // nolint:STYLE038 — a flat compile_pso per k g_pso_st2_adain_snake = compile_stamp(metal_st2_adain_snake_msl, ok) g_pso_st2_leaky = compile_stamp(metal_st2_leaky_msl, ok) g_pso_st2_axpy = compile_stamp(metal_st2_axpy_msl, ok) - g_pso_st2_add = compile_stamp(MetalSt2Add_metal_ew2_msl, ok) + g_pso_st2_add = compile_stamp(metal_st2_add_msl, ok) g_pso_st2_reflect1 = compile_stamp(metal_st2_reflect1_msl, ok) g_pso_dn_gate = compile_stamp(MetalDnGate_metal_dn_gate_msl, ok) g_pso_dn_gate_sig = compile_stamp(MetalDnGateSig_metal_dn_gate_msl, ok) @@ -4687,7 +4688,7 @@ def metal_kernels_release { release_handles(g_pso_dn_scan_h, g_pso_dn_prep) release_handles(g_pso_dn_gate_sig, g_pso_hc_norm, g_pso_hc_lo, g_pso_hc_init, g_pso_hc_mix, g_pso_hc_combine, g_pso_ple_gate, g_pso_ple_conv) release_handles(g_pso_st2_attn) - release_handles(g_pso_pk_rows, g_pso_pk_rows_elu, g_pso_pk_row_scale, g_pso_pk_attn) + release_handles(g_pso_pk_rows, g_pso_pk_rows_elu, g_pso_pk_row_scale, g_pso_pk_attn, g_pso_pk_rope_tab) release_handles(g_pso_pk_add_ln, g_pso_pk_gemv_f32, g_pso_pk_gemv_f32_ln_silu, g_pso_pk_gemv_f32_gate, g_pso_pk_gemv_f32_add_silu) release_handles(g_pso_pk_gemv_f32_tail, g_pso_pk_gemv_q8_ln_silu, g_pso_pk_gemv_q8_gate, g_pso_pk_gemv_q8_add_silu, g_pso_pk_gemv_q8_tail) release_handles(g_pso_st2_adain_leaky, g_pso_row_gather, g_pso_st2_concat, g_pso_st2_pool_dw, g_pso_st2_lstm, g_pso_st2_sigsum, g_pso_st2_gelu_tanh, @@ -10776,48 +10777,10 @@ class MetalSt2ColStats : MetalTgReduceBase { } } -struct St2AdainArgs { - c : uint - t : uint - eps : float - total : uint -} - -//! AdaIN as the CPU folds it, then Snake or LeakyReLU(0.2); an affine-less norm binds ones and zeros; Snake's alpha binds on the Snake stamp alone. -[ |> template_struct_instance] -class template MetalSt2AdainT { - @template_constant SNAKE : bool = true - @ssbo @binding = 0 @role = "read" x : array //!< [t x c] - @ssbo @binding = 1 @role = "write" y : array //!< [t x c], x itself allowed - @ssbo @binding = 2 @role = "read" stats : array //!< [2 x c] from the column stats - @ssbo @binding = 3 @role = "read" @off = "hoff" h : array //!< [2 x c] the AdaIN fc of the style: gamma, then beta - @ssbo @binding = 4 @role = "weight" @off = "gwoff" gw : array //!< [c] the norm's scale - @ssbo @binding = 5 @role = "weight" @off = "gboff" gb : array //!< [c] the norm's shift - @ssbo @binding = 6 @role = "weight" @off = "aoff" @template_gate = SNAKE alpha : array //!< [c] Snake's alpha - @uniform @binding = 7 ka : St2AdainArgs - - def adain_body { - let gid = gl_GlobalInvocationID.x - if (gid >= ka.t * ka.c) { - return - } - let ci = gid % ka.c - var scale = 0.0 - var shift = 0.0 - tts_adain_affine(stats[ci], stats[ka.c + ci], float(ka.t), ka.eps, gw[ci], gb[ci], h[ci], h[ka.c + ci], scale, shift) - let v = x[gid] * scale + shift - static_if (SNAKE) { - let a = alpha[ci] - let sn = sin(a * v) - y[gid] = v + (1.0 / a) * (sn * sn) - } else { - y[gid] = v > 0.0 ? v : v * 0.2 - } - } -} - [metal_dispatch(name = "enc_st2_adain_snake", pso = "g_pso_st2_adain_snake", tg = 64, grid = "total/64", params = "total : int64")] -class MetalSt2AdainSnake : MetalSt2AdainT { +class MetalSt2AdainSnake : GkAdain { + override SNAKE = true + [metal_kernel(name="metal_st2_adain_snake_msl")] def metal_st2_adain_snake { adain_body() @@ -10825,9 +10788,7 @@ class MetalSt2AdainSnake : MetalSt2AdainT { } [metal_dispatch(name = "enc_st2_adain_leaky", pso = "g_pso_st2_adain_leaky", tg = 64, grid = "total/64", params = "total : int64")] -class MetalSt2AdainLeaky : MetalSt2AdainT { - override SNAKE = false - +class MetalSt2AdainLeaky : GkAdain { [metal_kernel(name="metal_st2_adain_leaky_msl")] def metal_st2_adain_leaky { adain_body() @@ -10856,13 +10817,14 @@ class MetalSt2Axpy : GkAxpy { } } -//! o = (g + u) * ascale out of place: the AdaIN blocks' join (1 / sqrt 2) and the generator's averaged stage sum -//! into a third buffer; the in-place sums take the residual add (enc_add_c). -[metal_dispatch(name = "enc_st2_add", pso = "g_pso_st2_add", tg = 64, grid = "total/64", params = "total : int64")] -class MetalSt2Add : MetalEw2T { - override ACT = 2 - override HAS_ASCALE = true - override OUT = true +//! the residual join's add-and-scale: the AdaIN blocks' join (1 / sqrt 2) and the generator's averaged stage sum, in place where o is a +//! - off the `MetalEw2T` family onto the template both homes stamp (`GkAddScale`), whose body is the Vulkan twin's +[metal_dispatch(name = "enc_st2_add", pso = "g_pso_st2_add", tg = 64, grid = "nelem/64", params = "nelem : int64")] +class MetalSt2Add : GkAddScale { + [metal_kernel(name="metal_st2_add_msl")] + def metal_st2_add { + add_scale_body() + } } //! y [t + 1][c] = x [t][c] behind one reflected row (row 0 = x's row 1): the last stage's left reflect pad. @@ -10934,43 +10896,12 @@ class MetalSt2Concat : GkConcat { } } -//! The depthwise transposed pool on rows: each output row gathers the input rows its taps reach, w [c][k]. -struct St2PoolArgs { - c : uint - k : uint - stride : uint - pad_l : uint - dil : uint - t_in : uint - total : uint -} - +//! The depthwise transposed pool on rows - the transposed form of the depthwise conv both homes stamp. [metal_dispatch(name = "enc_st2_pool_dw", pso = "g_pso_st2_pool_dw", tg = 64, grid = "total/64", params = "total : int64")] -class MetalSt2PoolDw { - @ssbo @binding = 0 @role = "read" x : array - @ssbo @binding = 1 @role = "write" y : array - @ssbo @binding = 2 @role = "weight" @off = "woff" w : array - @ssbo @binding = 3 @role = "weight" @off = "boff" b : array - @uniform @binding = 4 ka : St2PoolArgs - +class MetalSt2PoolDw : GkDwConvRows { [metal_kernel(name="metal_st2_pool_dw_msl")] def metal_st2_pool_dw { - let gid = gl_GlobalInvocationID.x - if (gid >= ka.total) { - return - } - let p = gid / ka.c - let j = gid % ka.c - var acc = b[j] - var kk = 0u - while (kk < ka.k) { - let src = conv_src_tr_row(p, kk, 0u, ka.pad_l, ka.dil, ka.stride, ka.t_in) - if (src >= 0) { - acc += w[j * ka.k + kk] * x[uint(src) * ka.c + j] - } - kk++ - } - y[gid] = acc + dw_conv_rows_body() } } @@ -11120,384 +11051,109 @@ class MetalSt2GeluTanh : MetalEw1BaseT { } } +//! the duration sigmoid sums - the CPU `styletts2_durations`' raw durations [metal_dispatch(name = "enc_st2_sigsum", pso = "g_pso_st2_sigsum", tg = 64, grid = "t/64", params = "t : int64")] -class MetalSt2SigSum { - @ssbo @binding = 0 @role = "read" logits : array - @ssbo @binding = 1 @role = "write" raw : array - @uniform @binding = 2 nb : uint - @uniform @binding = 3 ls : uint - @uniform @binding = 4 speed : float - @uniform @binding = 5 t : uint - +class MetalSt2SigSum : GkSigSum { [metal_kernel(name="metal_st2_sigsum_msl")] def metal_st2_sigsum { - let gid = gl_GlobalInvocationID.x - if (gid >= t) { - return - } - var acc = 0.0 - var j = 0u - while (j < nb) { - acc += sigmoid_f32(logits[gid * ls + j]) - j++ - } - raw[gid] = acc / speed //! the CPU divides; a reciprocal multiply rounds a raw duration apart from it - } -} - -struct St2SourceArgs { - n : uint - nlow : uint - nh : uint - up : uint - sr : float - sine_amp : float - noise_std : float - voiced_thr : float - seed : uint - total : uint -} - -//! The low-rate phase increments per harmonic on either resample law (TORCH: the torch taps and -//! mix; else onnxruntime's), the first sample's initial phases folded in. -[ |> template_struct_instance] -class template MetalSt2SrcLowT { - @template_constant TORCH : bool = false - @ssbo @binding = 0 @role = "read" f0 : array - @ssbo @binding = 1 @role = "read" noise_u : array - @ssbo @binding = 2 @role = "write" low : array //!< [nh x nlow] - @uniform @binding = 3 ka : St2SourceArgs - - def st2_src_low_body { - let gid = gl_GlobalInvocationID.x - if (gid >= ka.nh * ka.nlow) { - return - } - let h = gid / ka.nlow - let k = gid % ka.nlow - var i0 = 0u - var i1 = 0u - var l0 = 0.0 - var l1 = 0.0 - let scale = 1.0 / float(ka.up) - static_if (TORCH) { - tts_resize_taps_torch(ka.n, k, 1.0 / scale, i0, i1, l0, l1) - } else { - tts_resize_taps_onnx(ka.n, (float(k) + 0.5) / scale, i0, i1, l0, l1) - } - let r0 = tts_src_rad(f0[i0 / ka.up] * float(h + 1u) / ka.sr, noise_u[h], i0, h) - let r1 = tts_src_rad(f0[i1 / ka.up] * float(h + 1u) / ka.sr, noise_u[h], i1, h) - static_if (TORCH) { - low[gid] = tts_resize_mix_torch(r0, r1, l0, l1) - } else { - low[gid] = tts_resize_mix_onnx(r0, r1, l0, l1, h + 1u == ka.nh) - } + sigsum_body() } } [metal_dispatch(name = "enc_st2_src_low_torch", pso = "g_pso_st2_src_low_torch", tg = 64, grid = "total/64", params = "total : int64")] -class MetalSt2SrcLowTorch : MetalSt2SrcLowT { +class MetalSt2SrcLowTorch : GkSrcLow { override TORCH = true [metal_kernel(name="metal_st2_src_low_torch_msl", fastmath=false)] def metal_st2_src_low_torch { - st2_src_low_body() + src_low_body() } } [metal_dispatch(name = "enc_st2_src_low_onnx", pso = "g_pso_st2_src_low_onnx", tg = 64, grid = "total/64", params = "total : int64")] -class MetalSt2SrcLowOnnx : MetalSt2SrcLowT { +class MetalSt2SrcLowOnnx : GkSrcLow { [metal_kernel(name="metal_st2_src_low_onnx_msl", fastmath=false)] def metal_st2_src_low_onnx { - st2_src_low_body() - } -} - -//! The phase per frame as the running sum of the increments: the torch law's double accumulator -//! as a two-float sum carrying its rounding error, the onnx law's plain float sum. -[ |> template_struct_instance] -class template MetalSt2SrcCumsumT { - @template_constant TORCH : bool = false - @ssbo @binding = 0 @role = "read" low : array - @ssbo @binding = 1 @role = "write" lowc : array //!< [nh x nlow] the phase per frame - @uniform @binding = 2 ka : St2SourceArgs - - def st2_src_cumsum_body { - let h = gl_LocalInvocationID.x - if (h >= ka.nh) { - return - } - var acc = 0.0 - var lo = 0.0 - var k = 0u - while (k < ka.nlow) { - let x = low[h * ka.nlow + k] - static_if (TORCH) { - tts_two_sum_add(acc, lo, x) - } else { - acc += x - } - lowc[h * ka.nlow + k] = tts_phase_rad(acc, ka.up) - k++ - } + src_low_body() } } [metal_dispatch(name = "enc_st2_src_cumsum_torch", pso = "g_pso_st2_src_cumsum_torch", tg = "nh", grid = "1", params = "nh : int64")] -class MetalSt2SrcCumsumTorch : MetalSt2SrcCumsumT { +class MetalSt2SrcCumsumTorch : GkSrcCumsum { override TORCH = true [metal_kernel(name="metal_st2_src_cumsum_torch_msl", fastmath=false)] //! fast math folds the compensation away def metal_st2_src_cumsum_torch { - st2_src_cumsum_body() + src_cumsum_body() } } [metal_dispatch(name = "enc_st2_src_cumsum_onnx", pso = "g_pso_st2_src_cumsum_onnx", tg = "nh", grid = "1", params = "nh : int64")] -class MetalSt2SrcCumsumOnnx : MetalSt2SrcCumsumT { +class MetalSt2SrcCumsumOnnx : GkSrcCumsum { [metal_kernel(name="metal_st2_src_cumsum_onnx_msl", fastmath=false)] def metal_st2_src_cumsum_onnx { - st2_src_cumsum_body() + src_cumsum_body() } } -//! The source's own noise draw, one normal per (sample, harmonic), into the rows the sines -//! kernel reads - the captured stream's seat when no stream was captured. +//! The source's own noise draw - the captured stream's seat when no stream was captured. [metal_dispatch(name = "enc_st2_src_noise", pso = "g_pso_st2_src_noise", tg = 64, grid = "total/64", params = "total : int64")] -class MetalSt2SrcNoise { - @ssbo @binding = 0 @role = "write" noise_n : array //!< [n x nh] - @uniform @binding = 1 ka : St2SourceArgs - +class MetalSt2SrcNoise : GkSrcNoise { [metal_kernel(name="metal_st2_src_noise_msl")] def metal_st2_src_noise { - let gid = gl_GlobalInvocationID.x - if (gid >= ka.n * ka.nh) { - return - } - noise_n[gid] = tts_hash_normal(ka.seed, gid) - } -} - -//! The mixed source signal per sample on either resample law: the harmonics' sines at the -//! interpolated phases, the noise rows scaled by voicing, the source linear and its tanh. -[ |> template_struct_instance] -class template MetalSt2SrcSinesT { - @template_constant TORCH : bool = false - @ssbo @binding = 0 @role = "read" f0 : array - @ssbo @binding = 1 @role = "read" lowc : array - @ssbo @binding = 2 @role = "read" noise_n : array //!< [n x nh] - @ssbo @binding = 3 @role = "weight" @off = "woff" lw : array //!< [nh] the source linear - @ssbo @binding = 4 @role = "weight" @off = "boff" lb : array //!< [1] its bias - @ssbo @binding = 5 @role = "write" har : array - @uniform @binding = 6 ka : St2SourceArgs - - def st2_src_sines_body { - let i = gl_GlobalInvocationID.x - if (i >= ka.n) { - return - } - var i0 = 0u - var i1 = 0u - var l0 = 0.0 - var l1 = 0.0 - static_if (TORCH) { - tts_resize_taps_torch(ka.nlow, i, 1.0 / float(ka.up), i0, i1, l0, l1) - } else { - tts_resize_taps_onnx(ka.nlow, (float(i) + 0.5) / float(ka.up), i0, i1, l0, l1) - } - let uv = f0[i / ka.up] > ka.voiced_thr ? 1.0 : 0.0 - let noise_amp = tts_noise_amp(uv, ka.noise_std, ka.sine_amp) - var acc = lb[0] - var h = 0u - while (h < ka.nh) { - var phase = 0.0 - static_if (TORCH) { - phase = tts_resize_mix_torch(lowc[h * ka.nlow + i0], lowc[h * ka.nlow + i1], l0, l1) - } else { - phase = tts_resize_mix_onnx(lowc[h * ka.nlow + i0], lowc[h * ka.nlow + i1], l0, l1, h + 1u == ka.nh) - } - acc += tts_src_term(lw[h], tts_sin_big(phase) * ka.sine_amp, uv, noise_amp, noise_n[i * ka.nh + h]) - h++ - } - har[i] = tts_tanh(acc) + src_noise_body() } } [metal_dispatch(name = "enc_st2_src_sines_torch", pso = "g_pso_st2_src_sines_torch", tg = 64, grid = "total/64", params = "total : int64")] -class MetalSt2SrcSinesTorch : MetalSt2SrcSinesT { +class MetalSt2SrcSinesTorch : GkSrcSines { override TORCH = true [metal_kernel(name="metal_st2_src_sines_torch_msl", fastmath=false)] def metal_st2_src_sines_torch { - st2_src_sines_body() + src_sines_body() } } [metal_dispatch(name = "enc_st2_src_sines_onnx", pso = "g_pso_st2_src_sines_onnx", tg = 64, grid = "total/64", params = "total : int64")] -class MetalSt2SrcSinesOnnx : MetalSt2SrcSinesT { +class MetalSt2SrcSinesOnnx : GkSrcSines { [metal_kernel(name="metal_st2_src_sines_onnx_msl", fastmath=false)] def metal_st2_src_sines_onnx { - st2_src_sines_body() + src_sines_body() } } //! spec rows [frames][2 bins], the magnitudes then the phases; a zero imaginary part on the negative real axis reads +pi. -struct St2StftArgs { - n : uint - frames : uint - bins : uint - k : uint - hop : uint - pad : uint - eps : float - total : uint -} - -//! The STFT as two convs over the padded signal: the pad law is the stamp's - REFLECT mirrors -//! the signal past its ends, else the edge sample repeats. -[ |> template_struct_instance] -class template MetalSt2StftT { - @template_constant REFLECT : bool = false - @ssbo @binding = 0 @role = "read" har : array - @ssbo @binding = 1 @role = "weight" @off = "reoff" wre : array //!< [bins x k] - @ssbo @binding = 2 @role = "weight" @off = "imoff" wim : array - @ssbo @binding = 3 @role = "write" spec : array - @uniform @binding = 4 ka : St2StftArgs - - def st2_stft_body { - let gid = gl_GlobalInvocationID.x - if (gid >= ka.frames * ka.bins) { - return - } - let f = gid / ka.bins - let b = gid % ka.bins - var re = 0.0 - var im = 0.0 - var tap = 0u - while (tap < ka.k) { - let p = int(f * ka.hop + tap) - int(ka.pad) - var src = 0 - static_if (REFLECT) { - src = tts_pad_reflect(p, int(ka.n)) - } else { - src = clamp(p, 0, int(ka.n) - 1) - } - let v = har[uint(src)] - re += wre[b * ka.k + tap] * v - im += wim[b * ka.k + tap] * v - tap++ - } - spec[f * 2u * ka.bins + b] = tts_stft_mag(re, im, ka.eps) - spec[f * 2u * ka.bins + ka.bins + b] = tts_stft_phase(re, im) - } -} - [metal_dispatch(name = "enc_st2_stft_reflect", pso = "g_pso_st2_stft_reflect", tg = 64, grid = "total/64", params = "total : int64")] -class MetalSt2StftReflect : MetalSt2StftT { +class MetalSt2StftReflect : GkStft { override REFLECT = true [metal_kernel(name="metal_st2_stft_reflect_msl")] def metal_st2_stft_reflect { - st2_stft_body() + stft_body() } } [metal_dispatch(name = "enc_st2_stft_edge", pso = "g_pso_st2_stft_edge", tg = 64, grid = "total/64", params = "total : int64")] -class MetalSt2StftEdge : MetalSt2StftT { +class MetalSt2StftEdge : GkStft { [metal_kernel(name="metal_st2_stft_edge_msl")] def metal_st2_stft_edge { - st2_stft_body() + stft_body() } } //! From conv_post's rows y [tt][stride] - a bin's log magnitude, then its phase; a non-negligible envelope divided out. -struct St2IstftArgs { - tt : uint - stride : uint - bins : uint - k : uint - hop : uint - pad : uint - envelope : uint - n : uint -} - +//! The inverse STFT from conv_post's rows. [metal_dispatch(name = "enc_st2_istft", pso = "g_pso_st2_istft", tg = 64, grid = "n/64", params = "n : int64")] -class MetalSt2Istft { - @ssbo @binding = 0 @role = "read" y : array - @ssbo @binding = 1 @role = "weight" @off = "reoff" wre : array //!< [k x bins] - the transposed conv's operand layout - @ssbo @binding = 2 @role = "weight" @off = "imoff" wim : array - @ssbo @binding = 3 @role = "weight" @off = "wnoff" window : array //!< [k] - @ssbo @binding = 4 @role = "write" wave : array - @uniform @binding = 5 ka : St2IstftArgs - +class MetalSt2Istft : GkIstft { [metal_kernel(name="metal_st2_istft_msl")] def metal_st2_istft { - let i = gl_GlobalInvocationID.x - if (i >= ka.n) { - return - } - let m = i + ka.pad - var acc = 0.0 - var env = 0.0 - var tap = 0u - while (tap < ka.k) { - if (m >= tap && (m - tap) % ka.hop == 0u) { - let q = (m - tap) / ka.hop - if (q < ka.tt) { - var b = 0u - while (b < ka.bins) { - acc += tts_istft_term(y[q * ka.stride + b], y[q * ka.stride + ka.bins + b], wre[tap * ka.bins + b], wim[tap * ka.bins + b]) - b++ - } - env += window[tap] * window[tap] - } - } - tap++ - } - wave[i] = (ka.envelope != 0u && env > 1e-11) ? acc / env : acc - } -} - -//! `rows` rows of `width` floats from src [.][ss] at column sc, rows sr0.., into dst [.][ds] at column dc, rows dr0.. -struct PkRowsArgs { - rows : uint - width : uint - ss : uint - sc : uint - sr0 : uint - ds : uint - dc : uint - dr0 : uint -} - -//! The row copy of the codec stream - a window into a conv's staging rows behind its carry, a -//! carry's tail to its head, the K and V columns into their caches, the sample column out; the -//! ELU stamp applies the activation on the way (the CPU's `elu_rows` before a conv). -[ |> template_struct_instance] -class template MetalPkRowsT { - @template_constant ELU : bool = false - @ssbo @binding = 0 @role = "read" src : array - @ssbo @binding = 1 @role = "write" dst : array - @uniform @binding = 2 ka : PkRowsArgs - - def pk_rows_body { - let gid = gl_GlobalInvocationID.x - if (gid >= ka.rows * ka.width) { - return - } - let r = gid / ka.width - let j = gid % ka.width - var v = src[(ka.sr0 + r) * ka.ss + ka.sc + j] - static_if (ELU) { - v = max(v, 0.0) + min(exp(v) - 1.0, 0.0) - } - dst[(ka.dr0 + r) * ka.ds + ka.dc + j] = v + istft_body() } } [metal_dispatch(name = "enc_pk_rows", pso = "g_pso_pk_rows", tg = 64, grid = "total/64", params = "total : int64")] -class MetalPkRows : MetalPkRowsT { +class MetalPkRows : GkPkRows { [metal_kernel(name="metal_pk_rows_msl")] def metal_pk_rows { pk_rows_body() @@ -11505,7 +11161,7 @@ class MetalPkRows : MetalPkRowsT { } [metal_dispatch(name = "enc_pk_rows_elu", pso = "g_pso_pk_rows_elu", tg = 64, grid = "total/64", params = "total : int64")] -class MetalPkRowsElu : MetalPkRowsT { +class MetalPkRowsElu : GkPkRows { override ELU = true [metal_kernel(name="metal_pk_rows_elu_msl")] @@ -11525,95 +11181,29 @@ class MetalPkRowScale : MetalRowMapT { } } -//! q [t][qs] at column qcol, k and v the caches' rows [.][d], o [t][d]; query i at position pos0 + i -//! sees the keys up to itself, the last `ctx` of them when ctx is nonzero, 2048 at most; heads of PK_ATTN_HS -//! (the literal 64 in the body - a kernel reads no module global); a threadgroup a (head, row) over staged scores. -let PK_ATTN_HS = 64 //!< the head size the attention row's body is stamped for -let PK_ATTN_MAX_KEYS = 2048l //!< the scores the attention row stages in threadgroup memory - the body's literal -struct PkAttnArgs { - t : uint - qs : uint - qcol : uint - d : uint - pos0 : uint - ctx : uint +//! the Pocket prompt's adjacent-pair rope over a span of rows from the voice slot's tables at an element base +//! (the position the rows start at): the Vulkan chain's rope, stamped here for the prompt stage +[metal_dispatch(name = "enc_pk_rope_tab", pso = "g_pso_pk_rope_tab", tg = 64, grid = "pairs/64", params = "pairs : int64")] +class MetalPkRopeTab : GkRopeTab { + [metal_kernel(name="metal_pk_rope_tab_msl")] + def metal_pk_rope_tab { + rope_tab_body() + } } -[metal_dispatch(name = "enc_pk_attn", pso = "g_pso_pk_attn", tgmem = "metal_pk_attn_msl_tgmem", tg = 256, grid = "total", params = "total : int64")] -class MetalPkAttn { - @ssbo @binding = 0 @role = "read" q : array - @ssbo @binding = 1 @role = "read" k : array - @ssbo @binding = 2 @role = "read" v : array - @ssbo @binding = 3 @role = "write" o : array - @uniform @binding = 4 ka : PkAttnArgs - @workgroup qh : float[128] - @workgroup sc : float[2048] // == PK_ATTN_MAX_KEYS - @workgroup po : float[512] - @workgroup red : float[16] +let PK_ATTN_HS = int(GK_PK_ATTN_HS) //!< the head size the attention row's body is stamped for +let PK_ATTN_MAX_KEYS = int64(GK_PK_ATTN_MAX_KEYS) //!< the scores the attention row stages in threadgroup memory - the body's literal - //! the scaled scores of keys j0.. into the stage, then their normalized softmax weights - def pk_attn_scores(h, j0, nk, lid : uint) { - var m = -1.0e30 - var j = lid - while (j < nk) { - let krow = (j0 + j) * ka.d + h * 64u - var sdot = 0.0 - var e = 0u - while (e < 64u) { - sdot += qh[e] * k[krow + e] - e++ - } - sc[j] = sdot - m = max(m, sdot) - j += 256u - } - let mall = tg_max_all(m, red) - barrier() - let inv = 1.0 / sq_exp_sum(sc, red, mall, nk, lid, 256u) - j = lid - while (j < nk) { - sc[j] *= inv - j += 256u - } - barrier() // fences the weights for the V loop - } +//! the subgroup primitives the shared reductions ride, Metal's spelling +def gk_subgroup_add(v : float) : float => simd_sum(v) +def gk_subgroup_max(v : float) : float => simd_max(v) +//! The Pocket chains' causal cached attention row; a threadgroup a (head, row) over staged scores. +[metal_dispatch(name = "enc_pk_attn", pso = "g_pso_pk_attn", tgmem = "metal_pk_attn_msl_tgmem", tg = 256, grid = "total", params = "total : int64")] +class MetalPkAttn : GkPkAttn { [metal_kernel(name="metal_pk_attn_msl")] def metal_pk_attn { - let wg = gl_WorkGroupID.x - let h = wg / ka.t - let i = wg % ka.t - let lid = gl_LocalInvocationID.x - let p = ka.pos0 + i - let j0 = tts_window_start(p, ka.ctx) - let nk = p + 1u - j0 - if (lid < 64u) { - qh[lid] = q[i * ka.qs + ka.qcol + h * 64u + lid] / sqrt(float(64u)) - } - barrier() - pk_attn_scores(h, j0, nk, lid) - let parts = 256u / 64u - if (lid < parts * 64u) { - let e = lid % 64u - let part = lid / 64u - var acc = 0.0 - var jj = part - while (jj < nk) { - acc += sc[jj] * v[(j0 + jj) * ka.d + h * 64u + e] - jj += parts - } - po[part * 64u + e] = acc - } - barrier() - if (lid < 64u) { - var tot = 0.0 - var pp = 0u - while (pp < parts) { - tot += po[pp * 64u + lid] - pp++ - } - o[i * ka.d + h * 64u + lid] = tot - } + pk_attn_body() } } diff --git a/modules/dasLLAMA/dasllama/dasllama_metal_tower.das b/modules/dasLLAMA/dasllama/dasllama_metal_tower.das index c38afffd5e..e4625d0ca0 100644 --- a/modules/dasLLAMA/dasllama/dasllama_metal_tower.das +++ b/modules/dasLLAMA/dasllama/dasllama_metal_tower.das @@ -3177,13 +3177,13 @@ struct private St2DecSlab : St2SlabBase { asr_res : St2ConvSlot encode : St2ResSlots decode : array - src_w : uint64 - src_b : uint64 - stft_re : uint64 - stft_im : uint64 - istft_re : uint64 - istft_im : uint64 - window : uint64 + src_w : int64 //!< the source linear's weight row and bias as element offsets - the sines kernel takes the slab whole + src_b : int64 + stft_re : int64 //!< the STFT and inverse STFT bases and the window as element offsets - their kernels take the slab whole + stft_im : int64 + istft_re : int64 + istft_im : int64 + window : int64 } struct private St2PredSlab : St2SlabBase { @@ -3306,13 +3306,13 @@ def private st2_dec_write(var w : TtsSlabWriter; d : St2Decoder; var s : St2DecS s.asr_res = st2_slot_of(w, t.asr_res) s.encode = st2_res_of(w, d.encode, t.encode) s.decode <- [for (b, r in d.decode, t.decode); st2_res_of(w, b, r)] - s.src_w = st2_byte_off(t.src_w) - s.src_b = st2_byte_off(t.src_b) - s.stft_re = st2_byte_off(t.stft_re) - s.stft_im = st2_byte_off(t.stft_im) - s.istft_re = st2_byte_off(t.istft_re) - s.istft_im = st2_byte_off(t.istft_im) - s.window = st2_byte_off(t.window) + s.src_w = t.src_w + s.src_b = t.src_b + s.stft_re = t.stft_re + s.stft_im = t.stft_im + s.istft_re = t.istft_re + s.istft_im = t.istft_im + s.window = t.window } @@ -3590,28 +3590,24 @@ def private st2_ln(enc : MetalComputeEncoder?; var c : St2GpuCtx; bx, by : Metal //! AdaIN and the site's activation from bin into bout; the two may be one buffer. def private st2_adain(enc : MetalComputeEncoder?; var c : St2GpuCtx; bin, bout : MetalBuffer?; t, cc : int64; a : TtsAdaIn; ns : St2NormSlot; s : array; snake : bool) { - let hoff = st2_style_rows(c, a.fc, s, cc, false) + let hoff = uint(st2_style_rows(c, a.fc, s, cc, false) / 4ul) // the style rows' element offset let u_t = st2_u(c, uint(t)) let u_c = st2_u(c, uint(cc)) enc_st2_colstats(enc, bin, c.bstats, u_t, u_c, cc) let n = t * cc - var ka = St2AdainArgs(c = uint(cc), t = uint(t), eps = ST2_ADAIN_EPS) + static_assert(GK_ADAIN_EPS == ST2_ADAIN_EPS, "the AdaIN stamps' eps is the model's") + var ka = GkAdainArgs(c = uint(cc), t = uint(t), hoff = hoff, gwoff = uint(ns.gwoff / 4ul), gboff = uint(ns.gboff / 4ul), + aoff = uint(ns.aux_off / 4ul)) if (snake) { - enc_st2_adain_snake(enc, bin, bout, c.bstats, c.bh, hoff, c.slab, ns.gwoff, c.slab, ns.gboff, c.slab, ns.aux_off, ka, n) + enc_st2_adain_snake(enc, bin, bout, c.bstats, c.bh, c.slab, ka, n) } else { - enc_st2_adain_leaky(enc, bin, bout, c.bstats, c.bh, hoff, c.slab, ns.gwoff, c.slab, ns.gboff, ka, n) + enc_st2_adain_leaky(enc, bin, bout, c.bstats, c.bh, c.slab, ka, n) } } //! o = (a + b) * scale: the in-place sum (o == a) is the residual add, another o the out-of-place stamp. -def private st2_add(enc : MetalComputeEncoder?; var c : St2GpuCtx; ba, bb, bo : MetalBuffer?; n : int64; scale : float) { - let u_scale = st2_uf(c, scale) - let u_n = st2_u(c, uint(n)) - if (bo == ba) { - enc_add_c(enc, ba, 0ul, bb, 0ul, u_n, u_scale, n) - } else { - enc_st2_add(enc, ba, 0ul, bb, 0ul, u_n, u_scale, bo, 0ul, n) - } +def private st2_add(enc : MetalComputeEncoder?; ba, bb, bo : MetalBuffer?; n : int64; scale : float) { + enc_st2_add(enc, ba, bb, bo, GkAddScaleArgs(nelem = uint(n), scale = scale), n) } def private st2_leaky(enc : MetalComputeEncoder?; var c : St2GpuCtx; bx : MetalBuffer?; n : int64; slope : float) { @@ -3659,9 +3655,9 @@ def private st2_res_blk(enc : MetalComputeEncoder?; var c : St2GpuCtx; b : St2Ad var t1 = t if (b.upsample) { t1 = conv1d_out_len(b.pool, t) - var ka = St2PoolArgs(c = uint(b.pool.cin), k = uint(b.pool.k), stride = uint(b.pool.stride), pad_l = uint(b.pool.pad_l), - dil = uint(b.pool.dilation), t_in = uint(t), total = uint(t1 * b.pool.cin)) - enc_st2_pool_dw(enc, r.br, r.bpool, c.slab, sl.pool_w, c.slab, sl.pool_b, ka, t1 * b.pool.cin) + var ka = GkDwConvRowsArgs(c = uint(b.pool.cin), k = uint(b.pool.k), stride = uint(b.pool.stride), pad_l = uint(b.pool.pad_l), + dil = uint(b.pool.dilation), t_in = uint(t), t_out = uint(t1), woff = uint(sl.pool_w / 4ul), boff = uint(sl.pool_b / 4ul)) + enc_st2_pool_dw(enc, r.br, r.bpool, c.slab, ka, t1 * b.pool.cin) st2_conv(enc, c, sl.c1, b.conv1, r.bpool, t1, r.bc1) } else { st2_conv(enc, c, sl.c1, b.conv1, r.br, t1, r.bc1) @@ -3674,15 +3670,15 @@ def private st2_res_blk(enc : MetalComputeEncoder?; var c : St2GpuCtx; b : St2Ad enc_row_gather(enc, bin, 0ul, r.bidx, bin, 0ul, bin, 0ul, r.bsc, kg, t1 * b.dim_in_s) if (b.learned_sc) { st2_conv(enc, c, sl.c1x1, b.conv1x1, r.bsc, t1, r.bsc2) - st2_add(enc, c, bout, r.bsc2, bout, n, ST2_JOIN) + st2_add(enc, bout, r.bsc2, bout, n, ST2_JOIN) } else { - st2_add(enc, c, bout, r.bsc, bout, n, ST2_JOIN) + st2_add(enc, bout, r.bsc, bout, n, ST2_JOIN) } } elif (b.learned_sc) { st2_conv(enc, c, sl.c1x1, b.conv1x1, bin, t1, r.bsc2) - st2_add(enc, c, bout, r.bsc2, bout, n, ST2_JOIN) + st2_add(enc, bout, r.bsc2, bout, n, ST2_JOIN) } else { - st2_add(enc, c, bout, bin, bout, n, ST2_JOIN) + st2_add(enc, bout, bin, bout, n, ST2_JOIN) } return t1 } @@ -3697,7 +3693,7 @@ def private st2_gen_res_block(enc : MetalComputeEncoder?; var c : St2GpuCtx; b : st2_conv(enc, c, sl.c1[j], b.convs1[j], bxt, t, by1) st2_adain(enc, c, by1, by1, t, b.c, b.adain2[j], sl.n2[j], s, true) st2_conv(enc, c, sl.c2[j], b.convs2[j], by1, t, bxt) - st2_add(enc, c, src, bxt, bout, n, 1.0) + st2_add(enc, src, bxt, bout, n, 1.0) } } @@ -3734,7 +3730,7 @@ def private st2_stage(enc : MetalComputeEncoder?; var c : St2GpuCtx; d : St2Deco let u_c = st2_u(c, uint(c2)) enc_st2_reflect1(enc, g.bxt, g.by, u_c, u_n, n) } - st2_add(enc, c, g.by, g.bsres, g.by, n, 1.0) + st2_add(enc, g.by, g.bsres, g.by, n, 1.0) let nk = d.n_kernels let nnoise = long_length(d.noise_res) for (j in range64(nk)) { @@ -3872,8 +3868,9 @@ let private ST2_MAX_DE_LSTMS = 8 //!< the duration encoder's LSTM stack the se def private st2_source_chain(enc : MetalComputeEncoder?; var c : St2GpuCtx; d : St2Decoder; cfg : SineSourceCfg; ss : TtsSourceShape; bf0, bnu, bnoise, blow, blowc, bhar, bspec : MetalBuffer?; captured : bool; seed : uint64) { let nh = cfg.n_harm - var ka = St2SourceArgs(n = uint(ss.samples), nlow = uint(ss.nlow), nh = uint(nh), up = uint(cfg.upsample), sr = cfg.sample_rate, - sine_amp = cfg.sine_amp, noise_std = cfg.noise_std, voiced_thr = cfg.voiced_thr, seed = uint(seed ^ (seed >> 32ul))) + var ka = GkSourceArgs(n = uint(ss.samples), nlow = uint(ss.nlow), nh = uint(nh), up = uint(cfg.upsample), sr = cfg.sample_rate, + sine_amp = cfg.sine_amp, noise_std = cfg.noise_std, voiced_thr = cfg.voiced_thr, seed = uint(seed ^ (seed >> 32ul)), + woff = uint(g_tw_st2_dec.src_w), boff = uint(g_tw_st2_dec.src_b)) let sl & = g_tw_st2_dec if (!captured) { enc_st2_src_noise(enc, bnoise, ka, ss.samples * nh) @@ -3881,19 +3878,19 @@ def private st2_source_chain(enc : MetalComputeEncoder?; var c : St2GpuCtx; d : if (cfg.torch_math) { enc_st2_src_low_torch(enc, bf0, bnu, blow, ka, nh * ss.nlow) enc_st2_src_cumsum_torch(enc, blow, blowc, ka, nh) - enc_st2_src_sines_torch(enc, bf0, blowc, bnoise, c.slab, sl.src_w, c.slab, sl.src_b, bhar, ka, ss.samples) + enc_st2_src_sines_torch(enc, bf0, blowc, bnoise, c.slab, bhar, ka, ss.samples) } else { enc_st2_src_low_onnx(enc, bf0, bnu, blow, ka, nh * ss.nlow) enc_st2_src_cumsum_onnx(enc, blow, blowc, ka, nh) - enc_st2_src_sines_onnx(enc, bf0, blowc, bnoise, c.slab, sl.src_w, c.slab, sl.src_b, bhar, ka, ss.samples) + enc_st2_src_sines_onnx(enc, bf0, blowc, bnoise, c.slab, bhar, ka, ss.samples) } let re & = unsafe(d.stft_fwd_re) - var ks = St2StftArgs(n = uint(ss.samples), frames = uint(ss.spec_frames), bins = uint(d.bins), k = uint(re.k), hop = uint(re.stride), - pad = uint(ST2_STFT_PAD), eps = d.stft_eps) + var ks = GkStftArgs(n = uint(ss.samples), frames = uint(ss.spec_frames), bins = uint(d.bins), k = uint(re.k), hop = uint(re.stride), + pad = uint(ST2_STFT_PAD), eps = d.stft_eps, reoff = uint(sl.stft_re), imoff = uint(sl.stft_im)) if (d.stft_pad_reflect) { - enc_st2_stft_reflect(enc, bhar, c.slab, sl.stft_re, c.slab, sl.stft_im, bspec, ks, ss.spec_frames * d.bins) + enc_st2_stft_reflect(enc, bhar, c.slab, bspec, ks, ss.spec_frames * d.bins) } else { - enc_st2_stft_edge(enc, bhar, c.slab, sl.stft_re, c.slab, sl.stft_im, bspec, ks, ss.spec_frames * d.bins) + enc_st2_stft_edge(enc, bhar, c.slab, bspec, ks, ss.spec_frames * d.bins) } } @@ -3960,7 +3957,7 @@ def private st2_decoder_chain(enc : MetalComputeEncoder?; var c : St2GpuCtx; d : cwr = blk.dim_out } if (!in_x) { //! the generator ping-pongs bx and bo from bx: the stream lands in bx first - st2_add(enc, c, g.bo, g.bo, b.bx, tt * cwr, 0.5) + st2_add(enc, g.bo, g.bo, b.bx, tt * cwr, 0.5) } } @@ -4008,9 +4005,9 @@ def private metal_styletts2_decode(d : St2Decoder; text_c : int64; source : Sine st2_decoder_chain(enc, c, d, b, g, r, text_c, frames, s_ref) st2_source_chain(enc, c, d, source, sh.ss, b.bf0r, b.bnu, b.bnoise, b.blow, b.blowc, b.bhar, b.bspec, noise.captured, seed) st2_generator_chain(enc, c, d, g, b.bx, b.bspec, r.bc1, sh.dec_width, sh.dec_rows, sh.ss.spec_frames, s_ref) - var ki = St2IstftArgs(tt = uint(sh.t_post), stride = uint(ST2_POST_WIDTH), bins = uint(d.bins), k = uint(bwd.k), hop = uint(bwd.stride), - pad = uint(ST2_STFT_PAD), envelope = d.istft_envelope ? 1u : 0u, n = uint(sh.n_out)) - enc_st2_istft(enc, r.bc1, c.slab, sl.istft_re, c.slab, sl.istft_im, c.slab, sl.window, b.bwave, ki, sh.n_out) + var ki = GkIstftArgs(tt = uint(sh.t_post), stride = uint(ST2_POST_WIDTH), bins = uint(d.bins), k = uint(bwd.k), hop = uint(bwd.stride), + pad = uint(ST2_STFT_PAD), envelope = d.istft_envelope ? 1u : 0u, n = uint(sh.n_out), reoff = uint(sl.istft_re), imoff = uint(sl.istft_im), wnoff = uint(sl.window)) + enc_st2_istft(enc, r.bc1, c.slab, b.bwave, ki, sh.n_out) } asr_prof_add("tts.gen.gpu.chain", tc) if (ran) { @@ -4219,11 +4216,7 @@ def private metal_styletts2_durations(p : St2Predictor; d : array; t : in let ran = with_compute_encoder(g_queue, err) $(enc : MetalComputeEncoder?) { // nolint:PERF026 — the error-text fetch runs only on the failure leg st2_bilstm(enc, c, p.lstm, sl.dur_lstm, bd, t, lb.xf, lb.xb, lb.y) st2_lin(enc, c, sl.dur_proj, lb.y, t, blog) - let u_nb = st2_u(c, uint(p.duration_proj.nout)) - let u_ls = st2_u(c, uint(sl.dur_proj.cout_p)) - let u_speed = st2_uf(c, speed) - let u_t = st2_u(c, uint(t)) - enc_st2_sigsum(enc, blog, braw, u_nb, u_ls, u_speed, u_t, t) + enc_st2_sigsum(enc, blog, braw, GkSigSumArgs(t = uint(t), nb = uint(p.duration_proj.nout), ls = uint(sl.dur_proj.cout_p), speed = speed), t) } if (ran) { unsafe(scratch_resize(raw, t)) @@ -4302,13 +4295,13 @@ def private metal_styletts2_albert(a : St2Albert; ids : array; @scratch @ex st2_lin(enc, c, sl.v, bx, t, bv) enc_st2_attn(enc, bq, bk, bv, bxb2, ka, a.heads * t) st2_lin(enc, c, sl.dense, bxb2, t, bxb) - st2_add(enc, c, bxb, bx, bxb, nd, 1.0) + st2_add(enc, bxb, bx, bxb, nd, 1.0) st2_ln(enc, c, bxb, bxb, t, d, ST2_ALBERT_EPS, c.slab, sl.attn_ln.gwoff, sl.attn_ln.gboff) st2_lin(enc, c, sl.ffn, bxb, t, bffn) let u_nf = st2_u(c, uint(t * ff)) enc_st2_gelu_tanh(enc, bffn, 0ul, u_nf, t * ff) st2_lin(enc, c, sl.ffn_out, bffn, t, bxb2) - st2_add(enc, c, bxb2, bxb, bx, nd, 1.0) + st2_add(enc, bxb2, bxb, bx, nd, 1.0) st2_ln(enc, c, bx, bx, t, d, ST2_ALBERT_EPS, c.slab, sl.full_ln.gwoff, sl.full_ln.gboff) } } @@ -4389,24 +4382,11 @@ def private metal_styletts2_text(te : St2TextEncoder; ids : array; @scratch return true } -struct private PkLayerSlots { - n1 : St2NormSlot - in_proj : St2ConvSlot - out_proj : St2ConvSlot - scale1_off : uint64 //!< the layer scale rows, 0 where the layer carries none - has_s1 : bool - n2 : St2NormSlot - ffn1 : St2ConvSlot - ffn2 : St2ConvSlot - scale2_off : uint64 - has_s2 : bool -} - struct private PkCodecSlab : St2SlabBase { quant : St2ConvSlot up_w : uint64 //!< the depthwise upsample's taps [c][k] and bias up_b : uint64 - layers : array + layers : array //!< the transformer's layers in the frames slab's form, every linear an f32 slot dec_in : St2ConvSlot ups : array res1 : array @@ -4420,11 +4400,21 @@ var private g_tw_pk_codec : PkCodecSlab let private PK_CODEC_MAX_FRAMES = 512l //!< the one-shot chain's row budget: past this the CPU's windows serve +//! an f32 slot of the shared writer as a frame-slab linear: the GEMM slot, no blob region +def private pk_lin_of(w : TtsSlabWriter; t : TtsConvSlot; nin, nout : int64) : PkLin { + var p = PkLin(nin = nin, nout = nout, q8 = false, f = st2_slot_of(w, t)) + if (!w.measure) { + p.u_nin = uniform_u32(uint(nin)) + p.u_nout = uniform_u32(uint(nout)) + } + return p +} + //! the shared writer's layer slots as the Metal chain binds them: byte offsets, the linears' uniforms acquired -def private pk_layer_of(w : TtsSlabWriter; t : TtsPkLayerSlots) : PkLayerSlots { - return PkLayerSlots(n1 = st2_norm_of(t.n1), in_proj = st2_slot_of(w, t.in_proj), out_proj = st2_slot_of(w, t.out_proj), - scale1_off = st2_byte_off(t.scale1_off), has_s1 = t.has_s1, n2 = st2_norm_of(t.n2), ffn1 = st2_slot_of(w, t.ffn1), ffn2 = st2_slot_of(w, t.ffn2), - scale2_off = st2_byte_off(t.scale2_off), has_s2 = t.has_s2) +def private pk_layer_of(w : TtsSlabWriter; t : TtsPkLayerSlots; l : PocketLayer) : PkFrLayer { + return PkFrLayer(n1 = st2_norm_of(t.n1), in_proj = pk_lin_of(w, t.in_proj, l.in_proj.nin, l.in_proj.nout), out_proj = pk_lin_of(w, t.out_proj, l.out_proj.nin, l.out_proj.nout), + scale1_off = st2_byte_off(t.scale1_off), has_s1 = t.has_s1, n2 = st2_norm_of(t.n2), ffn1 = pk_lin_of(w, t.ffn1, l.ffn1.nin, l.ffn1.nout), + ffn2 = pk_lin_of(w, t.ffn2, l.ffn2.nin, l.ffn2.nout), scale2_off = st2_byte_off(t.scale2_off), has_s2 = t.has_s2) } //! the shared codec layout with no rope rows in the slab: this driver's rope tables ride their own buffers (`pk_rope_tables`) @@ -4434,7 +4424,7 @@ def private pk_codec_write(var w : TtsSlabWriter; mm : PocketMimi; var s : PkCod s.quant = st2_slot_of(w, t.quant) s.up_w = st2_byte_off(t.up_w) s.up_b = st2_byte_off(t.up_b) - s.layers <- [for (l in t.layers); pk_layer_of(w, l)] + s.layers <- [for (tl, l in t.layers, mm.dec_tf.layers); pk_layer_of(w, tl, l)] s.dec_in = st2_slot_of(w, t.dec_in) s.ups <- [for (c in t.ups); st2_slot_of(w, c)] s.res1 <- [for (c in t.res1); st2_slot_of(w, c)] @@ -4487,7 +4477,7 @@ def private pk_conv(enc : MetalComputeEncoder?; var c : St2GpuCtx; slot : St2Con } def private pk_elu(enc : MetalComputeEncoder?; b : MetalBuffer?; rows, width : int64) { - var ka = PkRowsArgs(rows = uint(rows), width = uint(width), ss = uint(width), sc = 0u, sr0 = 0u, ds = uint(width), dc = 0u, dr0 = 0u) + var ka = GkPkRowsArgs(rows = uint(rows), width = uint(width), ss = uint(width), sc = 0u, sr0 = 0u, ds = uint(width), dc = 0u, dr0 = 0u) enc_pk_rows_elu(enc, b, b, ka, rows * width) } @@ -4522,9 +4512,12 @@ struct private PkTfBufs { } //! The transformer's buffers a chunk; the per-layer K/V handles land in the scratch's bk / bv. +//! a frame-slab linear's output row stride: the f32 slot's padded width, the q8 region's exact one +def private pk_lin_stride(p : PkLin) : int64 => p.q8 ? p.nout : p.f.cout_p + def private pk_tf_bufs(var c : St2GpuCtx; tf : PocketTransformer; sl : PkCodecSlab; t : int64) : PkTfBufs { - var b = PkTfBufs(bxb = st2_rows_buf(c, t, tf.d), bqkv = st2_rows_buf(c, t, sl.layers[0].in_proj.cout_p), bctx = st2_rows_buf(c, t, tf.d), - bh = st2_rows_buf(c, t, sl.layers[0].ffn1.cout_p)) + var b = PkTfBufs(bxb = st2_rows_buf(c, t, tf.d), bqkv = st2_rows_buf(c, t, pk_lin_stride(sl.layers[0].in_proj)), bctx = st2_rows_buf(c, t, tf.d), + bh = st2_rows_buf(c, t, pk_lin_stride(sl.layers[0].ffn1))) g_tw_pk_sc.bk |> clear() g_tw_pk_sc.bv |> clear() g_tw_pk_sc.bk |> reserve(length(tf.layers)) @@ -4536,59 +4529,27 @@ def private pk_tf_bufs(var c : St2GpuCtx; tf : PocketTransformer; sl : PkCodecSl return <- b } -//! The codec transformer over the residual rows bx [t][d] at positions 0..: the CPU's layer loop, dispatch for dispatch. -def private pk_transformer(enc : MetalComputeEncoder?; var c : St2GpuCtx; tf : PocketTransformer; sl : PkCodecSlab; b : PkTfBufs; bx : MetalBuffer?; t : int64) { - let d = tf.d - let dh = uint(d / tf.heads) - let nd = t * d - let u_d = st2_u(c, uint(d)) - let u_nd = st2_u(c, uint(nd)) - for (i in range(length(tf.layers))) { - let s & = unsafe(sl.layers[i]) - st2_ln(enc, c, bx, b.bxb, t, d, tf.eps, c.slab, s.n1.gwoff, s.n1.gboff) - st2_lin(enc, c, s.in_proj, b.bxb, t, b.bqkv) - let qs = uint(s.in_proj.cout_p) - let u_dh = st2_u(c, dh) - let u_np = st2_u(c, uint(t * (d / 2l))) - let u_neox = st2_u(c, 0u) - let u_qs = st2_u(c, qs) - enc_rope(enc, b.bqkv, 0ul, sl.bcos, sl.bsin, u_d, u_dh, u_np, u_neox, null, u_qs, t * (d / 2l)) - enc_rope(enc, b.bqkv, uint64(d * 4l), sl.bcos, sl.bsin, u_d, u_dh, u_np, u_neox, null, u_qs, t * (d / 2l)) - var kc = PkRowsArgs(rows = uint(t), width = uint(d), ss = qs, sc = uint(d), sr0 = 0u, ds = uint(d), dc = 0u, dr0 = 0u) - enc_pk_rows(enc, b.bqkv, g_tw_pk_sc.bk[i], kc, nd) - kc.sc = uint(2l * d) - enc_pk_rows(enc, b.bqkv, g_tw_pk_sc.bv[i], kc, nd) - var ka = PkAttnArgs(t = uint(t), qs = qs, qcol = 0u, d = uint(d), pos0 = 0u, ctx = uint(tf.context)) - enc_pk_attn(enc, b.bqkv, g_tw_pk_sc.bk[i], g_tw_pk_sc.bv[i], b.bctx, ka, tf.heads * t) - st2_lin(enc, c, s.out_proj, b.bctx, t, b.bxb) - if (s.has_s1) { - enc_pk_row_scale(enc, b.bxb, 0ul, c.slab, s.scale1_off, u_d, u_nd, nd) - } - st2_add(enc, c, bx, b.bxb, bx, nd, 1.0) - st2_ln(enc, c, bx, b.bxb, t, d, tf.eps, c.slab, s.n2.gwoff, s.n2.gboff) - st2_lin(enc, c, s.ffn1, b.bxb, t, b.bh) - let nf = t * s.ffn1.cout_p - enc_st2_gelu_tanh(enc, b.bh, 0ul, st2_u(c, uint(nf)), nf) - st2_lin(enc, c, s.ffn2, b.bh, t, b.bxb) - if (s.has_s2) { - enc_pk_row_scale(enc, b.bxb, 0ul, c.slab, s.scale2_off, u_d, u_nd, nd) - } - st2_add(enc, c, bx, b.bxb, bx, nd, 1.0) +//! The codec transformer over the residual rows bx [t][d] at positions 0..: the rows form over the codec slab's layers, its +//! keys and values into the scratch's per-layer rows. +def private pk_transformer(enc : MetalComputeEncoder?; var c : St2GpuCtx; tf : PocketTransformer; var sl : PkCodecSlab; b : PkTfBufs; bx : MetalBuffer?; t : int64) { + let f = PkRowsForm(pos0 = 0l, kv_row0 = 0l, bcos = sl.bcos, bsin = sl.bsin, kv_only_last = false) + pk_rows_tf(enc, c, tf, sl.layers, b, bx, null, t, f) $(li : int) { + return (bk = g_tw_pk_sc.bk[li], bv = g_tw_pk_sc.bv[li]) } } //! One ELU-conv-ELU-conv residual block over x [t][w] (a stage's rows) into y, through cc and cp. def private pk_res_block(enc : MetalComputeEncoder?; var c : St2GpuCtx; b : PocketResConv; s1, s2 : St2ConvSlot; bx, bcc, bcp, by : MetalBuffer?; t, w : int64) { - var ka = PkRowsArgs(rows = uint(t), width = uint(w), ss = uint(w), sc = 0u, sr0 = 0u, ds = uint(w), dc = 0u, dr0 = 0u) + var ka = GkPkRowsArgs(rows = uint(t), width = uint(w), ss = uint(w), sc = 0u, sr0 = 0u, ds = uint(w), dc = 0u, dr0 = 0u) enc_pk_rows_elu(enc, bx, bcc, ka, t * w) pk_conv(enc, c, s1, b.conv1, bcc, t, w, bcp) pk_elu(enc, bcp, t, s1.cout_p) pk_conv(enc, c, s2, b.conv2, bcp, t, s1.cout_p, by) - st2_add(enc, c, by, bx, by, t * w, 1.0) + st2_add(enc, by, bx, by, t * w, 1.0) } -def private pk_codec_chain(enc : MetalComputeEncoder?; var c : St2GpuCtx; mm : PocketMimi; sl : PkCodecSlab; blat : MetalBuffer?; +def private pk_codec_chain(enc : MetalComputeEncoder?; var c : St2GpuCtx; mm : PocketMimi; var sl : PkCodecSlab; blat : MetalBuffer?; rows : array; bwave : MetalBuffer?) { let d = mm.dec_tf.d let frames = rows[0] @@ -4596,9 +4557,9 @@ def private pk_codec_chain(enc : MetalComputeEncoder?; var c : St2GpuCtx; mm : P var bq = st2_rows_buf(c, frames, sl.quant.cout_p) pk_conv(enc, c, sl.quant, mm.quant_proj, blat, frames, mm.quant_proj.cin, bq) var bx = st2_rows_buf(c, t1, d) - var kp = St2PoolArgs(c = uint(d), k = uint(mm.upsample.k), stride = uint(mm.upsample.stride), pad_l = 0u, dil = uint(mm.upsample.dilation), - t_in = uint(frames), total = uint(t1 * d)) - enc_st2_pool_dw(enc, bq, bx, c.slab, sl.up_w, c.slab, sl.up_b, kp, t1 * d) + var kp = GkDwConvRowsArgs(c = uint(d), k = uint(mm.upsample.k), stride = uint(mm.upsample.stride), pad_l = 0u, dil = uint(mm.upsample.dilation), + t_in = uint(frames), t_out = uint(t1), woff = uint(sl.up_w / 4ul), boff = uint(sl.up_b / 4ul)) + enc_st2_pool_dw(enc, bq, bx, c.slab, kp, t1 * d) var tb = pk_tf_bufs(c, mm.dec_tf, sl, t1) pk_transformer(enc, c, mm.dec_tf, sl, tb, bx, t1) var bcur = st2_rows_buf(c, t1, sl.dec_in.cout_p) @@ -4623,11 +4584,11 @@ def private pk_codec_chain(enc : MetalComputeEncoder?; var c : St2GpuCtx; mm : P pk_elu(enc, bcur, t, w) var bo = st2_rows_buf(c, t, sl.dec_out.cout_p) pk_conv(enc, c, sl.dec_out, mm.dec_out, bcur, t, w, bo) - var kw = PkRowsArgs(rows = uint(t), width = 1u, ss = uint(sl.dec_out.cout_p), sc = 0u, sr0 = 0u, ds = 1u, dc = 0u, dr0 = 0u) + var kw = GkPkRowsArgs(rows = uint(t), width = 1u, ss = uint(sl.dec_out.cout_p), sc = 0u, sr0 = 0u, ds = 1u, dc = 0u, dr0 = 0u) enc_pk_rows(enc, bo, bwave, kw, t) } -[hot_path, unused_argument(sc), arch(at="../ARCHITECTURE_GPU_TOWER.md#tower-pocket-codec")] //! the carrier is the CPU chain's; the seat's rows live in the pool +[hot_path, unused_argument(sc), arch(at="../ARCHITECTURE_GPU_TOWER_POCKET.md#tower-pocket-codec")] //! the carrier is the CPU chain's; the seat's rows live in the pool def private metal_pocket_codec(m : PocketModel; latents : array; frames : int64; var sc : PocketScratch; @scratch @exact_size var wave : array) : bool { if (!g_tw_env_tower) { return tw_decline(MetalTowerDecline.knob) @@ -4737,6 +4698,7 @@ struct private PkVoiceSlot { bv : array bcos : MetalBuffer? bsin : MetalBuffer? + gen : uint64 //!< counted up at every rebuild: a prompt's residency record names the build it sits in } var private g_tw_pk_frames : PkFramesSlab @@ -4835,11 +4797,13 @@ def private pk_frames_write(var w : TtsSlabWriter; var qb : PkBlobWriter; m : Po } def private pk_voice_drop { + let gen = g_tw_pk_voice.gen st2_release_any(g_tw_pk_voice) unsafe { delete g_tw_pk_voice } g_tw_pk_voice.meta = TtsPkVoiceMeta() + g_tw_pk_voice.gen = gen + 1ul } //! A GEMV-served linear: q8 quants 32-wide, a K-quant plane 256-wide, f32 rows as they are. @@ -4894,7 +4858,7 @@ def private pk_frames_rebuild(m : PocketModel; key : uint64) : bool { } //! Host rows [r0, r1) of one cache as the device's [pos][d] rows, written straight into the shared buffers. -[arch(at="../ARCHITECTURE_GPU_TOWER.md#tower-pocket-frames")] +[arch(at="../ARCHITECTURE_GPU_TOWER_POCKET.md#tower-pocket-frames")] def private pk_kv_rows(kc : TtsKvCache; bk, bv : MetalBuffer?; r0, r1 : int64) { unsafe { var kp = reinterpret(metal_buffer_contents(bk)) @@ -4987,7 +4951,7 @@ def private pk_site(var c : St2GpuCtx; var bx : MetalBuffer?; xoff : uint64; var } //! The row GEMV at a site under one of the five stamp kinds. -[arch(at = "../ARCHITECTURE_GPU_TOWER.md#tower-pocket-frames")] +[arch(at = "../ARCHITECTURE_GPU_TOWER_POCKET.md#tower-pocket-frames")] def private pk_gemv(enc : MetalComputeEncoder?; var c : St2GpuCtx; p : PkLin; kind : int; s : PkGemvSite) { var ka = PkGemvArgs(nin = uint(p.nin), nout = uint(p.nout), stride = p.q8 ? 0u : uint(p.f.kdim_p), eps = POCKET_HEAD_LN_EPS) let woff = p.q8 ? 0ul : p.f.woff @@ -5050,7 +5014,7 @@ def private pk_fr_layer(enc : MetalComputeEncoder?; var c : St2GpuCtx; tf : Pock var kr = RopeStoreArgs(qd = uint(d), head_size = dh, npairs_q = uint(d / 2l), npairs = uint(d), neox = 0u, has_bias = 0u, rot = dh) enc_rpst32_c(enc, b.qkv, b.qkv, uint64(2l * d * 4l), g_tw_pk_voice.bk[li], uint64(p * d * 4l), g_tw_pk_voice.bv[li], uint64(p * d * 4l), g_tw_pk_voice.bcos, uint64(p * (d / tf.heads / 2l) * 4l), g_tw_pk_voice.bsin, g_one, kr, d) - var ka = PkAttnArgs(t = 1u, qs = uint(3l * d), qcol = 0u, d = uint(d), pos0 = uint(p), ctx = uint(tf.context)) + var ka = GkPkAttnArgs(t = 1u, qs = uint(3l * d), qcol = 0u, d = uint(d), pos0 = uint(p), ctx = uint(tf.context)) enc_pk_attn(enc, b.qkv, g_tw_pk_voice.bk[li], g_tw_pk_voice.bv[li], b.ctx, ka, tf.heads) pk_lin(enc, c, L.out_proj, b.ctx, 0ul, b.xb, 0ul) if (L.has_s1) { @@ -5141,10 +5105,11 @@ def private pk_frames_loop(m : PocketModel; var sl : PkFramesSlab; var c : St2Gp }) } -[hot_path, arch(at="../ARCHITECTURE_GPU_TOWER.md#tower-pocket-frames")] +[hot_path, arch(at="../ARCHITECTURE_GPU_TOWER_POCKET.md#tower-pocket-frames")] def private metal_pocket_frames(m : PocketModel; var vs : PocketVoiceState; n_txt : int64; std : float; frames_after_eos : int; captured, forced : array; max_frames : int64; var sc : PocketScratch) : tuple { let none = (frames = -1l, eos_at = -1l) + let dev_rows = tts_pk_prompt_dev_take(g_tw_pk_prompt_dev, vs, g_tw_pk_voice.gen, sc.rows, n_txt, m.backbone.d) //! spent on entry: a decline below must not leave it for a later chunk if (!g_tw_env_tower) { tw_decline(MetalTowerDecline.knob) return none @@ -5166,7 +5131,7 @@ def private metal_pocket_frames(m : PocketModel; var vs : PocketVoiceState; n_tx tc = ref_time_ticks() residency_flush() for (li, kc in iter_range(vs.caches), vs.caches) { - pk_kv_rows(kc, g_tw_pk_voice.bk[li], g_tw_pk_voice.bv[li], vs.len, base) + pk_kv_rows(kc, g_tw_pk_voice.bk[li], g_tw_pk_voice.bv[li], vs.len + dev_rows, base) } var sl & = g_tw_pk_frames let batch = min(pocket_frame_batch(), max_frames) @@ -5197,7 +5162,160 @@ def private metal_pocket_frames(m : PocketModel; var vs : PocketVoiceState; n_tx return r } -def private metal_pocket_driver : PocketGpuDriver => PocketGpuDriver(codec = @@metal_pocket_codec, frames = @@metal_pocket_frames) +var private g_tw_pk_prompt_dev : TtsPkPromptDev + +//! A frame-slab linear over `rows` rows: the f32 slot's exact GEMM, or the blob's q8 blocks under the prefill +//! ladder's q8 GEMM with the bias row after it (`bxh` the rows' half twin where that GEMM reads one). +def private pk_rows_lin(enc : MetalComputeEncoder?; var c : St2GpuCtx; p : PkLin; bx, bxh, by : MetalBuffer?; rows : int64) { + if (!p.q8) { + st2_lin(enc, c, p.f, bx, rows, by) + return + } + let mp = round_up(rows, 32l) + if (bxh != null) { + enc_cvt_half(enc, bx, bxh, mp * p.nin) + } + pf_enc_q8_mm(enc, g_tw_pk_frames.blob, p.qoff, bx, bxh, by, p.u_nin, p.u_nout, mp, p.nout, p.nin) + if (p.has_b) { + let n = rows * p.nout + let u_n = st2_u(c, uint(n)) + enc_add_bias_rows(enc, by, 0ul, c.slab, p.boff, p.u_nout, u_n, n) + } +} + +//! a transformer's rows form: every layer's keys and values written at row `kv_row0` of its cache rows, the rope tables' +//! row of position 0 at the first prompt position's row, the queries at `pos0`; `kv_only_last` ends the last layer at its cache rows +struct private PkRowsForm { + pos0 : int64 + kv_row0 : int64 + bcos : MetalBuffer? + bsin : MetalBuffer? + kv_only_last : bool +} + +//! A transformer over the rows x [t][d] of the buffers `b`: the CPU's layer loop, dispatch for dispatch, over frame-slab +//! layers (`pk_rows_lin` picks each linear's route; `bxh` the half twin the q8 route reads, null on an f32 slab); the +//! keys and values of layer li into the buffers `kv` names for it - the codec's scratch rows, or the voice slot. +def private pk_rows_tf(enc : MetalComputeEncoder?; var c : St2GpuCtx; tf : PocketTransformer; layers : array; b : PkTfBufs; bx, bxh : MetalBuffer?; + t : int64; f : PkRowsForm; kv : block<(li : int) : tuple>) { + let d = tf.d + let dh = uint(d / tf.heads) + let tab0 = uint(f.pos0 * (d / tf.heads / 2l)) // the tables' row of the first position + let nd = t * d + let u_d = st2_u(c, uint(d)) + let u_nd = st2_u(c, uint(nd)) + let last = length(layers) - 1 + let qs = uint(pk_lin_stride(layers[0].in_proj)) // the qkv rows' stride + var kc = GkPkRowsArgs(rows = uint(t), width = uint(d), ss = qs, sc = uint(d), sr0 = 0u, ds = uint(d), dc = 0u, dr0 = uint(f.kv_row0)) + var ra = GkRopeTabArgs(rows = uint(t), qs = qs, col = 0u, d = uint(d), dh = dh, coff = tab0, soff = tab0) + let pairs = t * (d / 2l) + for (i in range(length(layers))) { + let s & = unsafe(layers[i]) + let cache = invoke(kv, i) + let kv_only = f.kv_only_last && i == last + st2_ln(enc, c, bx, b.bxb, t, d, tf.eps, c.slab, s.n1.gwoff, s.n1.gboff) + pk_rows_lin(enc, c, s.in_proj, b.bxb, bxh, b.bqkv, t) + if (!kv_only) { + ra.col = 0u + enc_pk_rope_tab(enc, b.bqkv, f.bcos, f.bsin, ra, pairs) + } + ra.col = uint(d) + enc_pk_rope_tab(enc, b.bqkv, f.bcos, f.bsin, ra, pairs) + kc.sc = uint(d) + enc_pk_rows(enc, b.bqkv, cache.bk, kc, nd) + kc.sc = uint(2l * d) + enc_pk_rows(enc, b.bqkv, cache.bv, kc, nd) + continue if (kv_only) + var ka = GkPkAttnArgs(t = uint(t), qs = qs, qcol = 0u, d = uint(d), pos0 = uint(f.pos0), ctx = uint(tf.context)) + enc_pk_attn(enc, b.bqkv, cache.bk, cache.bv, b.bctx, ka, tf.heads * t) + pk_rows_lin(enc, c, s.out_proj, b.bctx, bxh, b.bxb, t) + if (s.has_s1) { + enc_pk_row_scale(enc, b.bxb, 0ul, c.slab, s.scale1_off, u_d, u_nd, nd) + } + st2_add(enc, bx, b.bxb, bx, nd, 1.0) + st2_ln(enc, c, bx, b.bxb, t, d, tf.eps, c.slab, s.n2.gwoff, s.n2.gboff) + pk_rows_lin(enc, c, s.ffn1, b.bxb, bxh, b.bh, t) + let nf = t * pk_lin_stride(s.ffn1) + enc_st2_gelu_tanh(enc, b.bh, 0ul, st2_u(c, uint(nf)), nf) + pk_rows_lin(enc, c, s.ffn2, b.bh, bxh, b.bxb, t) + if (s.has_s2) { + enc_pk_row_scale(enc, b.bxb, 0ul, c.slab, s.scale2_off, u_d, u_nd, nd) + } + st2_add(enc, bx, b.bxb, bx, nd, 1.0) + } +} + +//! The prompt's K/V rows [vs.len, vs.len + t) of every layer back from the voice slot's shared buffers into the host caches. +def private pk_prompt_readback(var vs : PocketVoiceState; t, d : int64) { + let v & = g_tw_pk_voice + unsafe { + for (li, kc in iter_range(vs.caches), vs.caches) { + let kp = reinterpret(metal_buffer_contents(v.bk[li])) + let vp = reinterpret(metal_buffer_contents(v.bv[li])) + kv_cache_append_rows(kc, kp + vs.len * d, vp + vs.len * d, t, d) + } + } +} + +//! The text prompt's rows through the backbone after the voice's: the frames slab's layers in the rows form at +//! positions vs.len.., the keys and values into the voice slot and back to the host caches, the slot's rows left +//! for the frames call after it. +[hot_path, arch(at="../ARCHITECTURE_GPU_TOWER_POCKET.md#tower-pocket-frames")] +def private metal_pocket_prompt(m : PocketModel; var vs : PocketVoiceState; n_txt : int64; var sc : PocketScratch) : bool { + if (!g_tw_env_tower) { + return tw_decline(MetalTowerDecline.knob) + } + let tf & = unsafe(m.backbone) + let d = tf.d + var lin_mm_ok = tf.ffn % 64l == 0l && d % 64l == 0l + for (l in tf.layers) { + lin_mm_ok = lin_mm_ok && l.in_proj.nin % 32l == 0l && l.out_proj.nin % 32l == 0l && l.ffn1.nin % 32l == 0l && l.ffn2.nin % 32l == 0l + } + if (!tts_pk_prompt_admit(m, vs, n_txt, PK_ATTN_MAX_KEYS) || !pk_frames_ok(m) || !lin_mm_ok) { + return tw_decline(MetalTowerDecline.shape) + } + var tc = ref_time_ticks() + if (!metal_tower_init() || !pk_frames_attach(m) || !pk_voice_attach(vs, tf)) { + return tw_decline(MetalTowerDecline.device) + } + asr_prof_add("tts.pocket.gpu.slab", tc) + tc = ref_time_ticks() + residency_flush() + var sl & = g_tw_pk_frames + var c = St2GpuCtx() + st2_ctx_make(c, sl.buf, 1l, 1l, 0ul) + c.exact = true + var bx = st2_rows_buf(c, n_txt, d) + st2_upload(bx, sc.rows, n_txt * d) + var b = PkTfBufs(bxb = st2_rows_buf(c, n_txt, d), bqkv = st2_rows_buf(c, n_txt, 3l * d), bctx = st2_rows_buf(c, n_txt, d), bh = st2_rows_buf(c, n_txt, tf.ffn)) + var bxh : MetalBuffer? = null + if (sl.blob != null && pf_q8_mm_half()) { + bxh = st2_rows_buf(c, n_txt, (max(d, tf.ffn) + 1l) / 2l) + } + asr_prof_add("tts.pocket.gpu.upload", tc) + tc = ref_time_ticks() + var err = "" + let f = PkRowsForm(pos0 = vs.len, kv_row0 = vs.len, bcos = g_tw_pk_voice.bcos, bsin = g_tw_pk_voice.bsin, kv_only_last = true) + let ran = with_compute_encoder(g_queue, err) $(enc : MetalComputeEncoder?) { // nolint:PERF026,LINT019 — the error-text fetch runs only on the failure leg + pk_rows_tf(enc, c, tf, sl.layers, b, bx, bxh, n_txt, f) $(li : int) { + return (bk = g_tw_pk_voice.bk[li], bv = g_tw_pk_voice.bv[li]) + } + } + asr_prof_add("tts.pocket.gpu.prompt", tc) + if (ran) { + pk_prompt_readback(vs, n_txt, d) + tts_pk_prompt_dev_set(g_tw_pk_prompt_dev, vs, g_tw_pk_voice.gen, sc.rows, n_txt, d) + g_tower_encodes++ + g_tower_rows += n_txt + } + st2_ctx_release(c) + if (!ran) { + return tw_decline(MetalTowerDecline.gpu_error) + } + return true +} + +def private metal_pocket_driver : PocketGpuDriver => PocketGpuDriver(codec = @@metal_pocket_codec, frames = @@metal_pocket_frames, prompt = @@metal_pocket_prompt) def private metal_styletts2_driver : St2GpuDriver { return St2GpuDriver(albert = @@metal_styletts2_albert, text = @@metal_styletts2_text, dur_enc = @@metal_styletts2_dur_enc, diff --git a/modules/dasLLAMA/dasllama/dasllama_pocket.das b/modules/dasLLAMA/dasllama/dasllama_pocket.das index e044ebe2a9..3cdfc263e0 100644 --- a/modules/dasLLAMA/dasllama/dasllama_pocket.das +++ b/modules/dasLLAMA/dasllama/dasllama_pocket.das @@ -695,16 +695,20 @@ def private proj(l : TtsLinear; x : array; t : int64; @scratch @exact_siz } } +//! `kv_only` ends the layer at its cache append - the last layer of a prompt, whose residual nothing reads def private layer_attention(l : PocketLayer; tf : PocketTransformer; var x : array; t, pos0 : int64; - var kc : TtsKvCache; var sc : PocketScratch) { + var kc : TtsKvCache; var sc : PocketScratch; kv_only : bool = false) { let d = tf.d let dh = d / tf.heads layernorm_rows_into(x, sc.xb, t, d, l.norm1, tf.eps) proj(l.in_proj, sc.xb, t, sc.qkv) split_qkv(sc.qkv, t, d, sc.q, sc.k, sc.v) - rope_rows(sc.q, t, d, dh, pos0, tf.period) + if (!kv_only) { + rope_rows(sc.q, t, d, dh, pos0, tf.period) + } rope_rows(sc.k, t, d, dh, pos0, tf.period) kv_cache_append(kc, sc.k, sc.v, t, d) + return if (kv_only) attention_causal_rows(sc.q, t, d, tf.context, kc, sc.ctx) proj(l.out_proj, sc.ctx, t, sc.xb) if (!empty(l.scale1.a)) { @@ -726,10 +730,15 @@ def private layer_ffn(l : PocketLayer; tf : PocketTransformer; var x : array; t, pos0 : int64; var caches : array; var sc : PocketScratch) { +//! values appended to its cache; x carries the residual out - unless `kv_only_last`, which ends +//! the last layer at its cache append (a prompt, whose residual nothing reads). +def transformer_rows(tf : PocketTransformer; var x : array; t, pos0 : int64; var caches : array; var sc : PocketScratch; + kv_only_last : bool = false) { + let last = length(tf.layers) - 1 for (i, l in iter_range(tf.layers), tf.layers) { - layer_attention(l, tf, x, t, pos0, caches[i], sc) + let kv_only = kv_only_last && i == last + layer_attention(l, tf, x, t, pos0, caches[i], sc, kv_only) + continue if (kv_only) layer_ffn(l, tf, x, t, sc) } } @@ -827,10 +836,16 @@ typedef PocketCodecGpuFn = function<(m : PocketModel; latents : array; fr typedef PocketFramesGpuFn = function<(m : PocketModel; var vs : PocketVoiceState; n_txt : int64; std : float; frames_after_eos : int; captured, forced : array; max_frames : int64; var sc : PocketScratch) : tuple> +//! The text prompt's rows through the backbone after the voice's: `prompt_rows`' contract - the +//! embedding rows of `sc.rows` (`pocket_prompt_embed`) at positions vs.len.., every layer's cache +//! appended on the host; false = declined, the caches untouched. +typedef PocketPromptGpuFn = function<(m : PocketModel; var vs : PocketVoiceState; n_txt : int64; var sc : PocketScratch) : bool> + //! A driver's seats, one per stage; a seat left empty keeps the CPU form for that stage. struct PocketGpuDriver { codec : PocketCodecGpuFn frames : PocketFramesGpuFn + prompt : PocketPromptGpuFn } var private g_pk_drv : PocketGpuDriver @@ -838,7 +853,8 @@ var private g_pk_drv_prev : PocketGpuDriver //! the registration a caller's se let private PK_SEAT_CODEC = 0 let private PK_SEAT_FRAMES = 1 -var private g_pk_seats <- tts_seats(["codec", "frames"]) +let private PK_SEAT_PROMPT = 2 +var private g_pk_seats <- tts_seats(["codec", "frames", "prompt"]) [arch(at = "../ARCHITECTURE_MEDIA.md#tower-gpu-hook")] def register_pocket_gpu(drv : PocketGpuDriver) { @@ -1005,7 +1021,7 @@ def voice_state_from_latents(m : PocketModel; lat : array; frames : int64 } vs.len = frames + bos caches_init(vs.caches, m.backbone, vs.len + chunk_rows(int64(POCKET_MAX_TOKENS))) - transformer_rows(m.backbone, rows, vs.len, 0l, vs.caches, sc) + transformer_rows(m.backbone, rows, vs.len, 0l, vs.caches, sc, kv_only_last = true) return <- vs } @@ -1417,13 +1433,33 @@ def pocket_max_frames(n_tokens : int64) : int64 { return int64(ceil((float(n_tokens) / MAX_GEN_TOKENS_PER_SECOND + GEN_SECONDS_PADDING) * FRAME_RATE)) } +//! The text prompt's embedding rows [n_txt][d] into `sc.rows` - the backbone's input on either rail. +def pocket_prompt_embed(m : PocketModel; ids : array; var sc : PocketScratch) { + let d = m.backbone.d + unsafe(scratch_resize(sc.rows, long_length(ids) * d)) + gather_rows(sc.rows, m.text_emb.a, ids, d) +} + //! The text prompt's rows through the backbone after the voice's, every layer's cache filled. def private prompt_rows(m : PocketModel; ids : array; var vs : PocketVoiceState; var sc : PocketScratch) { - let d = m.backbone.d let n_txt = long_length(ids) - unsafe(scratch_resize(sc.rows, n_txt * d)) - gather_rows(sc.rows, m.text_emb.a, ids, d) - transformer_rows(m.backbone, sc.rows, n_txt, vs.len, vs.caches, sc) + pocket_prompt_embed(m, ids, sc) + if (g_pk_drv.prompt != null) { + tts_seat_call(g_pk_seats, PK_SEAT_PROMPT) + if (invoke(g_pk_drv.prompt, m, vs, n_txt, sc)) { + tts_seat_served(g_pk_seats, PK_SEAT_PROMPT) + return + } + } + transformer_rows(m.backbone, sc.rows, n_txt, vs.len, vs.caches, sc, kv_only_last = true) +} + +//! The prompt's rows into the voice's caches alone, on whichever rail serves it - `pocket_synthesize`'s opening, the +//! caches reset to the voice's rows first. +[cold_path] +def pocket_prompt_rows(m : PocketModel; ids : array; var vs : PocketVoiceState; var sc : PocketScratch) { + caches_ready(vs, long_length(ids)) + prompt_rows(m, ids, vs, sc) } //! One prepared chunk in one voice: the text prompt after the voice's rows, frames until EOS plus diff --git a/modules/dasLLAMA/dasllama/dasllama_tts_blocks.das b/modules/dasLLAMA/dasllama/dasllama_tts_blocks.das index e697703662..826fa46f25 100644 --- a/modules/dasLLAMA/dasllama/dasllama_tts_blocks.das +++ b/modules/dasLLAMA/dasllama/dasllama_tts_blocks.das @@ -1360,6 +1360,13 @@ def kv_cache_init(var kc : TtsKvCache; heads, dh, cap : int64) { //! Append the `t` token-major rows of k and v [t][hidden] at positions kc.n .. kc.n + t - 1. def kv_cache_append(var kc : TtsKvCache; k, v : array; t, hidden : int64) { + unsafe { + kv_cache_append_rows(kc, addr(k[0]), addr(v[0]), t, hidden) + } +} + +//! `kv_cache_append` over the rows at `ks` and `vs` - a device readback's staging +def kv_cache_append_rows(var kc : TtsKvCache; ks, vs : float const?; t, hidden : int64) { let hh = kc.heads let dh = kc.dh let cap = kc.cap @@ -1370,8 +1377,6 @@ def kv_cache_append(var kc : TtsKvCache; k, v : array; t, hidden : int64) unsafe { var kp = addr(kc.k[0]) var vp = addr(kc.v[0]) - let ks = addr(k[0]) - let vs = addr(v[0]) for (r in range64(t)) { for (h in range64(hh)) { for (j in range64(dh)) { diff --git a/modules/dasLLAMA/dasllama/dasllama_tts_slab.das b/modules/dasLLAMA/dasllama/dasllama_tts_slab.das index 3451e0a53b..feb28a810d 100644 --- a/modules/dasLLAMA/dasllama/dasllama_tts_slab.das +++ b/modules/dasLLAMA/dasllama/dasllama_tts_slab.das @@ -979,10 +979,53 @@ def tts_pk_frames_ok(m : PocketModel; hs, max_nin : int64; lin_ok : block<(l : T return invoke(lin_ok, h.final_adaln) && h.final_adaln.nin == n && h.final_adaln.nout == 2l * n && invoke(lin_ok, h.final_linear) && h.final_linear.nin == n && h.final_linear.nout == ld } -//! the voice's caches hold this frames call: a cache a layer, voice + text + frames within capacity, capacity within `max_keys` -def tts_pk_frames_admit(m : PocketModel; vs : PocketVoiceState; n_txt, max_frames, max_keys : int64) : bool { - return (max_frames >= 1l && n_txt >= 0l && !empty(vs.caches) && length(vs.caches) == length(m.backbone.layers) - && vs.len + n_txt + max_frames <= vs.caches[0].cap && vs.caches[0].cap <= max_keys) +//! the voice's caches hold `rows` more rows past the voice's: a cache a layer, the rows within capacity, capacity within `max_keys` +def tts_pk_caches_admit(m : PocketModel; vs : PocketVoiceState; rows, max_keys : int64) : bool { + return (!empty(vs.caches) && length(vs.caches) == length(m.backbone.layers) && vs.len + rows <= vs.caches[0].cap && vs.caches[0].cap <= max_keys) +} + +//! the frames call's admission: at least one frame, the text and the frames within the caches +def tts_pk_frames_admit(m : PocketModel; vs : PocketVoiceState; n_txt, max_frames, max_keys : int64) : bool + => max_frames >= 1l && n_txt >= 0l && tts_pk_caches_admit(m, vs, n_txt + max_frames, max_keys) + +//! the prompt call's admission: one row or more, the text within the caches +def tts_pk_prompt_admit(m : PocketModel; vs : PocketVoiceState; n_txt, max_keys : int64) : bool + => n_txt >= 1l && tts_pk_caches_admit(m, vs, n_txt, max_keys) + +//! the prompt rows a served prompt seat left in the voice slot for the frames call after it: the voice they belong +//! to, the slot build they sit in (`gen`, counted up at every rebuild), the rows [len, len + n) and a sample of the +//! embedding rows they came from; a frames call over another voice, slot, fill or prompt uploads every row itself +struct TtsPkPromptDev { + key : uint64 + gen : uint64 + len : int64 + n : int64 + emb : uint64 +} + +//! a sample of the prompt's embedding rows [n][d] - the first and last row's leading floats and the whole's length - the +//! key two prompts of one length part on +def tts_pk_prompt_emb_key(rows : array; n, d : int64) : uint64 { + var key = 31ul + uint64(n) * 11ul + let last = max(n - 1l, 0l) * d + for (i in range64(min(d, 32l))) { + key = hash_combine64(key, uint64(unsafe(reinterpret(rows[i])))) + key = hash_combine64(key, uint64(unsafe(reinterpret(rows[last + i])))) + } + return key +} + +//! the rows [vs.len, vs.len + n_txt) the record holds on the device for this frames call, 0 where it holds another voice, +//! slot build or prompt; the record is spent either way, so a frames call that never came cannot leave it to the next +def tts_pk_prompt_dev_take(var rec : TtsPkPromptDev; vs : PocketVoiceState; gen : uint64; rows : array; n_txt, d : int64) : int64 { + let n = (rec.key == tts_pk_voice_key(vs) && rec.gen == gen && rec.len == vs.len && rec.n == n_txt + && rec.emb == tts_pk_prompt_emb_key(rows, n_txt, d)) ? rec.n : 0l + rec = TtsPkPromptDev() + return n +} + +def tts_pk_prompt_dev_set(var rec : TtsPkPromptDev; vs : PocketVoiceState; gen : uint64; rows : array; n_txt, d : int64) { + rec = TtsPkPromptDev(key = tts_pk_voice_key(vs), gen = gen, len = vs.len, n = n_txt, emb = tts_pk_prompt_emb_key(rows, n_txt, d)) } //! the conditioning rows a frames call reads back after the loop stepped `stepped` frames: every frame made and the breaking frame's row, as the CPU loop leaves them diff --git a/modules/dasLLAMA/dasllama/dasllama_vulkan_classes.das b/modules/dasLLAMA/dasllama/dasllama_vulkan_classes.das index 7f9ecd1ef6..4a439970a1 100644 --- a/modules/dasLLAMA/dasllama/dasllama_vulkan_classes.das +++ b/modules/dasLLAMA/dasllama/dasllama_vulkan_classes.das @@ -160,7 +160,7 @@ def private q8k_amax5(m0 : float) : float { //! (a stamp off the gate carries no such binding and none of the text): eight consecutive lanes hold a block, `blk_pass` quantizing //! a row's blocks from `normed_at`, the stamp's value for a plane element under the row's statistics [ |> template_struct_instance] -class template Q8BlockStoreT { +class template Q8BlockStoreT : GkWgReduce { @template_constant OUT_Q8 : bool = false //!< the stamp quantizes its rows into Q8_0 blocks: `outq` and the block scales `outs` [arch(at="../ARCHITECTURE_GPU_VULKAN.md#q8-requant-byte-store")] @@ -220,44 +220,24 @@ class template Q8BlockStoreT { } } -//! the workgroup reduces the row kernels share (one wg reduces one row): one barrier a reduce, every -//! thread summing the subgroup partials in the same order, so every thread holds the same sum, max or inverse rms +//! the row kernels' reduces in value form over the shared `GkWgReduce` (one wg reduces one row), and the inverse rms [ |> template_struct_instance] class template WgReduceBase : Q8BlockStoreT { - @workgroup part : float[128] //! two slots of 64 partials, one per subgroup (subgroupSize >= 4 floor) - //! `slot` alternates 0 / 1 between consecutive reduces of one kernel def wg_rms_inv(ss0 : float; n : uint; eps : float; slot : uint) : float => 1.0 / sqrt(wg_sum(ss0, slot) / float(n) + eps) // the CPU rmsnorm's own form - //! the workgroup sum every thread holds alike (the layernorm's mean and variance ride it, and `wg_rms_inv`) + //! the workgroup sum every thread holds alike (the layernorm's mean and variance ride it, and `wg_rms_inv`) - the + //! shared base's reduce in value form, which the SPIR-V emitter takes def wg_sum(v0 : float; slot : uint) : float { - let ss = subgroupAdd(v0) - if (gl_SubgroupInvocationID == 0u) { - part[slot * 64u + gl_SubgroupID] = ss - } - barrier() var tot = 0.0 - var sgi = 0u - while (sgi < gl_NumSubgroups) { - tot += part[slot * 64u + sgi] - sgi++ - } + wg_sum_into(v0, slot, tot) return tot } //! the workgroup max every thread holds alike (the chunked attention's score max) def wg_max(v0 : float; slot : uint) : float { - let sm = subgroupMax(v0) - if (gl_SubgroupInvocationID == 0u) { - part[slot * 64u + gl_SubgroupID] = sm - } - barrier() - var m = -3.0e38 - var sgi = 0u - while (sgi < gl_NumSubgroups) { - m = max(m, part[slot * 64u + sgi]) - sgi++ - } + var m = 0.0 + wg_max_into(v0, slot, m) return m } } @@ -10155,99 +10135,16 @@ class TowerGluSig { } } -struct DwConvRowsArgs { - c : uint //! the channels, one tap row each - k : uint //! the taps, where the stamp does not fix them - stride : uint //! the transposed stamp's stride - pad_l : uint //! the transposed stamp's cropped left edge - dil : uint //! the transposed stamp's dilation - t_in : uint //! the input rows, and the output rows of the causal and the centered stamps - t_out : uint //! the transposed stamp's output rows - woff : uint //! w [c][k]'s element base in wn - boff : uint //! the transposed stamp's bias row in wn - bnwoff : uint //! the centered stamp's folded BatchNorm scale row in wn - bnboff : uint //! the centered stamp's folded BatchNorm shift row in wn -} - -//! a depthwise conv on rows, y [t_out][c] from x [t_in][c] and w [c][k], one element an invocation; FORM is the source -//! rule and the epilogue: 0 the transposed conv gathered, the bias row first and one mad a tap (the CPU -//! `conv1d_rows_transposed_depthwise`); 1 the causal conv, a row reading the k - 1 rows before it (gemma4a's conv module); -//! 2 the centered conv, the folded BatchNorm and silu after it (canary's conv module tail) -[ |> template_struct_instance] -class template DwConvRowsT { - @template_constant FORM : int = 0 //!< spelled as its literal: a template constant folds only a literal - @template_constant K : uint = 0u //!< the taps where the stamp fixes them, else 0 and the taps are `pa.k` - @ssbo @binding = 0 x : array - @ssbo @binding = 1 y : array - @ssbo @binding = 2 @role = "weight" wn : array - @push_constant pa : DwConvRowsArgs - - def rows_out : uint { - static_if (FORM == 0) { - return pa.t_out - } else { - return pa.t_in - } - } - - def taps : uint { - static_if (K != 0u) { - return K - } else { - return pa.k - } - } - - def src_row(t, kk : uint) : int { - static_if (FORM == 0) { - return conv_src_tr_row(t, kk, 0u, pa.pad_l, pa.dil, max(pa.stride, 1u), pa.t_in) - } else { - static_if (FORM == 1) { - return conv_src_fwd_row(t, kk, 0u, taps() - 1u, 1u, 1u, pa.t_in) - } else { - return conv_src_fwd_row(t, kk, 0u, (taps() - 1u) / 2u, 1u, 1u, pa.t_in) - } - } - } - - [spirv_kernel(local_size_x = 256)] - def run { - let e = gl_GlobalInvocationID.x - if (e < rows_out() * pa.c) { - let c = max(pa.c, 1u) - let t = e / c - let j = e % c - let nk = taps() - var acc = 0.0 - static_if (FORM == 0) { - acc = wn[pa.boff + j] - } - var kk = 0u - while (kk < nk) { - let src = src_row(t, kk) - if (src >= 0) { - static_if (FORM == 0) { - acc = mad(wn[pa.woff + j * nk + kk], x[uint(src) * pa.c + j], acc) - } else { - acc += wn[pa.woff + j * nk + kk] * x[uint(src) * pa.c + j] - } - } - kk++ - } - static_if (FORM == 2) { - y[e] = silu_f32(acc * wn[pa.bnwoff + j] + wn[pa.bnboff + j]) - } else { - y[e] = acc - } - } - } -} - //! gemma4a's causal depthwise conv, k = 5: the chunk's first rows read nothing before the start [vk_dispatch(name = "tower_dw_conv5_cls", kernel = "run", grid = "wgs", params = "wgs : int64")] -class TowerDwConv5 : DwConvRowsT { +class TowerDwConv5 : GkDwConvRows { override FORM = 1 override K = 5u + + [spirv_kernel(local_size_x = 256)] + def run { + dw_conv_rows_body() + } } struct TowerCnAttnRArgs { @@ -10361,8 +10258,13 @@ class TowerCnAttnR64 : TowerCnAttnRT { //! canary's conv module tail: the centered depthwise conv over time (k taps, odd - canary 9), the folded BatchNorm and //! silu in one pass, the CPU loop's form [vk_dispatch(name = "tower_cn_dw_cls", kernel = "run", grid = "wgs", params = "wgs : int64")] -class TowerCnDw : DwConvRowsT { +class TowerCnDw : GkDwConvRows { override FORM = 2 + + [spirv_kernel(local_size_x = 256)] + def run { + dw_conv_rows_body() + } } struct TowerGegluArgs { @@ -10913,81 +10815,17 @@ class template TowerWinAttnT { [vk_dispatch(name = "tower_win_attn_cls", kernel = "run", grid = "wgs", params = "wgs : int64"), arch(at = "../ARCHITECTURE_GPU_TOWER_VULKAN.md#vk-tower-routes")] class TowerWinAttn : TowerWinAttnT {} -struct RopeTabArgs { - rows : uint //! the rows roped, row r at position r - qs : uint //! x's row stride - col : uint //! the roped span's first column (a fused [q | k | v] row's q or k slot) - d : uint //! the span's width, its heads side by side - dh : uint //! the head size - coff : uint //! the cos table [pos][dh / 2]'s element base in its binding - soff : uint //! the sin table's element base in its binding, laid out as the cos table -} - -//! the rotary embedding over a span of rows in place from per-position tables: row r at position r turns each pair -//! of each head by the tables' angle j at that position - NEOX the pairs (j, j + dh / 2), the CPU `rope_neox_tab_rows` (the -//! vision mrope tables), else the adjacent pairs (2j, 2j + 1), the CPU `rope_rows`; one pair an invocation over rows * d / 2 -[ |> template_struct_instance] -class template RopeTabT { - @template_constant NEOX : bool = false - @template_constant SLAB_TABS : bool = false //!< the tables sit in a weight slab, not in per-encode planes - @ssbo @binding = 0 x : array - @ssbo @binding = 1 @template_gate = "!SLAB_TABS" cos_t : array - @ssbo @binding = 2 @template_gate = "!SLAB_TABS" sin_t : array - @ssbo @binding = 1 @role = "weight" @template_gate = SLAB_TABS cos_w : array - @ssbo @binding = 2 @role = "weight" @template_gate = SLAB_TABS sin_w : array - @push_constant pa : RopeTabArgs - - def pair_first(hb, j : uint) : uint { - static_if (NEOX) { - return hb + j - } else { - return hb + 2u * j - } - } - - def pair_partner_dist(hh : uint) : uint { - static_if (NEOX) { - return hh - } else { - return 1u - } - } +//! the vision towers' full-head NEOX rope over q or k rows, one table row a position shared by every head +[vk_dispatch(name = "tower_rope_tab_cls", kernel = "run", grid = "wgs", params = "wgs : int64")] +class TowerRopeTab : GkRopeTab { + override NEOX = true [spirv_kernel(local_size_x = 256)] def run { - let gid = gl_GlobalInvocationID.x - let half = pa.d / 2u - if (gid < pa.rows * half) { - let r = gid / max(half, 1u) - let pi = gid % max(half, 1u) - let hh = max(pa.dh / 2u, 1u) - let j = pi % hh - let i0 = pair_first(r * pa.qs + pa.col + (pi / hh) * pa.dh, j) - let i1 = i0 + pair_partner_dist(hh) - let tb = r * hh + j - var cv = 0.0 - var sv = 0.0 - static_if (SLAB_TABS) { - cv = cos_w[pa.coff + tb] - sv = sin_w[pa.soff + tb] - } else { - cv = cos_t[pa.coff + tb] - sv = sin_t[pa.soff + tb] - } - let v0 = x[i0] - let v1 = x[i1] - x[i0] = v0 * cv - v1 * sv - x[i1] = v0 * sv + v1 * cv - } + rope_tab_body() } } -//! the vision towers' full-head NEOX rope over q or k rows, one table row a position shared by every head -[vk_dispatch(name = "tower_rope_tab_cls", kernel = "run", grid = "wgs", params = "wgs : int64")] -class TowerRopeTab : RopeTabT { - override NEOX = true -} - struct RopeKvArgs { qd : uint // query width (the kv rows start here in kvsrc) kvd : uint // kv width (v rows start at qd + kvd) @@ -13674,33 +13512,12 @@ class TtsConcat : GkConcat { } } -struct TtsSigSumArgs { - t : uint //! the rows - nb : uint //! the logits summed a row - ls : uint //! the logit rows' stride - speed : float //! the divisor of each row's sum -} - -//! raw[r] = (the sum of sigmoid(logits[r][j]) over j < nb, in j order) / speed - the CPU `styletts2_durations`' raw -//! durations; one row an invocation +//! the duration sigmoid sums - the CPU `styletts2_durations`' raw durations [vk_dispatch(name = "tts_sigsum_cls", kernel = "run", grid = "t/256", params = "t : int64")] -class TtsSigSum { - @ssbo @binding = 0 logits : array - @ssbo @binding = 1 raw : array - @push_constant pa : TtsSigSumArgs - +class TtsSigSum : GkSigSum { [spirv_kernel(local_size_x = 256)] def run { - let r = gl_GlobalInvocationID.x - if (r < pa.t) { - var acc = 0.0 - var j = 0u - while (j < pa.nb) { - acc += sigmoid_f32(logits[r * pa.ls + j]) - j++ - } - raw[r] = acc / pa.speed - } + sigsum_body() } } @@ -13803,88 +13620,41 @@ def tts_stats_blocks(t : int64) : int64 => (t + int64(TTS_STATS_ROWS) - 1l) / in //! the stats rows a seat holds for `c_max` channels over `t_max` rows: the mean and squared-deviation rows, then the partials def tts_stats_floats(t_max, c_max : int64) : int64 => 2l * c_max + tts_stats_blocks(t_max) * c_max -let TTS_ADAIN_EPS = 1e-5 //! the instance norms' eps, the kernel body's literal; the driver holds it to the model's `ST2_ADAIN_EPS` - -struct TtsAdainArgs { - c : uint //! the channels - t : uint //! the rows the statistics ran over - xoff : uint //! x's element base - yoff : uint //! y's element base - hoff : uint //! the style fc rows' base in h: gamma [c], then beta [c] - gwoff : uint //! the norm's scale row in wn - gboff : uint //! the norm's shift row in wn - aoff : uint //! Snake's alpha row in wn, read by the Snake stamp alone -} - -//! AdaIN as the CPU `adain_affine` folds it - the column statistics, the norm's own scale and shift, then (1 + gamma) and -//! beta - followed by Snake or LeakyReLU(0.2); y may be x itself; one element an invocation over t * c -[ |> template_struct_instance] -class template TtsAdainT { - @template_constant SNAKE : bool = false - @ssbo @binding = 0 x : array - @ssbo @binding = 1 y : array - @ssbo @binding = 2 stats : array - @ssbo @binding = 3 h : array - @ssbo @binding = 4 @role = "weight" wn : array - @push_constant pa : TtsAdainArgs - +//! AdaIN then LeakyReLU(0.2) / Snake +[vk_dispatch(name = "tts_adain_leaky_cls", kernel = "run", grid = "total/256", params = "total : int64")] +class TtsAdainLeaky : GkAdain { [spirv_kernel(local_size_x = 256)] def run { - let gid = gl_GlobalInvocationID.x - if (gid < pa.t * pa.c) { - let ci = gid % max(pa.c, 1u) - var scale = 0.0 - var shift = 0.0 - tts_adain_affine(stats[ci], stats[pa.c + ci], float(pa.t), TTS_ADAIN_EPS, wn[pa.gwoff + ci], wn[pa.gboff + ci], h[pa.hoff + ci], - h[pa.hoff + pa.c + ci], scale, shift) - let v = mad(x[pa.xoff + gid], scale, shift) - static_if (SNAKE) { - let a = wn[pa.aoff + ci] - let sn = sin(a * v) - y[pa.yoff + gid] = mad(1.0 / a, sn * sn, v) - } else { - y[pa.yoff + gid] = v > 0.0 ? v : v * 0.2 - } - } + adain_body() } } -[vk_dispatch(name = "tts_adain_leaky_cls", kernel = "run", grid = "total/256", params = "total : int64")] -class TtsAdainLeaky : TtsAdainT {} - [vk_dispatch(name = "tts_adain_snake_cls", kernel = "run", grid = "total/256", params = "total : int64")] -class TtsAdainSnake : TtsAdainT { +class TtsAdainSnake : GkAdain { override SNAKE = true + + [spirv_kernel(local_size_x = 256)] + def run { + adain_body() + } } //! the depthwise transposed pool on rows: y [t_out][c] gathers the rows of x [t_in][c] its taps reach - the CPU //! `conv1d_rows_transposed_depthwise` [vk_dispatch(name = "tts_pool_dw_cls", kernel = "run", grid = "total/256", params = "total : int64")] -class TtsPoolDw : DwConvRowsT {} - -struct TtsAddScaleArgs { - nelem : uint - scale : float - aoff : uint //! each plane's element base - boff : uint - ooff : uint +class TtsPoolDw : GkDwConvRows { + [spirv_kernel(local_size_x = 256)] + def run { + dw_conv_rows_body() + } } -//! o = (a + b) * scale over nelem elements - the residual join's 1 / sqrt(2) and the averaged stage sums; o may be a; -//! one element an invocation +//! the residual join's add-and-scale, in place where o is a [vk_dispatch(name = "tts_add_scale_cls", kernel = "run", grid = "nelem/256", params = "nelem : int64")] -class TtsAddScale { - @ssbo @binding = 0 a : array - @ssbo @binding = 1 b : array - @ssbo @binding = 2 o : array - @push_constant pa : TtsAddScaleArgs - +class TtsAddScale : GkAddScale { [spirv_kernel(local_size_x = 256)] def run { - let i = gl_GlobalInvocationID.x - if (i < pa.nelem) { - o[pa.ooff + i] = (a[pa.aoff + i] + b[pa.boff + i]) * pa.scale - } + add_scale_body() } } @@ -13897,338 +13667,116 @@ class TtsAxpy : GkAxpy { } } -struct TtsSourceArgs { - n : uint //! the samples, frames * up - nlow : uint //! the phase frames, `resize_len(n, 1 / up)` - nh : uint //! the harmonics - up : uint - sr : float - sine_amp : float - noise_std : float - voiced_thr : float - seed : uint //! the own noise draw's key - woff : uint //! the source linear's weight row [nh] in wn - boff : uint //! its bias [1] in wn -} - -//! a / b rounded as the CPU's IEEE divide: a SPIR-V divide sits up to 2.5 ulp off, which the phase sums carry (Metal's `/` is IEEE) -def private tts_div(a, b : float) : float { - var r = 1.0 / b - r = mad(mad(-b, r, 1.0), r, r) - let q = a * r - return mad(mad(-q, b, a), r, q) -} - -//! the source's resample law - TORCH: torch's taps and mix, else onnxruntime's - the base of the source stamps; the taps' -//! quotients divide through `tts_div` -[ |> template_struct_instance] -class template TtsSrcLawT { - @template_constant TORCH : bool = false - - def law_taps(n, to : uint; scale : float; var i0, i1 : uint&; var l0, l1 : float&) { - static_if (TORCH) { - tts_resize_taps_torch(n, to, tts_div(1.0, scale), i0, i1, l0, l1) - } else { - tts_resize_taps_onnx(n, tts_div(float(to) + 0.5, scale), i0, i1, l0, l1) - } - } +//! the low-rate phase increments on either resample law +[vk_dispatch(name = "tts_src_low_torch_cls", kernel = "run", grid = "total/256", params = "total : int64")] +class TtsSrcLowTorch : GkSrcLow { + override TORCH = true - def law_mix(x0, x1, l0, l1 : float; last : bool) : float { - static_if (TORCH) { - return tts_resize_mix_torch(x0, x1, l0, l1) - } else { - return tts_resize_mix_onnx(x0, x1, l0, l1, last) - } + [spirv_kernel(local_size_x = 256)] + def run { + src_low_body() } } -//! the low-rate phase increments low [nh][nlow] on the stamp's resample law, the first sample's initial phases folded in; -//! one element an invocation -[ |> template_struct_instance] -class template TtsSrcLowT : TtsSrcLawT { - @ssbo @binding = 0 f0 : array - @ssbo @binding = 1 noise_u : array - @ssbo @binding = 2 low : array - @push_constant pa : TtsSourceArgs - +[vk_dispatch(name = "tts_src_low_onnx_cls", kernel = "run", grid = "total/256", params = "total : int64")] +class TtsSrcLowOnnx : GkSrcLow { [spirv_kernel(local_size_x = 256)] def run { - let gid = gl_GlobalInvocationID.x - if (gid < pa.nh * pa.nlow) { - let nlow = max(pa.nlow, 1u) - let up = max(pa.up, 1u) - let h = gid / nlow - var i0 = 0u - var i1 = 0u - var l0 = 0.0 - var l1 = 0.0 - law_taps(pa.n, gid % nlow, tts_div(1.0, float(pa.up)), i0, i1, l0, l1) - let r0 = tts_src_rad(tts_div(f0[i0 / up] * float(h + 1u), pa.sr), noise_u[h], i0, h) - let r1 = tts_src_rad(tts_div(f0[i1 / up] * float(h + 1u), pa.sr), noise_u[h], i1, h) - low[gid] = law_mix(r0, r1, l0, l1, h + 1u == pa.nh) - } + src_low_body() } } -[vk_dispatch(name = "tts_src_low_torch_cls", kernel = "run", grid = "total/256", params = "total : int64")] -class TtsSrcLowTorch : TtsSrcLowT { +//! the phase per frame on the torch law, wgs = 1; `precise` keeps the two-float sum's compensation +[vk_dispatch(name = "tts_src_cumsum_torch_cls", kernel = "run", grid = "wgs", params = "wgs : int64")] +class TtsSrcCumsumTorch : GkSrcCumsum { override TORCH = true -} - -[vk_dispatch(name = "tts_src_low_onnx_cls", kernel = "run", grid = "total/256", params = "total : int64")] -class TtsSrcLowOnnx : TtsSrcLowT {} - -//! the phase per frame lowc [nh][nlow] as the running sum of the increments - TORCH: the double accumulator as a two-float -//! sum carrying its rounding error, else the plain float sum; one workgroup, invocation h < nh walks its harmonic, wgs = 1 -[ |> template_struct_instance] -class template TtsSrcCumsumT : TtsSrcLawT { - @ssbo @binding = 0 low : array - @ssbo @binding = 1 lowc : array - @push_constant pa : TtsSourceArgs [spirv_kernel(local_size_x = 256, precise = true)] def run { - let h = gl_LocalInvocationID.x - if (h < pa.nh) { - var acc = 0.0 - var lo = 0.0 - var k = 0u - while (k < pa.nlow) { - let x = low[h * pa.nlow + k] - static_if (TORCH) { - tts_two_sum_add(acc, lo, x) - } else { - acc += x - } - lowc[h * pa.nlow + k] = tts_phase_rad(acc, pa.up) - k++ - } - } + src_cumsum_body() } } -[vk_dispatch(name = "tts_src_cumsum_torch_cls", kernel = "run", grid = "wgs", params = "wgs : int64")] -class TtsSrcCumsumTorch : TtsSrcCumsumT { - override TORCH = true -} - [vk_dispatch(name = "tts_src_cumsum_onnx_cls", kernel = "run", grid = "wgs", params = "wgs : int64")] -class TtsSrcCumsumOnnx : TtsSrcCumsumT {} - -//! the source's own noise draw noise_n [n][nh], `tts_hash_normal` at the seed - the captured stream's seat when no stream -//! was captured; one element an invocation -[vk_dispatch(name = "tts_src_noise_cls", kernel = "run", grid = "total/256", params = "total : int64")] -class TtsSrcNoise { - @ssbo @binding = 0 noise_n : array - @push_constant pa : TtsSourceArgs - - [spirv_kernel(local_size_x = 256)] +class TtsSrcCumsumOnnx : GkSrcCumsum { + [spirv_kernel(local_size_x = 256, precise = true)] def run { - let gid = gl_GlobalInvocationID.x - if (gid < pa.n * pa.nh) { - noise_n[gid] = tts_hash_normal(pa.seed, gid) - } + src_cumsum_body() } } -//! the mixed source signal har [n] on the stamp's resample law: the harmonics' sines at the interpolated phases, the noise -//! rows scaled by voicing, the source linear (lw [nh] at woff, lb [1] at boff) and its tanh; one sample an invocation -[ |> template_struct_instance] -class template TtsSrcSinesT : TtsSrcLawT { - @ssbo @binding = 0 f0 : array - @ssbo @binding = 1 lowc : array - @ssbo @binding = 2 noise_n : array - @ssbo @binding = 3 @role = "weight" wn : array - @ssbo @binding = 4 har : array - @push_constant pa : TtsSourceArgs - +//! the source's own noise draw +[vk_dispatch(name = "tts_src_noise_cls", kernel = "run", grid = "total/256", params = "total : int64")] +class TtsSrcNoise : GkSrcNoise { [spirv_kernel(local_size_x = 256)] def run { - let i = gl_GlobalInvocationID.x - if (i < pa.n) { - var i0 = 0u - var i1 = 0u - var l0 = 0.0 - var l1 = 0.0 - law_taps(pa.nlow, i, float(pa.up), i0, i1, l0, l1) - let uv = f0[i / max(pa.up, 1u)] > pa.voiced_thr ? 1.0 : 0.0 - let noise_amp = tts_noise_amp(uv, pa.noise_std, pa.sine_amp) - var acc = wn[pa.boff] - var h = 0u - while (h < pa.nh) { - let phase = law_mix(lowc[h * pa.nlow + i0], lowc[h * pa.nlow + i1], l0, l1, h + 1u == pa.nh) - acc += tts_src_term(wn[pa.woff + h], tts_sin_big(phase) * pa.sine_amp, uv, noise_amp, noise_n[i * pa.nh + h]) - h++ - } - har[i] = tts_tanh(acc) - } + src_noise_body() } } +//! the mixed source signal on either resample law [vk_dispatch(name = "tts_src_sines_torch_cls", kernel = "run", grid = "n/256", params = "n : int64")] -class TtsSrcSinesTorch : TtsSrcSinesT { +class TtsSrcSinesTorch : GkSrcSines { override TORCH = true -} -[vk_dispatch(name = "tts_src_sines_onnx_cls", kernel = "run", grid = "n/256", params = "n : int64")] -class TtsSrcSinesOnnx : TtsSrcSinesT {} - -struct TtsStftArgs { - n : uint //! the signal's samples - frames : uint - bins : uint - k : uint //! the frame length, each basis row's taps - hop : uint - pad : uint //! the samples padded on each side - eps : float - reoff : uint //! the real basis wre [bins][k] in wn - imoff : uint //! the imaginary basis, the same layout + [spirv_kernel(local_size_x = 256)] + def run { + src_sines_body() + } } -//! the STFT as two convs over the padded signal - REFLECT mirrors the signal past its ends, else the edge sample repeats - -//! into spec [frames][2 bins], the magnitudes sqrt(re^2 + im^2 + eps) then the phases, a zero imaginary part on the -//! negative real axis reading +pi (the CPU `magnitude_phase`); one (frame, bin) an invocation -[ |> template_struct_instance] -class template TtsStftT { - @template_constant REFLECT : bool = false - @ssbo @binding = 0 har : array - @ssbo @binding = 1 @role = "weight" wn : array - @ssbo @binding = 2 spec : array - @push_constant pa : TtsStftArgs - +[vk_dispatch(name = "tts_src_sines_onnx_cls", kernel = "run", grid = "n/256", params = "n : int64")] +class TtsSrcSinesOnnx : GkSrcSines { [spirv_kernel(local_size_x = 256)] def run { - let gid = gl_GlobalInvocationID.x - if (gid < pa.frames * pa.bins) { - let f = gid / pa.bins - let b = gid % pa.bins - var re = 0.0 - var im = 0.0 - var tap = 0u - while (tap < pa.k) { - let p = int(f * pa.hop + tap) - int(pa.pad) - var src = 0 - static_if (REFLECT) { - src = tts_pad_reflect(p, int(pa.n)) - } else { - src = clamp(p, 0, int(pa.n) - 1) - } - let v = har[uint(src)] - re += wn[pa.reoff + b * pa.k + tap] * v - im += wn[pa.imoff + b * pa.k + tap] * v - tap++ - } - let row = f * 2u * pa.bins - spec[row + b] = tts_stft_mag(re, im, pa.eps) - spec[row + pa.bins + b] = tts_stft_phase(re, im) - } + src_sines_body() } } +//! the STFT on the reflect pad law [vk_dispatch(name = "tts_stft_reflect_cls", kernel = "run", grid = "total/256", params = "total : int64")] -class TtsStftReflect : TtsStftT { +class TtsStftReflect : GkStft { override REFLECT = true + + [spirv_kernel(local_size_x = 256)] + def run { + stft_body() + } } [vk_dispatch(name = "tts_stft_edge_cls", kernel = "run", grid = "total/256", params = "total : int64")] -class TtsStftEdge : TtsStftT {} - -struct TtsIstftArgs { - tt : uint //! the frames - stride : uint //! y's row width - bins : uint - k : uint //! the frame length - hop : uint - pad : uint //! the samples trimmed from each end - envelope : uint //! 1 = the summed squared window divided out where it is not negligible - n : uint //! the samples out - reoff : uint //! the real basis wre [k][bins] in wn - the transposed conv's operand layout - imoff : uint //! the imaginary basis, the same layout - wnoff : uint //! the window [k] in wn -} - -//! the inverse STFT from conv_post's rows y [tt][stride] - a bin's log magnitude, then its phase - as the overlap-added -//! transposed conv, a non-negligible envelope divided out under `envelope`; one sample an invocation -[vk_dispatch(name = "tts_istft_cls", kernel = "run", grid = "n/256", params = "n : int64")] -class TtsIstft { - @ssbo @binding = 0 y : array - @ssbo @binding = 1 @role = "weight" wn : array - @ssbo @binding = 2 wave : array - @push_constant pa : TtsIstftArgs - +class TtsStftEdge : GkStft { [spirv_kernel(local_size_x = 256)] def run { - let i = gl_GlobalInvocationID.x - if (i < pa.n) { - let m = i + pa.pad - let hop = max(pa.hop, 1u) - var acc = 0.0 - var env = 0.0 - var tap = 0u - while (tap < pa.k) { - if (m >= tap && (m - tap) % hop == 0u) { - let q = (m - tap) / hop - if (q < pa.tt) { - var b = 0u - while (b < pa.bins) { - let wa = tap * pa.bins + b - acc += tts_istft_term(y[q * pa.stride + b], y[q * pa.stride + pa.bins + b], wn[pa.reoff + wa], wn[pa.imoff + wa]) - b++ - } - let wv = wn[pa.wnoff + tap] - env += wv * wv - } - } - tap++ - } - wave[i] = (pa.envelope != 0u && env > 1e-11) ? acc / env : acc - } + stft_body() } } -struct TtsPkRowsArgs { - rows : uint - width : uint //! the floats a row copies - ss : uint //! src's row stride - sc : uint //! src's first column - sr0 : uint //! src's first row - ds : uint //! dst's row stride - dc : uint //! dst's first column - dr0 : uint //! dst's first row -} - -//! `rows` rows of `width` floats from src [.][ss] at column sc, rows sr0.., into dst [.][ds] at column dc, rows dr0.. - the -//! codec stream's row copy (a conv's window behind its carry, a carry's tail to its head, the K and V columns into their -//! caches); the ELU stamp maps each value on the way (the CPU `elu_rows` before a conv); one element an invocation -[ |> template_struct_instance] -class template TtsPkRowsT { - @template_constant ELU : bool = false - @ssbo @binding = 0 src : array - @ssbo @binding = 1 dst : array - @push_constant pa : TtsPkRowsArgs - +//! the inverse STFT from conv_post's rows +[vk_dispatch(name = "tts_istft_cls", kernel = "run", grid = "n/256", params = "n : int64")] +class TtsIstft : GkIstft { [spirv_kernel(local_size_x = 256)] def run { - let gid = gl_GlobalInvocationID.x - if (gid < pa.rows * pa.width) { - let w = max(pa.width, 1u) - let r = gid / w - let j = gid % w - var v = src[(pa.sr0 + r) * pa.ss + pa.sc + j] - static_if (ELU) { - v = max(v, 0.0) + min(exp(v) - 1.0, 0.0) - } - dst[(pa.dr0 + r) * pa.ds + pa.dc + j] = v - } + istft_body() } } [vk_dispatch(name = "tts_pk_rows_cls", kernel = "run", grid = "total/256", params = "total : int64")] -class TtsPkRows : TtsPkRowsT {} +class TtsPkRows : GkPkRows { + [spirv_kernel(local_size_x = 256)] + def run { + pk_rows_body() + } +} [vk_dispatch(name = "tts_pk_rows_elu_cls", kernel = "run", grid = "total/256", params = "total : int64")] -class TtsPkRowsElu : TtsPkRowsT { +class TtsPkRowsElu : GkPkRows { override ELU = true + + [spirv_kernel(local_size_x = 256)] + def run { + pk_rows_body() + } } //! x [rows][d] *= the scale row at boff in place - a transformer layer's per-channel layer scale (the CPU @@ -14238,132 +13786,20 @@ class TtsPkRowScale : TowerBiasActT { def override biased(e : uint) : float => x[e] * wn[pa.boff + e % max(1u, pa.d)] } -let TTS_PK_ATTN_HS = 64u //! the head size the attention row is stamped for, the body's literal: another head size declines at the driver -let TTS_PK_ATTN_MAX_KEYS = 2048u //! the scores a row stages, the body's literal: a longer key range declines at the driver +let TTS_PK_ATTN_HS = GK_PK_ATTN_HS +let TTS_PK_ATTN_MAX_KEYS = GK_PK_ATTN_MAX_KEYS -struct TtsPkAttnArgs { - t : uint //! the query rows - qs : uint //! q's row stride - qcol : uint //! the query span's first column in q's rows - d : uint //! the heads' width, the cache rows' and o's row stride - pos0 : uint //! the first query's position - ctx : uint //! the keys a query sees at most, 0 = every key up to itself - qoff : uint //! each plane's element base - koff : uint - voff : uint - ooff : uint -} +//! the subgroup primitives the shared reductions ride, SPIR-V's spelling +def gk_subgroup_add(v : float) : float => subgroupAdd(v) +def gk_subgroup_max(v : float) : float => subgroupMax(v) -//! the causal cached attention of one query row over the caches' rows (the CPU `attention_causal_rows`): query i at -//! position pos0 + i sees the keys up to itself, the last `ctx` of them when ctx is nonzero; the scaled query staged, -//! the scores into the stage, the softmax, then the value sum four parts a column (`value_sum`); one workgroup a -//! (head, query row), wgs = `tts_attn_wgs` +//! the Pocket chains' causal cached attention row - off `WgReduceBase` onto the template both homes stamp (`GkPkAttn`, which +//! derives the same reductions through `GkWgReduce`) [vk_dispatch(name = "tts_pk_attn_cls", kernel = "run", grid = "wgs", params = "wgs : int64")] -class TtsPkAttn : WgReduceBase { - @ssbo @binding = 0 q : array - @ssbo @binding = 1 k : array - @ssbo @binding = 2 v : array - @ssbo @binding = 3 o : array - @push_constant pa : TtsPkAttnArgs - @workgroup qh : float[64] - @workgroup sc : float[int(TTS_PK_ATTN_MAX_KEYS)] - @workgroup po : float[256] - - //! a column's quarter of the value sum: the keys `quarter`, `quarter` + 4, ... in order, eight a step with their loads - //! in flight together, the column at `vb` in the first key's row - def value_sum(vb, quarter, nk : uint) : float { - var acc = 0.0 - var jj = quarter - while (jj + 28u < nk) { - let v0 = v[vb + jj * pa.d] - let v1 = v[vb + (jj + 4u) * pa.d] - let v2 = v[vb + (jj + 8u) * pa.d] - let v3 = v[vb + (jj + 12u) * pa.d] - let v4 = v[vb + (jj + 16u) * pa.d] - let v5 = v[vb + (jj + 20u) * pa.d] - let v6 = v[vb + (jj + 24u) * pa.d] - let v7 = v[vb + (jj + 28u) * pa.d] - acc += sc[jj] * v0 - acc += sc[jj + 4u] * v1 - acc += sc[jj + 8u] * v2 - acc += sc[jj + 12u] * v3 - acc += sc[jj + 16u] * v4 - acc += sc[jj + 20u] * v5 - acc += sc[jj + 24u] * v6 - acc += sc[jj + 28u] * v7 - jj += 32u - } - while (jj < nk) { - acc += sc[jj] * v[vb + jj * pa.d] - jj += 4u - } - return acc - } - +class TtsPkAttn : GkPkAttn { [spirv_kernel(local_size_x = 256)] def run { - let wg = gl_WorkGroupID.x - let t = max(pa.t, 1u) - let h = wg / t - let i = wg % t - let lid = gl_LocalInvocationID.x - let p = pa.pos0 + i - let j0 = tts_window_start(p, pa.ctx) - let nk = p + 1u - j0 - if (lid < 64u) { - qh[lid] = q[pa.qoff + i * pa.qs + pa.qcol + h * 64u + lid] / sqrt(float(64u)) - } - barrier() - var m = -1.0e30 - var j = lid - while (j < nk) { - let krow = pa.koff + (j0 + j) * pa.d + h * 64u - var sdot = 0.0 - var e = 0u - while (e < 64u) { - let k0 = k[krow + e] - let k1 = k[krow + e + 1u] - let k2 = k[krow + e + 2u] - let k3 = k[krow + e + 3u] - sdot += qh[e] * k0 - sdot += qh[e + 1u] * k1 - sdot += qh[e + 2u] * k2 - sdot += qh[e + 3u] * k3 - e += 4u - } - sc[j] = sdot - m = max(m, sdot) - j += 256u - } - let mall = wg_max(m, 0u) - var l = 0.0 - j = lid - while (j < nk) { - let pw = exp(sc[j] - mall) - sc[j] = pw - l += pw - j += 256u - } - let inv = 1.0 / wg_sum(l, 1u) - j = lid - while (j < nk) { - sc[j] = sc[j] * inv - j += 256u - } - barrier() - let col = lid % 64u - let quarter = lid / 64u - po[quarter * 64u + col] = value_sum(pa.voff + j0 * pa.d + h * 64u + col, quarter, nk) - barrier() - if (lid < 64u) { - var tot = 0.0 - var pp = 0u - while (pp < 4u) { - tot += po[pp * 64u + lid] - pp++ - } - o[pa.ooff + i * pa.d + h * 64u + lid] = tot - } + pk_attn_body() } } @@ -14645,8 +14081,11 @@ class TtsPkGemvLnGelu : TtsPkGemvT { override EPI = 8 } -//! the Pocket codec's adjacent-pair rope over the q and k spans, the tables in the slab (cos_w and sin_w bind one buffer) +//! the Pocket chains' adjacent-pair rope over the q and k spans, the tables in a slab or the voice slot (cos_t and sin_t bind one buffer) [vk_dispatch(name = "tts_pk_rope_cls", kernel = "run", grid = "pairs/256", params = "pairs : int64")] -class TtsPkRope : RopeTabT { - override SLAB_TABS = true +class TtsPkRope : GkRopeTab { + [spirv_kernel(local_size_x = 256)] + def run { + rope_tab_body() + } } diff --git a/modules/dasLLAMA/dasllama/dasllama_vulkan_tower.das b/modules/dasLLAMA/dasllama/dasllama_vulkan_tower.das index 94c0231ad4..1969254091 100644 --- a/modules/dasLLAMA/dasllama/dasllama_vulkan_tower.das +++ b/modules/dasLLAMA/dasllama/dasllama_vulkan_tower.das @@ -1235,7 +1235,7 @@ def private vulkan_gemma4a_blocks(t : Gemma4aEncoder; var s : Gemma4aState; npos var pc_rq_rel = RqArgs(inbase = 0u, nblk = uint(rel_cap_rows * d / q8e), lo = -VT_NO_CLAMP, hi = VT_NO_CLAMP) var pc_rel = BatchArgs(n = uint(d), d = uint(d), map_off = 0u) var pc_glu = TowerConvModArgs(d = uint(d), nelem = uint(nd)) - var pc_dw = DwConvRowsArgs(c = uint(d), t_in = uint(npos)) + var pc_dw = GkDwConvRowsArgs(c = uint(d), t_in = uint(npos)) var pc_silu_ff = TowerBiasArgs(d = uint(ff), nelem = uint(nff), boff = 0u, act = BIAS_ACT_SILU) var pc_silu_d = TowerBiasArgs(d = uint(d), nelem = uint(nd), boff = 0u, act = BIAS_ACT_SILU) var pc_attn = TowerG4aAttnArgs(d = uint(d), npos = uint(npos), pdsoff = 0u) @@ -1558,7 +1558,7 @@ def private vulkan_canary_blocks(t : CanaryEncoder; var s : CanaryState; npos, w var pc_rq_xb = RqArgs(inbase = 0u, nblk = uint(nd / q8e), lo = -VT_NO_CLAMP, hi = VT_NO_CLAMP) var pc_rel = BatchArgs(n = uint(d), d = uint(d), map_off = 0u) var pc_glu = TowerConvModArgs(d = uint(d), nelem = uint(nd)) - var pc_dw = DwConvRowsArgs(c = uint(d), k = uint(t.conv_kernel), t_in = uint(npos)) + var pc_dw = GkDwConvRowsArgs(c = uint(d), k = uint(t.conv_kernel), t_in = uint(npos)) var pc_attn_r = TowerCnAttnRArgs(d = uint(d), npos = uint(npos), buoff = 0u, hoff = 0u) var pc_ln = TowerLnArgs(d = uint(d), rows = uint(npos), eps = t.eps, ln_on = 1u) var pc_post = TowerLnArgs(d = uint(d), rows = uint(npos), woff = 0u, boff = 0u, eps = t.eps, ln_on = 1u, lnoff = 0u, lnboff = 0u, ascale = 1.0) @@ -2595,7 +2595,7 @@ def private vulkan_qwen3v_blocks(t : Qwen3vTower; var s : Qwen3vState; npos : in var pc_fa = FaArgs(rows = uint(npos), w0 = 0u, qd = uint(dp), kvd = uint(dp), kv_mul = 1u, kvbase = 0u, scale = 1.0 / sqrt(float(hs)), window = 0u, sinkoff = 0u) var pc_pad_qkv = TowerRestrideArgs(rows = uint(npos), rows_pad = uint(npos_pad), nheads = uint(t.n_head), hs = uint(hs), hsp = uint(hsp), sstride = uint(3l * d), soff = 0u) var pc_unpad = TowerRestrideArgs(rows = uint(npos), rows_pad = uint(npos_pad), nheads = uint(t.n_head), hs = uint(hs), hsp = uint(hsp), sstride = uint(d), soff = 0u) //! the attention rows are d-wide - var pc_rope = RopeTabArgs(rows = uint(npos), qs = uint(3l * d), col = 0u, d = uint(t.n_head * hs), dh = uint(hs), coff = 0u, soff = 0u) + var pc_rope = GkRopeTabArgs(rows = uint(npos), qs = uint(3l * d), col = 0u, d = uint(t.n_head * hs), dh = uint(hs), coff = 0u, soff = 0u) var pc_rq = RqArgs(inbase = 0u, nblk = uint(nd / q8e), lo = -VT_NO_CLAMP, hi = VT_NO_CLAMP) var pc_rqh = RqArgs(inbase = 0u, nblk = uint(nff / q8e), lo = -VT_NO_CLAMP, hi = VT_NO_CLAMP) var pc_gqkv = BatchArgs(n = uint(d), d = uint(3l * d), map_off = 0u) @@ -2740,7 +2740,7 @@ def private vulkan_qwen25v_blocks(t : Qwen25vTower; var s : Qwen25vState; npos : let scale = 1.0 / sqrt(float(hs)) var pc_fa = FaArgs(rows = uint(npos), w0 = 0u, qd = uint(dp), kvd = uint(dp), kv_mul = 1u, kvbase = 0u, scale = scale, window = 0u, sinkoff = 0u) var pc_rs = TowerRestrideArgs(rows = uint(npos), rows_pad = uint(npos_pad), nheads = uint(t.n_head), hs = uint(hs), hsp = uint(hsp), sstride = uint(d), soff = 0u) - var pc_rope = RopeTabArgs(rows = uint(npos), qs = uint(d), col = 0u, d = uint(t.n_head * hs), dh = uint(hs), coff = 0u, soff = 0u) + var pc_rope = GkRopeTabArgs(rows = uint(npos), qs = uint(d), col = 0u, d = uint(t.n_head * hs), dh = uint(hs), coff = 0u, soff = 0u) var pc_gd = F16GemmArgs(n = uint(d), d = uint(d), rows = uint(npos), woff = 0u, ybase = 0u, ystride = uint(d), ksplit = 0u) var pc_gff = F16GemmArgs(n = uint(d), d = uint(ff), rows = uint(npos), woff = 0u, ybase = 0u, ystride = uint(ff), ksplit = 0u) var pc_dn = F16GemmArgs(n = uint(ff), d = uint(d), rows = uint(npos), woff = 0u, ybase = 0u, ystride = uint(d), ksplit = 0u) diff --git a/modules/dasLLAMA/dasllama/dasllama_vulkan_tts.das b/modules/dasLLAMA/dasllama/dasllama_vulkan_tts.das index 494977eb04..b61d97e20e 100644 --- a/modules/dasLLAMA/dasllama/dasllama_vulkan_tts.das +++ b/modules/dasLLAMA/dasllama/dasllama_vulkan_tts.das @@ -21,7 +21,7 @@ require dasllama/dasllama_vulkan_classes require dasllama/dasllama_vulkan_tower //! The Vulkan TTS driver: the kitten and kokoro families' synthesis seats and the Pocket family's -//! codec and frames seats on the tower's knob (`ARCHITECTURE_GPU_TOWER_VULKAN_TTS.md`); the Metal +//! codec, frames and prompt seats on the tower's knob (`ARCHITECTURE_GPU_TOWER_VULKAN_TTS.md`); the Metal //! tower driver fills the same seats on an Apple build. var private g_ts_declines : VkDeclineCounter @@ -319,6 +319,7 @@ def private ts_forget { ts_scratch_release(g_ts_dec_scratch) ts_scratch_release(g_ts_pkc_scratch) ts_scratch_release(g_ts_pkf_scratch) + ts_scratch_release(g_ts_pkp_scratch) ts_pk_voice_drop() } @@ -700,7 +701,7 @@ def private ts_concat(raw : VkCommandBuffer; var h : VkHaz; x, res, f0, nd, o : //! raw[t] = the sigmoid sum of nb logits a row over speed def private ts_sigsum(raw : VkCommandBuffer; var h : VkHaz; logits, out : TsRows; t, nb : int64; speed : float) { var s = set_tts_sigsum_cls(fixed_array(logits.dev, out.dev), fixed_array(logits.bytes, out.bytes), fixed_array(logits.bit, out.bit)) - var pa = TtsSigSumArgs(t = uint(t), nb = uint(nb), ls = uint(nb), speed = speed) + var pa = GkSigSumArgs(t = uint(t), nb = uint(nb), ls = uint(nb), speed = speed) enc_tts_sigsum_cls(raw, h, s, pa, t) vt_pt(raw, VtProfRole.tts_rows) } @@ -726,8 +727,8 @@ def private ts_adain(raw : VkCommandBuffer; var h : VkHaz; slab : TsRows; sc_ : let bufs = fixed_array(x.dev, y.dev, sc_.stats.dev, sc_.aux.dev, slab.dev) let sizes = fixed_array(x.bytes, y.bytes, sc_.stats.bytes, sc_.aux.bytes, slab.bytes) let bits = fixed_array(x.bit, y.bit, sc_.stats.bit, sc_.aux.bit, 0u) - static_assert(TTS_ADAIN_EPS == ST2_ADAIN_EPS, "the AdaIN stamps' eps is the model's") - var p2 = TtsAdainArgs(c = uint(c), t = uint(t), xoff = 0u, yoff = 0u, hoff = uint(hoff), gwoff = uint(ns.gwoff), gboff = uint(ns.gboff), + static_assert(GK_ADAIN_EPS == ST2_ADAIN_EPS, "the AdaIN stamps' eps is the model's") + var p2 = GkAdainArgs(c = uint(c), t = uint(t), hoff = uint(hoff), gwoff = uint(ns.gwoff), gboff = uint(ns.gboff), aoff = uint(ns.aux_off)) if (snake) { var s2 = set_tts_adain_snake_cls(bufs, sizes, bits) @@ -742,7 +743,7 @@ def private ts_adain(raw : VkCommandBuffer; var h : VkHaz; slab : TsRows; sc_ : //! the depthwise transposed pool on rows: y [t_out][c] from x [t_in][c] over the slab's w [c][k] at `woff` and b [c] at `boff` def private ts_pool_dw(raw : VkCommandBuffer; var h : VkHaz; slab : TsRows; cv : TtsConv1d; woff, boff : int64; x, y : TsRows; t_in, t_out : int64) { var s = set_tts_pool_dw_cls(fixed_array(x.dev, y.dev, slab.dev), fixed_array(x.bytes, y.bytes, slab.bytes), fixed_array(x.bit, y.bit, 0u)) - var pa = DwConvRowsArgs(c = uint(cv.cin), k = uint(cv.k), stride = uint(cv.stride), pad_l = uint(cv.pad_l), dil = uint(cv.dilation), t_in = uint(t_in), + var pa = GkDwConvRowsArgs(c = uint(cv.cin), k = uint(cv.k), stride = uint(cv.stride), pad_l = uint(cv.pad_l), dil = uint(cv.dilation), t_in = uint(t_in), t_out = uint(t_out), woff = uint(woff), boff = uint(boff)) enc_tts_pool_dw_cls(raw, h, s, pa, t_out * cv.cin) vt_pt(raw, VtProfRole.tts_rows) @@ -759,7 +760,7 @@ def private ts_gather_rows(raw : VkCommandBuffer; var h : VkHaz; src, idx, out : //! o = (a + b) * scale over n elements (o may be a) def private ts_add_scale(raw : VkCommandBuffer; var h : VkHaz; a, b, o : TsRows; n : int64; scale : float) { var s = set_tts_add_scale_cls(fixed_array(a.dev, b.dev, o.dev), fixed_array(a.bytes, b.bytes, o.bytes), fixed_array(a.bit, b.bit, o.bit)) - var pa = TtsAddScaleArgs(nelem = uint(n), scale = scale, aoff = 0u, boff = 0u, ooff = 0u) + var pa = GkAddScaleArgs(nelem = uint(n), scale = scale) enc_tts_add_scale_cls(raw, h, s, pa, n) vt_pt(raw, VtProfRole.tts_rows) } @@ -1363,9 +1364,9 @@ def private ts_stage(raw : VkCommandBuffer; var h : VkHaz; slab : TsRows; sc_ : let n = t2 * c2 if (last) { var s = set_tts_pk_rows_cls(fixed_array(r_bxt.dev, r_by.dev), fixed_array(r_bxt.bytes, r_by.bytes), fixed_array(r_bxt.bit, r_by.bit)) - var shift = TtsPkRowsArgs(rows = uint(t2 - 1l), width = uint(c2), ss = uint(c2), sc = 0u, sr0 = 0u, ds = uint(c2), dc = 0u, dr0 = 1u) + var shift = GkPkRowsArgs(rows = uint(t2 - 1l), width = uint(c2), ss = uint(c2), sc = 0u, sr0 = 0u, ds = uint(c2), dc = 0u, dr0 = 1u) enc_tts_pk_rows_cls(raw, h, s, shift, n - c2) - var head = TtsPkRowsArgs(rows = 1u, width = uint(c2), ss = uint(c2), sc = 0u, sr0 = 1u, ds = uint(c2), dc = 0u, dr0 = 0u) + var head = GkPkRowsArgs(rows = 1u, width = uint(c2), ss = uint(c2), sc = 0u, sr0 = 1u, ds = uint(c2), dc = 0u, dr0 = 0u) enc_tts_pk_rows_cls(raw, h, s, head, c2) vt_pt(raw, VtProfRole.tts_rows) } @@ -1413,7 +1414,7 @@ def private ts_source_chain(raw : VkCommandBuffer; var h : VkHaz; slab : TsRows; let r_lowc = ts_rows(sc_, TS_D_LOWC) let r_har = ts_rows(sc_, TS_D_HAR) let r_spec = ts_rows(sc_, TS_D_SPEC) - var ka = TtsSourceArgs(n = uint(sh.ss.samples), nlow = uint(sh.ss.nlow), nh = uint(nh), up = uint(cfg.upsample), sr = cfg.sample_rate, sine_amp = cfg.sine_amp, + var ka = GkSourceArgs(n = uint(sh.ss.samples), nlow = uint(sh.ss.nlow), nh = uint(nh), up = uint(cfg.upsample), sr = cfg.sample_rate, sine_amp = cfg.sine_amp, noise_std = cfg.noise_std, voiced_thr = cfg.voiced_thr, seed = uint(seed ^ (seed >> 32ul)), woff = uint(sl.src_w), boff = uint(sl.src_b)) if (!captured) { var s0 = set_tts_src_noise_cls(fixed_array(r_noise.dev), fixed_array(r_noise.bytes), fixed_array(r_noise.bit)) @@ -1446,7 +1447,7 @@ def private ts_source_chain(raw : VkCommandBuffer; var h : VkHaz; slab : TsRows; vt_pt(raw, VtProfRole.tts_source) } let re & = unsafe(d.stft_fwd_re) - var ks = TtsStftArgs(n = uint(sh.ss.samples), frames = uint(sh.ss.spec_frames), bins = uint(d.bins), k = uint(re.k), hop = uint(re.stride), pad = uint(ST2_STFT_PAD), + var ks = GkStftArgs(n = uint(sh.ss.samples), frames = uint(sh.ss.spec_frames), bins = uint(d.bins), k = uint(re.k), hop = uint(re.stride), pad = uint(ST2_STFT_PAD), eps = d.stft_eps, reoff = uint(sl.stft_re), imoff = uint(sl.stft_im)) let st_bufs = fixed_array(r_har.dev, slab.dev, r_spec.dev) let st_sizes = fixed_array(r_har.bytes, slab.bytes, r_spec.bytes) @@ -1566,7 +1567,7 @@ def private vulkan_styletts2_decode(d : St2Decoder; text_c : int64; source : Sin ts_source_chain(raw, h, slab, sc_, d, source, sh, noise.captured, seed) ts_generator_chain(raw, h, slab, sc_, d, sh.dec_width, sh.dec_rows, sh.ss.spec_frames) let r_post = ts_rows(sc_, TS_D_POST) - var ki = TtsIstftArgs(tt = uint(sh.t_post), stride = uint(d.conv_post.cout), bins = uint(d.bins), k = uint(bwd.k), hop = uint(bwd.stride), + var ki = GkIstftArgs(tt = uint(sh.t_post), stride = uint(d.conv_post.cout), bins = uint(d.bins), k = uint(bwd.k), hop = uint(bwd.stride), pad = uint(ST2_STFT_PAD), envelope = d.istft_envelope ? 1u : 0u, n = uint(sh.n_out), reoff = uint(sl.istft_re), imoff = uint(sl.istft_im), wnoff = uint(sl.window)) var s = set_tts_istft_cls(fixed_array(r_post.dev, slab.dev, r_wave.dev), fixed_array(r_post.bytes, slab.bytes, r_wave.bytes), fixed_array(r_post.bit, 0u, r_wave.bit)) @@ -1698,6 +1699,7 @@ struct private TsPkVoice { v : array rope : TsRows sin_off : int64 + gen : uint64 //! counted up at every rebuild: a prompt's residency record names the build it sits in } var private g_ts_pk_voice : TsPkVoice @@ -1737,6 +1739,7 @@ def private ts_pk_voice_attach(vs : PocketVoiceState; tf : PocketTransformer) : def private ts_pk_voice_rebuild(vs : PocketVoiceState; tf : PocketTransformer) : bool { var v & = g_ts_pk_voice ts_pk_voice_drop() + v.gen++ let t0 = ref_time_ticks() let cap = vs.caches[0].cap let d = tf.d @@ -1778,7 +1781,7 @@ def private ts_pk_conv(raw : VkCommandBuffer; var h : VkHaz; slab : TsRows; slot } //! a window of rows copied: `rows` x `width` from src [.][ss] at column sc, rows sr0.., into dst [.][ds] at column dc, rows dr0.., the ELU on the way where asked -def private ts_pk_rows(raw : VkCommandBuffer; var h : VkHaz; src, dst : TsRows; pa : TtsPkRowsArgs; elu : bool) { +def private ts_pk_rows(raw : VkCommandBuffer; var h : VkHaz; src, dst : TsRows; pa : GkPkRowsArgs; elu : bool) { var pc = pa let total = int64(pa.rows) * int64(pa.width) if (elu) { @@ -1791,8 +1794,8 @@ def private ts_pk_rows(raw : VkCommandBuffer; var h : VkHaz; src, dst : TsRows; vt_pt(raw, VtProfRole.tts_pk) } -def private ts_pk_rows_args(rows, width : int64) : TtsPkRowsArgs - => TtsPkRowsArgs(rows = uint(rows), width = uint(width), ss = uint(width), sc = 0u, sr0 = 0u, ds = uint(width), dc = 0u, dr0 = 0u) +def private ts_pk_rows_args(rows, width : int64) : GkPkRowsArgs + => GkPkRowsArgs(rows = uint(rows), width = uint(width), ss = uint(width), sc = 0u, sr0 = 0u, ds = uint(width), dc = 0u, dr0 = 0u) //! the ELU over x's first rows x width elements in place def private ts_pk_elu(raw : VkCommandBuffer; var h : VkHaz; x : TsRows; rows, width : int64) { @@ -1810,7 +1813,7 @@ def private ts_pk_row_scale(raw : VkCommandBuffer; var h : VkHaz; slab : TsRows; //! the adjacent-pair rope over the span `d` wide at column `col` of x's rows [rows][qs], positions 0.., on the tables in `tabs` (cos at coff, sin at soff) def private ts_pk_rope(raw : VkCommandBuffer; var h : VkHaz; tabs : TsRows; coff, soff : int64; x : TsRows; rows, qs, col, d, dh : int64) { var s = set_tts_pk_rope_cls(fixed_array(x.dev, tabs.dev, tabs.dev), fixed_array(x.bytes, tabs.bytes, tabs.bytes), fixed_array(x.bit, tabs.bit, tabs.bit)) - var pa = RopeTabArgs(rows = uint(rows), qs = uint(qs), col = uint(col), d = uint(d), dh = uint(dh), coff = uint(coff), soff = uint(soff)) + var pa = GkRopeTabArgs(rows = uint(rows), qs = uint(qs), col = uint(col), d = uint(d), dh = uint(dh), coff = uint(coff), soff = uint(soff)) enc_tts_pk_rope_cls(raw, h, s, pa, rows * d / 2l) vt_pt(raw, VtProfRole.tts_pk) } @@ -1818,7 +1821,7 @@ def private ts_pk_rope(raw : VkCommandBuffer; var h : VkHaz; tabs : TsRows; coff //! the causal cached attention: q from the qkv rows [t][qs] at column 0, k and v the cache rows [pos][d], o [t][d]; query i at position pos0 + i def private ts_pk_attn(raw : VkCommandBuffer; var h : VkHaz; qkv, k, v, o : TsRows; t, qs, d, heads, pos0, ctx : int64) { var s = set_tts_pk_attn_cls(fixed_array(qkv.dev, k.dev, v.dev, o.dev), fixed_array(qkv.bytes, k.bytes, v.bytes, o.bytes), fixed_array(qkv.bit, k.bit, v.bit, o.bit)) - var pa = TtsPkAttnArgs(t = uint(t), qs = uint(qs), qcol = 0u, d = uint(d), pos0 = uint(pos0), ctx = uint(ctx), qoff = 0u, koff = 0u, voff = 0u, ooff = 0u) + var pa = GkPkAttnArgs(t = uint(t), qs = uint(qs), qcol = 0u, d = uint(d), pos0 = uint(pos0), ctx = uint(ctx)) enc_tts_pk_attn_cls(raw, h, s, pa, tts_attn_wgs(t, heads)) vt_pt(raw, VtProfRole.tts_pk_attn) } @@ -1964,37 +1967,57 @@ def private ts_pk_codec_floats(mm : PocketMimi; frames : int64) : TsFloats { return fl } -//! the codec transformer over the residual rows X [t][d] at positions 0..: the CPU's layer loop, dispatch for dispatch -def private ts_pk_transformer(raw : VkCommandBuffer; var h : VkHaz; slab : TsRows; sc_ : TsScratch; tf : PocketTransformer; sl : TtsPkCodecSlab; t : int64) { +//! a transformer's rows form: every layer's keys and values written at row `kv_row0` of its cache rows, the rope tables' +//! row of position 0 at `coff` / `soff`, the queries at `pos0`; `kv_only_last` ends the last layer at its cache rows +struct private TsPkRowsForm { + pos0 : int64 + kv_row0 : int64 + tabs : TsRows + coff : int64 + soff : int64 + zero : int64 + kv_only_last : bool +} + +//! a transformer over the residual rows X [t][d] of the scratch's slot set: the CPU's layer loop, dispatch for dispatch +def private ts_pk_rows_tf(raw : VkCommandBuffer; var h : VkHaz; slab : TsRows; sc_ : TsScratch; tf : PocketTransformer; layers : array; + t : int64; f : TsPkRowsForm; slot_x, slot_xb, slot_qkv, slot_ctx, slot_h : int; kv : block<(li : int) : tuple>) { let d = tf.d let dh = d / tf.heads let nd = t * d - let r_x = ts_rows(sc_, TS_C_X) - let r_xb = ts_rows(sc_, TS_C_XB) - let r_qkv = ts_rows(sc_, TS_C_QKV) - let r_ctx = ts_rows(sc_, TS_C_CTX) - let r_h = ts_rows(sc_, TS_C_H) - let r_k = ts_rows(sc_, TS_C_K) - let r_v = ts_rows(sc_, TS_C_V) - var kc = TtsPkRowsArgs(rows = uint(t), width = uint(d), ss = uint(3l * d), sc = uint(d), sr0 = 0u, ds = uint(d), dc = 0u, dr0 = 0u) + let half = dh / 2l + let r_x = ts_rows(sc_, slot_x) + let r_xb = ts_rows(sc_, slot_xb) + let r_qkv = ts_rows(sc_, slot_qkv) + let r_ctx = ts_rows(sc_, slot_ctx) + let r_h = ts_rows(sc_, slot_h) + var kc = GkPkRowsArgs(rows = uint(t), width = uint(d), ss = uint(3l * d), sc = uint(d), sr0 = 0u, ds = uint(d), dc = 0u, dr0 = uint(f.kv_row0)) var pa_gelu = TtsElemArgs(nelem = uint(t * tf.ffn), off = 0u, slope = 0.0) var s_gelu = set_tts_gelu_tanh_cls(fixed_array(r_h.dev), fixed_array(r_h.bytes), fixed_array(r_h.bit)) - for (i in range(length(tf.layers))) { - let s & = unsafe(sl.layers[i]) + let coff = f.coff + f.pos0 * half + let soff = f.soff + f.pos0 * half + let last = length(layers) - 1 + for (i in range(length(layers))) { + let s & = unsafe(layers[i]) + let cache = invoke(kv, i) + let kv_only = f.kv_only_last && i == last ts_ln(raw, h, slab, s.n1, r_x, r_xb, t, d, tf.eps) ts_lin(raw, h, slab, s.in_proj, d, 3l * d, r_xb, r_qkv, t) - ts_pk_rope(raw, h, slab, sl.rope_cos, sl.rope_sin, r_qkv, t, 3l * d, 0l, d, dh) - ts_pk_rope(raw, h, slab, sl.rope_cos, sl.rope_sin, r_qkv, t, 3l * d, d, d, dh) + if (!kv_only) { + ts_pk_rope(raw, h, f.tabs, coff, soff, r_qkv, t, 3l * d, 0l, d, dh) + } + ts_pk_rope(raw, h, f.tabs, coff, soff, r_qkv, t, 3l * d, d, d, dh) kc.sc = uint(d) - ts_pk_rows(raw, h, r_qkv, r_k, kc, false) + ts_pk_rows(raw, h, r_qkv, cache.k, kc, false) kc.sc = uint(2l * d) - ts_pk_rows(raw, h, r_qkv, r_v, kc, false) - ts_pk_attn(raw, h, r_qkv, r_k, r_v, r_ctx, t, 3l * d, d, tf.heads, 0l, tf.context) + ts_pk_rows(raw, h, r_qkv, cache.v, kc, false) + continue if (kv_only) + ts_pk_attn(raw, h, r_qkv, cache.k, cache.v, r_ctx, t, 3l * d, d, tf.heads, f.pos0, tf.context) ts_lin(raw, h, slab, s.out_proj, d, d, r_ctx, r_xb, t) if (s.has_s1) { ts_pk_row_scale(raw, h, slab, r_xb, t, d, s.scale1_off) } - ts_pk_add_ln(raw, h, slab, sl.zero, s.n2, r_x, r_xb, r_xb, t, d, tf.eps) + ts_pk_add_ln(raw, h, slab, f.zero, s.n2, r_x, r_xb, r_xb, t, d, tf.eps) ts_lin(raw, h, slab, s.ffn1, d, tf.ffn, r_xb, r_h, t) enc_tts_gelu_tanh_cls(raw, h, s_gelu, pa_gelu, t * tf.ffn) vt_pt(raw, VtProfRole.tts_act) @@ -2006,6 +2029,16 @@ def private ts_pk_transformer(raw : VkCommandBuffer; var h : VkHaz; slab : TsRow } } +//! the codec transformer over the residual rows X [t][d] at positions 0.., its keys and values in the scratch's one row set +def private ts_pk_transformer(raw : VkCommandBuffer; var h : VkHaz; slab : TsRows; sc_ : TsScratch; tf : PocketTransformer; sl : TtsPkCodecSlab; t : int64) { + let r_k = ts_rows(sc_, TS_C_K) + let r_v = ts_rows(sc_, TS_C_V) + let f = TsPkRowsForm(pos0 = 0l, kv_row0 = 0l, tabs = slab, coff = sl.rope_cos, soff = sl.rope_sin, zero = sl.zero, kv_only_last = false) + ts_pk_rows_tf(raw, h, slab, sc_, tf, sl.layers, t, f, TS_C_X, TS_C_XB, TS_C_QKV, TS_C_CTX, TS_C_H) $(_li : int) { + return (k = r_k, v = r_v) + } +} + //! one ELU-conv-ELU-conv residual block over x [t][w] (a stage's rows) into y, through the block's two temporaries def private ts_pk_res_block(raw : VkCommandBuffer; var h : VkHaz; slab : TsRows; sc_ : TsScratch; b : PocketResConv; s1, s2 : TtsConvSlot; x, y : TsRows; t, w : int64) { let r_cc = ts_rows(sc_, TS_C_CC) @@ -2028,7 +2061,7 @@ def private ts_pk_codec_chain(raw : VkCommandBuffer; var h : VkHaz; slab : TsRow let r_x = ts_rows(sc_, TS_C_X) ts_pk_conv(raw, h, slab, sl.quant, mm.quant_proj, frames, ts_rows(sc_, TS_C_LAT), r_col, r_q) var sp = set_tts_pool_dw_cls(fixed_array(r_q.dev, r_x.dev, slab.dev), fixed_array(r_q.bytes, r_x.bytes, slab.bytes), fixed_array(r_q.bit, r_x.bit, 0u)) - var pp = DwConvRowsArgs(c = uint(d), k = uint(mm.upsample.k), stride = uint(mm.upsample.stride), pad_l = 0u, dil = uint(mm.upsample.dilation), t_in = uint(frames), + var pp = GkDwConvRowsArgs(c = uint(d), k = uint(mm.upsample.k), stride = uint(mm.upsample.stride), pad_l = 0u, dil = uint(mm.upsample.dilation), t_in = uint(frames), t_out = uint(t1), woff = uint(sl.up_w), boff = uint(sl.up_b)) enc_tts_pool_dw_cls(raw, h, sp, pp, t1 * d) vt_pt(raw, VtProfRole.tts_rows) @@ -2052,7 +2085,7 @@ def private ts_pk_codec_chain(raw : VkCommandBuffer; var h : VkHaz; slab : TsRow ts_pk_elu(raw, h, cur, t, w) let r_o = ts_rows(sc_, TS_C_O) ts_pk_conv(raw, h, slab, sl.dec_out, mm.dec_out, t, cur, r_col, r_o) - let kw = TtsPkRowsArgs(rows = uint(t), width = 1u, ss = uint(mm.dec_out.cout), sc = 0u, sr0 = 0u, ds = 1u, dc = 0u, dr0 = 0u) + let kw = GkPkRowsArgs(rows = uint(t), width = 1u, ss = uint(mm.dec_out.cout), sc = 0u, sr0 = 0u, ds = 1u, dc = 0u, dr0 = 0u) ts_pk_rows(raw, h, r_o, ts_rows(sc_, TS_C_WAVE), kw, false) } @@ -2203,7 +2236,7 @@ def private ts_pk_frame(raw : VkCommandBuffer; var h : VkHaz; slab : TsRows; sc_ } ts_ln(raw, h, slab, sl.out_norm, r_x, r_xb, 1l, d, m.backbone.eps) let coff = f * d - let kc = TtsPkRowsArgs(rows = 1u, width = uint(d), ss = uint(d), sc = 0u, sr0 = 0u, ds = uint(d), dc = 0u, dr0 = uint(f)) + let kc = GkPkRowsArgs(rows = 1u, width = uint(d), ss = uint(d), sc = 0u, sr0 = 0u, ds = uint(d), dc = 0u, dr0 = uint(f)) ts_pk_rows(raw, h, r_xb, r_conds, kc, false) ts_pk_lin(raw, h, slab, sl.out_eos, d, 1l, r_conds, coff, ts_rows(sc_, TS_F_EOS), f) let noff = i * ld @@ -2268,6 +2301,7 @@ def private ts_pk_frames_loop(m : PocketModel; sl : TtsPkFramesSlab; slab : TsRo def private vulkan_pocket_frames(m : PocketModel; var vs : PocketVoiceState; n_txt : int64; std : float; frames_after_eos : int; captured, forced : array; max_frames : int64; var sc : PocketScratch) : tuple { let none = (frames = -1l, eos_at = -1l) + let dev_rows = tts_pk_prompt_dev_take(g_ts_pk_prompt_dev, vs, g_ts_pk_voice.gen, sc.rows, n_txt, m.backbone.d) //! spent on entry: a decline below must not leave it for a later chunk if (!vulkan_tower_on()) { ts_decline("frames", VulkanTtsDecline.knob) return none @@ -2295,7 +2329,7 @@ def private vulkan_pocket_frames(m : PocketModel; var vs : PocketVoiceState; n_t return none } for (li, kc in iter_range(vs.caches), vs.caches) { - ts_pk_kv_upload(kc, g_ts_pk_voice.k[li], g_ts_pk_voice.v[li], vs.len, base) + ts_pk_kv_upload(kc, g_ts_pk_voice.k[li], g_ts_pk_voice.v[li], vs.len + dev_rows, base) } let sl & = g_ts_pk_frames_slots let sc_ & = g_ts_pkf_scratch @@ -2312,6 +2346,76 @@ def private vulkan_pocket_frames(m : PocketModel; var vs : PocketVoiceState; n_t return r } +var private g_ts_pkp_scratch : TsScratch +var private g_ts_pk_prompt_dev : TtsPkPromptDev + +//! the prompt seat's slots: the rows form's five at the prompt's rows (the frames slots' indices, its own scratch) +def private ts_pk_prompt_floats(m : PocketModel; t : int64) : TsFloats { + let d = m.backbone.d + var fl : TsFloats + fl[TS_F_X] = t * d + fl[TS_F_XB] = t * d + fl[TS_F_QKV] = t * 3l * d + fl[TS_F_CTX] = t * d + fl[TS_F_H] = t * m.backbone.ffn + return fl +} + +//! the prompt's K/V rows [vs.len, vs.len + t) of every layer back from the voice slot into the host caches +def private ts_pk_prompt_readback(var vs : PocketVoiceState; sc_ : TsScratch; t, d : int64) { + let v & = g_ts_pk_voice + g_ts_pk_kv_host |> ensure_length(2l * t * d) + unsafe { + var kp = addr(g_ts_pk_kv_host[0]) + var vp = addr(g_ts_pk_kv_host[t * d]) + for (li, kc in iter_range(vs.caches), vs.caches) { + readback_range(v.k[li].dev, v.k[li].bit, vs.len * d * 4l, sc_.hb, kp, t * d * 4l) + readback_range(v.v[li].dev, v.v[li].bit, vs.len * d * 4l, sc_.hb, vp, t * d * 4l) + kv_cache_append_rows(kc, kp, vp, t, d) + } + } +} + +//! the text prompt's rows through the backbone after the voice's: the frames slab's layers in the rows form at +//! positions vs.len.., the keys and values into the voice slot and back to the host caches, the slot's rows left for +//! the frames call after it +[hot_path, arch(at = "../ARCHITECTURE_GPU_TOWER_VULKAN_TTS.md#vk-pocket-chain")] +def private vulkan_pocket_prompt(m : PocketModel; var vs : PocketVoiceState; n_txt : int64; var sc : PocketScratch) : bool { + if (!vulkan_tower_on()) { + return ts_decline("prompt", VulkanTtsDecline.knob) + } + let tf & = unsafe(m.backbone) + let d = tf.d + let shape_ok = tts_pk_frames_ok(m, TS_PK_ATTN_HS, TS_PK_GEMV_MAX_NIN) $(l : TtsLinear) : bool { + return tts_gemm_lin_ok(l) + } + if (!tts_pk_prompt_admit(m, vs, n_txt, TS_PK_ATTN_MAX_KEYS) || d % 4l != 0l || !shape_ok) { + return ts_decline("prompt", VulkanTtsDecline.shape) + } + if (!ts_arm_ensure("pocket prompt", @@ts_pk_codec_ensure)) { + return ts_decline("prompt", VulkanTtsDecline.device) + } + let fl = ts_pk_prompt_floats(m, n_txt) + if (!ts_pk_frames_attach(m) || !ts_pk_voice_attach(vs, tf) || !ts_scratch_ready(g_ts_pkp_scratch, fl, 0l, 0l, n_txt * d)) { + return ts_decline("prompt", VulkanTtsDecline.memory) + } + let sl & = g_ts_pk_frames_slots + let sc_ & = g_ts_pkp_scratch + let slab = g_ts_pk_frames.rows + let v & = g_ts_pk_voice + ts_upload(ts_rows(sc_, TS_F_X), sc.rows, n_txt * d) + let f = TsPkRowsForm(pos0 = vs.len, kv_row0 = vs.len, tabs = v.rope, coff = 0l, soff = v.sin_off, zero = sl.zero, kv_only_last = true) + vt_chain_one_shot("pocket prompt", long_length(tf.layers), n_txt) $(raw : VkCommandBuffer; var h : VkHaz) { + ts_pk_rows_tf(raw, h, slab, sc_, tf, sl.layers, n_txt, f, TS_F_X, TS_F_XB, TS_F_QKV, TS_F_CTX, TS_F_H) $(li : int) { + return (k = v.k[li], v = v.v[li]) + } + } + g_ts_encodes++ + ts_pk_prompt_readback(vs, sc_, n_txt, d) + tts_pk_prompt_dev_set(g_ts_pk_prompt_dev, vs, v.gen, sc.rows, n_txt, d) + return true +} + [init] def dasllama_vulkan_tts_register { register_vk_drop_hook(@@ts_forget) @@ -2320,6 +2424,6 @@ def dasllama_vulkan_tts_register { register_styletts2_gpu(St2GpuDriver(albert = @@vulkan_styletts2_albert, text = @@vulkan_styletts2_text, dur_enc = @@vulkan_styletts2_dur_enc, durations = @@vulkan_styletts2_durations, prosody = @@vulkan_styletts2_prosody, decode = @@vulkan_styletts2_decode, gen = @@vulkan_styletts2_generator)) - register_pocket_gpu(PocketGpuDriver(codec = @@vulkan_pocket_codec, frames = @@vulkan_pocket_frames)) + register_pocket_gpu(PocketGpuDriver(codec = @@vulkan_pocket_codec, frames = @@vulkan_pocket_frames, prompt = @@vulkan_pocket_prompt)) } } diff --git a/modules/dasLLAMA/followup_metal.md b/modules/dasLLAMA/followup_metal.md index 32e364c39d..22008eb588 100644 --- a/modules/dasLLAMA/followup_metal.md +++ b/modules/dasLLAMA/followup_metal.md @@ -745,15 +745,14 @@ Snake blocks at 512 channels, where a wider N tile or the AdaIN pass folded into loader are the A/Bs. The served-lane synthesis cell (`tests/_tts_parity.das`, `tts_gpu_synthesis`) gates counters and sample counts and logs its sample-wise waveform figure without a bar; a phase-insensitive instrument - a per-window spectral compare tolerant of one -frame of shift - would gate the served lane end to end. Pocket TTS rides the tower in two seats -(`ARCHITECTURE_GPU_TOWER.md#tower-pocket-codec` and `ARCHITECTURE_GPU_TOWER.md#tower-pocket-frames`); what its frame loop still costs is the +frame of shift - would gate the served lane end to end. Pocket TTS rides the tower in three seats +(`ARCHITECTURE_GPU_TOWER_POCKET.md#tower-pocket-codec` and `ARCHITECTURE_GPU_TOWER_POCKET.md#tower-pocket-frames`); what its frame loop still costs is the GPU's own time, `debug-jit` on the M5 Max 0.6 ms a frame on the q8 file against the CPU's 1.3 (`PERF_LEDGER.md`, `harness/pocket_stage_probe.das`), 0.44 of it the backbone's 57 dispatches and 0.19 the head's 22, the encode under 0.03: the levers are the head as one threadgroup over q8 weights (nine million parameters, one dispatch in place of -22 - where the CPU's head is q8 already, the kq files), the first norm folded into the q8 GEMV's -prologue as the f32 route already folds it, and the text prompt's rows on the tower (the -`prompt` stage clock of `test_pocket_synthesis_metal`, six milliseconds a chunk on the CPU). The frames' K/V never return to the host, so a chunk +22 - where the CPU's head is q8 already, the kq files) and the first norm folded into the q8 GEMV's +prologue as the f32 route already folds it. The frames' K/V never return to the host, so a chunk whose command buffer fails reruns whole on the CPU. The codec transformer's layer (`pk_transformer`) and the frame loop's (`pk_fr_layer`) are two bodies of one layer: one body waits on a rows form of the rope-and-store (`MetalRopeStoreBKvT` with a row table is the diff --git a/modules/dasLLAMA/followup_vulkan.md b/modules/dasLLAMA/followup_vulkan.md index 924a62b747..a36df694c9 100644 --- a/modules/dasLLAMA/followup_vulkan.md +++ b/modules/dasLLAMA/followup_vulkan.md @@ -1735,7 +1735,7 @@ module) is independent and can land any time - it is pure structure. the compensated sums they protect can fold. Done = the arm sets the knob where the user has not, and the source kernels' cells run green on the Mac. 112. **`tts_div`'s Newton step relies on `mad` being fused.** The source kernels divide as the - CPU's IEEE division rounds through `tts_div` (`dasllama_vulkan_classes.das`): a reciprocal + CPU's IEEE division rounds through `tts_div` (`dasllama_gpu_math.das`, both homes' bodies): a reciprocal refined once and the quotient corrected by its exact remainder, each correction a `mad` whose exactness needs one rounding. `mad` lowers to `Fma`, which the `precise` mark does not decorate and which Vulkan lets a driver evaluate as a multiply then an add @@ -1770,12 +1770,6 @@ module) is independent and can land any time - it is pure structure. windowed the same way, or the sample-rate convs' columns gathered inside the tile so no column slot exists; the instrument is the attach ledger's `vk tts scratch` line and the codec served cells. -107. **The Pocket prompt stays on the CPU.** The text rows' backbone prefill over the voice's - caches runs the CPU chain, since the family's hook record carries a codec and a frames seat - only (the Metal twin's too); its share of a steady Pocket sentence on the pod is the prompt - bucket of the `PERF_LEDGER.md` Vulkan TTS seats entry. A prompt seat is the frames seat's layer - chain at t = n_txt rows on the f32 tile, writing the caches' rows the frames seat reads; the - instrument is `harness/tts_synth.das`'s prompt bucket beside the two rows. 106. **The TTS LSTM recurrence walks one SM.** `TtsLstmDir` runs a direction in one workgroup (four lanes a hidden unit over the recurrence transposed to [H][4H]), and a step costs the same whatever the k loop's shape (four independent accumulators read the same as one): the @@ -1818,7 +1812,7 @@ module) is independent and can land any time - it is pure structure. `vk_prof()` reports) under `DASLLAMA_GPU_PROF=1` - one cell reading the report text; the bench's `--asr-clips` override and the served-clip stamping, and the `DASLLAMA_VK_WDEC` knob read; the lower-side guards of `TowerDwConv5` and `TowerCnDw` (`st >= 0` in - `conv_src_fwd_row`, which their `DwConvRowsT.src_row` calls) and `TowerDwConv` (`iy >= 0`, + `conv_src_fwd_row`, which their `GkDwConvRows.src_row` calls) and `TowerDwConv` (`iy >= 0`, `ix >= 0`) - a wrapped index lands gigabytes past any buffer and robust buffer access reads zero, so the fixtures cannot reach them without a source base offset in the args that lets a garbage prefix sit before the block (the upper-side guards are @@ -1936,19 +1930,25 @@ module) is independent and can land any time - it is pure structure. skip on it as they do on a missing device) or the offending opcode is spelled the way MoltenVK's converter accepts, and the three files read green under MoltenVK. -115. **Seventeen TTS kernels are one algorithm under two class shells.** `TtsSrcCumsumT`, `TtsStftT`, - `TtsIstft`, `TtsSrcSinesT`, `TtsSrcLowT`, `TtsAdainT`, `TtsPkAttn`, `TtsPkGemvT`, - `TtsPoolDw`, `TtsIm2colT`, `TtsElemT`, `TtsAddScale`, `TtsPkRowScale`, `TtsRowGather`, - `TtsPkRowsT` (the reflect pad's two copies included), `TtsSigSum` and - `TtsSrcNoise` each have a `MetalSt2*` / `MetalPk*` twin whose body is the same arithmetic, held in - `dasllama_gpu_math.das`; what stays twice is the shell - the binding declarations, the entry, the - workgroup count. The shell folds as `TtsConcat` and `TtsAxpy` did: one `class template` in - `dasllama_gpu_kernels_common.das` both homes stamp (`ARCHITECTURE_GPU.md#gpu-shared-kernels`). A pair - whose two bodies were tuned apart (`TtsElemT` maps four elements an invocation, its Metal twin one a - thread) keeps both and shares the per-element function. The seat chains the two - drivers write per backend are the other half of the same picture and a separate arc: the host flow of - every TTS seat - the stage ping-pong, the concat when the width differs, the head-block loop - is - identical and could run once over an encoder interface, about 450 lines. +115. **The TTS kernel pairs still under two shells, and what each fold costs.** The kernels both GPU + homes run as one `class template` (`ARCHITECTURE_GPU.md#gpu-shared-kernels`) cover the Pocket + row copy, rope and attention, the duration sums, the whole harmonic source, the STFT pair, AdaIN, + the depthwise rows conv, the residual add-and-scale, Axpy and Concat. What stays twice, by kind: + (a) the same output under a different thread-to-work split - `TtsElemT` / `TtsPkRowScale` / + `TtsGeluTanh` (Vulkan four elements an invocation, Metal one a thread), the reflect pad (two + `TtsPkRows` copies against `MetalSt2Reflect1`), the column statistics (four dispatches against + one), the Pocket row GEMV (`TtsPkGemvT`'s float4 lane against `MetalPkGemvT`'s strided lane, five + stamps) - each folds onto one body and is profiled on both boxes, two bodies kept only where the + profile parts; (b) a different algorithm - the LSTM direction (each home's slab writer lays `w_hh` + out its own way) and the ALBERT attention (a two-pass softmax against an online one) - which stay + two; (c) pairs whose Metal twin serves another family's chain and folds with that family's PR - + `TtsRowGather` (`MetalRowGather` is the decode greedy embed's and the whisper stem's) and + `TtsIm2colT` (`MetalSt2Im2colCm` is the whisper twin). Vulkan's four fused frame-loop GEMV stamps + (`LnQkv`, `AddScale`, `Add`, `LnGelu`) have no Metal twin: once the GEMV template is shared, Metal + gains them and the five-dispatch layer. The seat chains the two drivers write per backend are the + other half of the same picture and a separate arc: the host flow of every TTS seat - the stage + ping-pong, the concat when the width differs, the head-block loop - is identical and could run + once over an encoder interface, about 450 lines. 116. **The speculative round has no Vulkan round seat.** The round's four backend seats (`register_mtp_spec_override`, `_spec_batch_`, `_round_`, `_seam_` in `dasllama_common.das`) diff --git a/modules/dasLLAMA/tests/CLAUDE.md b/modules/dasLLAMA/tests/CLAUDE.md index 9e0622da6b..721d6e95f2 100644 --- a/modules/dasLLAMA/tests/CLAUDE.md +++ b/modules/dasLLAMA/tests/CLAUDE.md @@ -435,7 +435,9 @@ a panel; the rel-shift softmax (`fc_pexp`) over two heads at 24 and 300 keys - t one half step of the oracle's, the stat within 1e-5 of a double sum, pad rows and pad columns left as they were, a poisoned rel score reddening both. The Pocket chain's cells (the same file): the row copies with and without the ELU, the layer scale, the rows rope over a row stride -from a column with the tables bound at a position's row, the attention +from a column with the tables bound at a position's row, the prompt's rope from the voice tables at a +position base (`pk_rope_tab_gate`: the k span of 23 rows from position 37 against `rope_rows`, the q +and v spans untouched, a poisoned k element reddening the compare), the attention row against `attention_causal_rows` over a `TtsKvCache` (every key and an 8-key window, an unseen key's poison staying silent), the rope-and-store kernel's f32 stamp at the frame loop's binds (no bias, the whole head, the tables and the caches at the position's row) against @@ -562,11 +564,11 @@ the pad columns stay out of the sum), `test_vkt_colstats` the column statistics double-precision sums at t 45 / 72 channels, t 300 / 20 and t 1000 / 300 (eight blocks, a lane's second channel), `test_vkt_adain` the statistics then both fused AdaIN stamps (`TtsAdainLeaky`, `TtsAdainSnake`) against -the CPU `adain_rows_into` followed by `leaky_relu` or `snake_rows`, every plane at an element base +the CPU `adain_rows_into` followed by `leaky_relu` or `snake_rows`, the style rows at an element base off zero, plus the leaky stamp in place (bit for bit the out-of-place rows, the input overwritten), -the prefixes before the bases kept, and `test_vkt_add_scale` the residual join (`TtsAddScale`) -bit-exact against (a + b) / sqrt(2) in f32 over 1030 elements at offsets, out of place and in place, -the elements outside the run kept. +the elements past the run kept, and `test_vkt_add_scale` the residual join (`TtsAddScale`) +bit-exact against (a + b) / sqrt(2) in f32 over 1030 elements, out of place and in place, +the elements past the run kept. `test_vulkan_tts_source_kernels.das` - model-free (a Vulkan device, else skips): the TTS tower's Vulkan decoder tail against the CPU chain - `test_vkt_axpy` holds the stage sum's axpy (`TtsAxpy`) bit-exact against o + 0.25 y in f32 over 1030 elements at offsets, accumulating onto a filled o and @@ -601,7 +603,7 @@ non-negative ones pass bit for bit), every element outside the window left at it `test_vkt_pk_attn` the causal cached attention row (`TtsPkAttn`) against `attention_causal_rows` over a `TtsKvCache` of two 64-wide heads, the device's key and value rows written from the cache's own layouts with three poisoned rows past the appended 45 - a 45-query prompt over every key, a -decode step at position 44 over an 8-key window (the query at a row offset of the q plane) and the +decode step at position 44 over an 8-key window (the query its own one-row plane) and the prompt over a 16-key window - at the approx bar (controls: a poisoned key just before the decode step's window leaves its output bit for bit, the 16-key window moves the prompt's rows); `test_vkt_pk_gemv` the row GEMV's five stamps (`TtsPkGemvDot`, `TtsPkGemvLnSilu`, `TtsPkGemvGate`, @@ -968,8 +970,9 @@ facade lint trips DASLLAMA001 (code 50503) on a direct engine require with no es guard trips on a path-require resolving into modules/dasLLAMA, name prefix or not. `test_dasllama_lint_escape.das` - model-free: `options _dasllama_internal = true` admits a direct engine require (the lint's escape hatch). -`test_dasllama_lint_contracts.das` - model-free: the lint's ALLOWED set (a facade-only program -with no escape compiles; an internal require does not) via spawned compiles, +`test_dasllama_lint_contracts.das` - model-free: the lint's ALLOWED set (one program on the +facade, scheduler, exchange-schema and bench entries with no escape compiles - one engine-wide +spawn, since each costs the windows nightly runner two minutes; an internal require does not), `load_audio_16k_mono`'s empty-on-failure contract, `decode_audio_16k_mono`'s frame cap (a synthetic `sampleRate=1` WAV bomb is refused before decode, an uncapped call still works), and `gemma4a_probe_proj_dim`'s 0-not-panic contract on `.dlim` / missing / non-GGUF inputs. @@ -2382,7 +2385,13 @@ pod) with the generator check, batches of three (the encode count within one of the batch the EOS frame lands in may split the tail - the latents within the bar, a batch below one clamping to one), and the seats taken by an empty record and given back (`register_pocket_gpu` / `unregister_pocket_gpu`, -the hook unreached then serving again); the frames seat after the LLM tier's model drop in the same +the hook unreached then serving again); the prompt seat's model-free rules (`test_pocket_prompt_rules`: the cache admission's four refusals on a +bare voice state, and the residency record spent by one take and missing on another slot build, row +count, fill or embedding rows); the prompt seat (`test_pocket_prompt_gpu`: every layer's K and V rows of one case's text, the +seat against the CPU chain, within `GPU_PROMPT_BAR` on the f32 lane and `GPU_PROMPT_SERVED_BAR` on +the q8 and kq files' served planes, each with the bar's one-element control, the text rotated by one +token as the compare's control, one hook call and one encode served, the knob-off leg bit-equal with +its decline); the frames seat after the LLM tier's model drop in the same process (`test_pocket_frames_after_model_drop`: one served leg on the q8 file, `moe_gpu_drop_model`, the same leg again serving and reading the same frame count - the slabs, the scratch and the voice slot rebuilt behind the drop); the served frames cell takes both oracle voices in turn @@ -2390,8 +2399,8 @@ on each file, so the second voice's slot displaces the first's, within `GPU_FRAM with the x3-scaled noise as the compare's control; the seat record's refusal of a name no seat carries and its seat names in order (`test_pocket_seat_stats`, model-free); the long chunk's codec seat declining by shape; and the served synthesis across the knob, every -chunk's codec and frame loop served, the encodes -past one a chunk, the knob-off chunks declining at both seats - where no GPU device serves all +chunk's prompt, codec and frame loop served, the encodes +past one a chunk, the knob-off chunks declining at every seat - where no GPU device serves all three skip loudly, a present device that declines is a red; the parity, stream and frames cells pin the tower off, since the CPU chain is what they hold; the published Q8_0 file (`pocket-tts-en-q8.gguf`) against the f16 file's load-time quants - every backbone GEMM arrived diff --git a/modules/dasLLAMA/tests/run.das b/modules/dasLLAMA/tests/run.das index a210353bc4..0f896b8737 100644 --- a/modules/dasLLAMA/tests/run.das +++ b/modules/dasLLAMA/tests/run.das @@ -355,7 +355,7 @@ let MODULE_AREAS <- { "dasllama_vulkan_tower" => "audio,vision,tts", // the Vulkan TTS driver rides the tower's classes "dasllama_tts_slab" => "audio,vision,tts", // the Metal tower requires it for every seat "dasllama_vulkan_tts" => "tts", - "dasllama_gpu_kernels_common" => "tts", // the kernels both GPU homes stamp: the TTS chains' so far + "dasllama_gpu_kernels_common" => "audio,vision,tts", // the kernels both GPU homes stamp: the TTS chains', and the depthwise rows and rope the audio and vision towers stamp "dasllama_asr" => "audio", "dasllama_asr_types" => "audio", "dasllama_audio" => "audio", "dasllama_audio_embedder" => "audio", "dasllama_audio_io" => "audio", "dasllama_canary" => "audio", "dasllama_gemma4a" => "audio", "dasllama_metal_asr_dec" => "audio", "dasllama_vulkan_asr_dec" => "audio", "dasllama_parakeet" => "audio", diff --git a/modules/dasLLAMA/tests/test_dasllama_lint_contracts.das b/modules/dasLLAMA/tests/test_dasllama_lint_contracts.das index 25a0047a90..10a8646be6 100644 --- a/modules/dasLLAMA/tests/test_dasllama_lint_contracts.das +++ b/modules/dasLLAMA/tests/test_dasllama_lint_contracts.das @@ -18,8 +18,12 @@ def fixture_source(requires : array) : string { return join(requires, "\n") + "\n\n[export]\ndef main() " + "\{\n print(\"ok\\n\")\n\}\n" } +//! an engine-wide compile: the facade pulls every engine module, and the windows nightly runner compiles it in about two +//! minutes, so the cap sits well past that +let ENGINE_COMPILE_CAP_SEC = 480.0 + def spawn_compile(requires : array; tag : string) : tuple { - return spawn_compile_only(fixture_source(requires), "lint_{tag}", 120.0) + return spawn_compile_only(fixture_source(requires), "lint_{tag}", ENGINE_COMPILE_CAP_SEC) } //! a module with a function global above its `[init, boot_restore]` function, and one below it where `below` is set: the annotation @@ -61,15 +65,11 @@ def test_boot_restore_order_gate(t : T?) { [test] def test_allowed_set_admits_the_entry_modules(t : T?) { - t |> run("a facade-only program with no escape compiles") @(t : T?) { - let r = spawn_compile(["require dasllama/dasllama"], "facade") - t |> equal(0, r.rc, "facade require is allowed: {r.out}") - t |> success(find(r.out, "DASLLAMA001") < 0, "no lint finding on the facade") - } - t |> run("the scheduler, exchange and bench entries are allowed too") @(t : T?) { + t |> run("a program on the facade, scheduler, exchange and bench entries with no escape compiles") @(t : T?) { let r = spawn_compile(["require dasllama/dasllama", "require dasllama/dasllama_scheduler", "require dasllama/dasllama_exchange_schema", "require dasllama/dasllama_bench"], "entries") - t |> equal(0, r.rc, "scheduler + exchange_schema + bench allowed: {r.out}") + t |> equal(0, r.rc, "facade + scheduler + exchange_schema + bench allowed: {r.out}") + t |> success(find(r.out, "DASLLAMA001") < 0, "no lint finding on the entry modules") } t |> run("an engine internal with no escape does not compile") @(t : T?) { let r = spawn_compile(["require dasllama/dasllama_gguf"], "internal") diff --git a/modules/dasLLAMA/tests/test_metal_prefill_kernels.das b/modules/dasLLAMA/tests/test_metal_prefill_kernels.das index 565e47d5af..9a237d4e49 100644 --- a/modules/dasLLAMA/tests/test_metal_prefill_kernels.das +++ b/modules/dasLLAMA/tests/test_metal_prefill_kernels.das @@ -1316,7 +1316,7 @@ def private st2_adain_gate(t : T?; var dev, queue; nrows : int; snake : bool) { gb[i] = f16_at((i * 2) % 11, 0.01, -0.05) alpha[i] = f16_at((i * 3) % 7, 0.1, 0.3) } - var ka = St2AdainArgs(c = uint(c), t = uint(nrows), eps = 1e-5) + var ka = GkAdainArgs(c = uint(c), t = uint(nrows), hoff = 0u, gwoff = 0u, gboff = uint(c), aoff = uint(2 * c)) var inscope want_st : array var inscope want : array want_st |> resize(2 * c) @@ -1347,9 +1347,10 @@ def private st2_adain_gate(t : T?; var dev, queue; nrows : int; snake : bool) { let bst = mc |> out_words(2 * c, -1.0e30) let by = mc |> out_words(n, -1.0e30) let bh = mc |> in_(h) - let bgw = mc |> in_(gw) - let bgb = mc |> in_(gb) - let balpha = mc |> in_(alpha) + var inscope wn <- clone_to_move(gw) + wn |> push_from(gb) + wn |> push_from(alpha) + let bwn = mc |> in_(wn) let bx = mc |> in_(x) let bt = mc |> in_plain([uint(nrows)]) let bc = mc |> in_plain([uint(c)]) @@ -1366,11 +1367,11 @@ def private st2_adain_gate(t : T?; var dev, queue; nrows : int; snake : bool) { } if (snake) { stamp_run(ssn, pso) { - invoke(ssn.enc, enc, px, mc.buf[by], pst, mc.buf[bh], 0ul, mc.buf[bgw], 0ul, mc.buf[bgb], 0ul, mc.buf[balpha], 0ul, ka, int64(n)) + invoke(ssn.enc, enc, px, mc.buf[by], pst, mc.buf[bh], mc.buf[bwn], ka, int64(n)) } } else { stamp_run(slk, pso) { - invoke(slk.enc, enc, px, mc.buf[by], pst, mc.buf[bh], 0ul, mc.buf[bgw], 0ul, mc.buf[bgb], 0ul, ka, int64(n)) + invoke(slk.enc, enc, px, mc.buf[by], pst, mc.buf[bh], mc.buf[bwn], ka, int64(n)) } } } @@ -1445,7 +1446,6 @@ def private st2_rows_gate(t : T?; var dev, queue) { var bo = mc.buf[po] var (bn, bslope, bc, bnrf) = (mc.buf[mc |> in_plain([uint(n)])], mc.buf[mc |> in_plain([0.1])], mc.buf[mc |> in_plain([uint(c)])], mc.buf[mc |> in_plain([uint(n + c)])]) - var (bsc1, bscj) = (mc.buf[mc |> in_plain([1.0])], mc.buf[mc |> in_plain([0.70710678])]) let run_leaky <- @(bxx : MetalBuffer?) : bool { return queue_record(t, queue, "{tag} leaky") $(enc : MetalComputeEncoder?) { stamp_run(slk, pso_lk) { @@ -1461,10 +1461,9 @@ def private st2_rows_gate(t : T?; var dev, queue) { } } let run_add <- @(bax, bbx : MetalBuffer?; joined : bool) : bool { - let bs = joined ? bscj : bsc1 return queue_record(t, queue, "{tag} add") $(enc : MetalComputeEncoder?) { stamp_run(sad, pso_add) { - invoke(sad.enc, enc, bax, 0ul, bbx, 0ul, bn, bs, bo, 0ul, int64(n)) + invoke(sad.enc, enc, bax, bbx, bo, GkAddScaleArgs(nelem = uint(n), scale = joined ? 0.70710678 : 1.0), int64(n)) } } } @@ -1678,8 +1677,9 @@ def private st2_pool_gate(t : T?; var dev, queue) { want[p * c + j] = float(acc) } } - let pw = mc |> in_(w) - let pb = mc |> in_(b) + var inscope wn <- clone_to_move(w) + wn |> push_from(b) + let pwn = mc |> in_(wn) let py = mc |> out_words(t_out * c, -1.0e30) let px = mc |> in_(x) for (leg in range(2)) { @@ -1687,10 +1687,11 @@ def private st2_pool_gate(t : T?; var dev, queue) { x[4 * c + 6] += 2.0 mc |> plane_upload(px, x) } - var ka = St2PoolArgs(c = uint(c), k = uint(k), stride = uint(stride), pad_l = uint(pad_l), dil = 1u, t_in = uint(t_in), total = uint(t_out * c)) + var ka = GkDwConvRowsArgs(c = uint(c), k = uint(k), stride = uint(stride), pad_l = uint(pad_l), dil = 1u, t_in = uint(t_in), t_out = uint(t_out), + woff = 0u, boff = uint(length(w))) let ran = mc |> cell_record(t, tag) $(enc : MetalComputeEncoder?) { stamp_run(st, pso) { - invoke(st.enc, enc, mc.buf[px], mc.buf[py], mc.buf[pw], 0ul, mc.buf[pb], 0ul, ka, int64(t_out * c)) + invoke(st.enc, enc, mc.buf[px], mc.buf[py], mc.buf[pwn], ka, int64(t_out * c)) } } continue if (!ran) @@ -1832,7 +1833,7 @@ def private pk_rows_gate(t : T?; var dev, queue; elu : bool) { return if (!pso_ok(t, tag, pso, err)) let rows = 10 let width = 40 - let ka = PkRowsArgs(rows = uint(rows), width = uint(width), ss = 48u, sc = 8u, sr0 = 2u, ds = 64u, dc = 5u, dr0 = 3u) + let ka = GkPkRowsArgs(rows = uint(rows), width = uint(width), ss = 48u, sc = 8u, sr0 = 2u, ds = 64u, dc = 5u, dr0 = 3u) var inscope x <- f16_fill((2 + rows) * 48, 3, 0.5, -1.0) var inscope want : array want |> resize((3 + rows) * 64) @@ -1910,6 +1911,70 @@ def private pk_row_scale_gate(t : T?; var dev, queue) { //! rows [7][96] rotated over 64 columns from column 16 (two heads of 32) at positions 5..: the CPU's `rope_rows` on the same columns //! The codec's rope: MetalRope over rows at a stride from a column, the tables' first row position 0 +//! the Pocket prompt's rope over the k span of t rows from a position base, the tables bound whole and read at that +//! position's row - the position-based rope both homes stamp (`GkRopeTab`); the q and v spans untouched +def private pk_rope_tab_gate(t : T?; var dev, queue) { + let tag = "pk_rope_tab" + with_metal_cell(dev, queue) $(var mc : MetalCell) { + var err : string + let st = stamp_of(@@enc_pk_rope_tab) + var pso = mc |> keep_stamp(st, err) + return if (!pso_ok(t, tag, pso, err)) + let tt = 23 + let d = 128 + let dh = 64 + let qs = 3 * d + let pos0 = 37 + let cap = 64 + let half = dh / 2 + let period = 10000.0 + var inscope x <- f16_fill(tt * qs, 41, 0.05, -2.1) + var inscope tcos : array + var inscope tsin : array + let none : array + build_rope_tabs(tcos, tsin, period, 1.0, 1.0, none, false, 0l, int64(cap), int64(dh)) + t |> equal(length(tcos), cap * half, "{tag}: the tables carry {cap} positions of {half} angles") + var inscope want <- [for (e in range(tt * d)); x[(e / d) * qs + d + e % d]] + rope_rows(want, int64(tt), int64(d), int64(dh), int64(pos0), period) + let ka = GkRopeTabArgs(rows = uint(tt), qs = uint(qs), col = uint(d), d = uint(d), dh = uint(dh), coff = uint(pos0 * half), soff = uint(pos0 * half)) + let pc = mc |> in_(tcos) + let ps = mc |> in_(tsin) + var bc = mc.buf[pc] + var bs = mc.buf[ps] + let run <- @(bxx : MetalBuffer?) : bool { + var kl = ka + return queue_record(t, queue, tag) $(enc : MetalComputeEncoder?) { + stamp_run(st, pso) { + invoke(st.enc, enc, bxx, bc, bs, kl, int64(tt * d / 2)) + } + } + } + let px = mc |> in_(x) + if (invoke(run, mc.buf[px])) { + var inscope rows <- mc |> plane_whole(px) + var inscope got <- [for (e in range(tt * d)); rows[(e / d) * qs + d + e % d]] + check_rows(t, "{tag} k span vs the CPU rope_rows", got, want, Bar(rel = 1e-5, abs_floor = 1e-6, poison_scaled = true), d, dh) + t |> success(got[2] != x[d + 2], "{tag} in place control: a pair at a nonzero angle is turned off its input") + var kept = 0 + for (r in range(tt)) { + for (j in range(qs)) { + if ((j < d || j >= 2 * d) && rows[r * qs + j] == x[r * qs + j]) { + kept++ + } + } + } + t |> equal(kept, tt * 2 * d, "{tag}: the q and v spans are untouched") + } + x[3 * qs + d + 5] += 1.5 + let px2 = mc |> in_(x) + if (invoke(run, mc.buf[px2])) { + var inscope rows2 <- mc |> plane_whole(px2) + var inscope got2 <- [for (e in range(tt * d)); rows2[(e / d) * qs + d + e % d]] + t |> success(mismatch_rel("{tag} poison (MUST mismatch)", got2, want, 1e-5, 1e-6) > 0, "{tag}: a poisoned k element reds the compare") + } + } +} + def private rope_stride_gate(t : T?; var dev, queue) { let tag = "rope stride" with_metal_cell(dev, queue) $(var mc : MetalCell) { @@ -2018,7 +2083,7 @@ def private pk_attn_gate(t : T?; var dev, queue; ctx, n_old : int) { with_dasllama_jobque_() { attention_causal_rows(qn, int64(tt), int64(d), int64(ctx), kc, want) } - let ka = PkAttnArgs(t = uint(tt), qs = uint(qs), qcol = uint(qcol), d = uint(d), pos0 = uint(n_old), ctx = uint(ctx)) + let ka = GkPkAttnArgs(t = uint(tt), qs = uint(qs), qcol = uint(qcol), d = uint(d), pos0 = uint(n_old), ctx = uint(ctx)) let pq = mc |> in_(q) let pv = mc |> in_(vv) var bq = mc.buf[pq] @@ -2416,7 +2481,7 @@ def private st2_sigsum_gate(t : T?; var dev, queue) { } let praw = mc |> out_words(nrows, -1.0e30) let plg = mc |> in_(lg) - let (pnb, pls, pspeed, pt) = (mc |> in_plain([uint(nb)]), mc |> in_plain([uint(ls)]), mc |> in_plain([1.3]), mc |> in_plain([uint(nrows)])) + var ka = GkSigSumArgs(t = uint(nrows), nb = uint(nb), ls = uint(ls), speed = 1.3) for (leg in range(2)) { if (leg == 1) { lg[5 * ls + 7] += 6.0 @@ -2424,7 +2489,7 @@ def private st2_sigsum_gate(t : T?; var dev, queue) { } let ran = mc |> cell_record(t, tag) $(enc : MetalComputeEncoder?) { stamp_run(st, pso) { - invoke(st.enc, enc, mc.buf[plg], mc.buf[praw], mc.buf[pnb], mc.buf[pls], mc.buf[pspeed], mc.buf[pt], int64(nrows)) + invoke(st.enc, enc, mc.buf[plg], mc.buf[praw], ka, int64(nrows)) } } continue if (!ran) @@ -2494,24 +2559,24 @@ def private st2_source_gate(t : T?; var dev, queue; torch_math : bool) { let n = src.n let nlow = src.nlow var inscope f0 := src.f0 - let ka = St2SourceArgs(n = uint(n), nlow = uint(nlow), nh = uint(nh), up = uint(up), sr = TTS_SRC_SR, sine_amp = 0.1, noise_std = 0.003, - voiced_thr = 10.0, seed = 77u) + let ka = GkSourceArgs(n = uint(n), nlow = uint(nlow), nh = uint(nh), up = uint(up), sr = TTS_SRC_SR, sine_amp = 0.1, noise_std = 0.003, + voiced_thr = 10.0, seed = 77u, woff = 0u, boff = uint(length(src.lin_w))) t |> equal(src.rows, int64(n), "{tag}: the CPU source mixes every sample") let pnu = mc |> in_(src.nu) let pnoise = mc |> in_(src.noise) - let plw = mc |> in_(src.lin_w) - let plb = mc |> in_(src.lin_b) + var inscope wn <- clone_to_move(src.lin_w) + wn |> push_from(src.lin_b) + let plw = mc |> in_(wn) let plow = mc |> out_words(nh * nlow, -1.0e30) let plowc = mc |> out_words(nh * nlow, -1.0e30) let phar = mc |> out_words(n, -1.0e30) var bnu = mc.buf[pnu] var bnoise = mc.buf[pnoise] var blw = mc.buf[plw] - var blb = mc.buf[plb] var blow = mc.buf[plow] var blowc = mc.buf[plowc] var bhar = mc.buf[phar] - let run <- @(bf0 : MetalBuffer?; kax : St2SourceArgs; own : bool; bnz : MetalBuffer?; bhar : MetalBuffer?) : bool { + let run <- @(bf0 : MetalBuffer?; kax : GkSourceArgs; own : bool; bnz : MetalBuffer?; bhar : MetalBuffer?) : bool { var kk = kax return queue_record(t, queue, "{tag} dispatch") $(enc : MetalComputeEncoder?) { if (own) { // the driver's own draw: the noise rows hashed on the device before the sines read them @@ -2526,7 +2591,7 @@ def private st2_source_gate(t : T?; var dev, queue; torch_math : bool) { invoke(scum.enc, enc, blow, blowc, kk, int64(nh)) } stamp_run(ssin, pso_sin) { - invoke(ssin.enc, enc, bf0, blowc, bnz, blw, 0ul, blb, 0ul, bhar, kk, int64(n)) + invoke(ssin.enc, enc, bf0, blowc, bnz, blw, bhar, kk, int64(n)) } } } @@ -2614,8 +2679,9 @@ def private st2_stft_gate(t : T?; var dev, queue; reflect : bool) { var inscope wim <- f16_fill(bins * k, 3, 0.15, -0.5) wim[2 * k + 5] = 0.0 var inscope want <- stft_ref(har, wre, wim, n, k, hop, pad, bins, reflect) - let pre = mc |> in_(wre) - let pim = mc |> in_(wim) + var inscope wn <- clone_to_move(wre) + wn |> push_from(wim) + let pwn = mc |> in_(wn) let pspec = mc |> out_words(frames * 2 * bins, -1.0e30) let ph = mc |> in_(har) for (leg in range(2)) { @@ -2623,10 +2689,11 @@ def private st2_stft_gate(t : T?; var dev, queue; reflect : bool) { har[1] += 2.0 //! the first frame's reflected tap reads it mc |> plane_upload(ph, har) } - var ka = St2StftArgs(n = uint(n), frames = uint(frames), bins = uint(bins), k = uint(k), hop = uint(hop), pad = uint(pad), eps = 1e-9) + var ka = GkStftArgs(n = uint(n), frames = uint(frames), bins = uint(bins), k = uint(k), hop = uint(hop), pad = uint(pad), eps = 1e-9, + reoff = 0u, imoff = uint(bins * k)) let ran = mc |> cell_record(t, tag) $(enc : MetalComputeEncoder?) { stamp_run(st, pso) { - invoke(st.enc, enc, mc.buf[ph], mc.buf[pre], 0ul, mc.buf[pim], 0ul, mc.buf[pspec], ka, int64(frames * bins)) + invoke(st.enc, enc, mc.buf[ph], mc.buf[pwn], mc.buf[pspec], ka, int64(frames * bins)) } } continue if (!ran) @@ -2659,9 +2726,10 @@ def private st2_istft_gate(t : T?; var dev, queue; envelope : bool) { var inscope wim <- f16_fill(k * bins, 3, 0.15, -0.5) var inscope window <- f16_fill(k, 11, 0.1, 0.3) var inscope ref <- istft_ref(y, wre, wim, window, tt, stride, bins, k, hop, pad, envelope) - let pre = mc |> in_(wre) - let pim = mc |> in_(wim) - let pwn = mc |> in_(window) + var inscope wn <- clone_to_move(wre) + wn |> push_from(wim) + wn |> push_from(window) + let pwn = mc |> in_(wn) let pw = mc |> out_words(n, -1.0e30) let py = mc |> in_(y) for (leg in range(2)) { @@ -2669,11 +2737,11 @@ def private st2_istft_gate(t : T?; var dev, queue; envelope : bool) { y[4 * stride + 1] += 2.0 //! a log-magnitude: every sample its frame overlaps moves mc |> plane_upload(py, y) } - var ka = St2IstftArgs(tt = uint(tt), stride = uint(stride), bins = uint(bins), k = uint(k), hop = uint(hop), pad = uint(pad), - envelope = envelope ? 1u : 0u, n = uint(n)) + var ka = GkIstftArgs(tt = uint(tt), stride = uint(stride), bins = uint(bins), k = uint(k), hop = uint(hop), pad = uint(pad), + envelope = envelope ? 1u : 0u, n = uint(n), reoff = 0u, imoff = uint(k * bins), wnoff = uint(2 * k * bins)) let ran = mc |> cell_record(t, tag) $(enc : MetalComputeEncoder?) { stamp_run(st, pso) { - invoke(st.enc, enc, mc.buf[py], mc.buf[pre], 0ul, mc.buf[pim], 0ul, mc.buf[pwn], 0ul, mc.buf[pw], ka, int64(n)) + invoke(st.enc, enc, mc.buf[py], mc.buf[pwn], mc.buf[pw], ka, int64(n)) } } continue if (!ran) @@ -2931,6 +2999,7 @@ def test_metal_prefill_kernels(t : T?) { pk_rows_gate(t, dev, queue, true) pk_row_scale_gate(t, dev, queue) rope_stride_gate(t, dev, queue) + pk_rope_tab_gate(t, dev, queue) pk_attn_gate(t, dev, queue, 0, 10) pk_attn_gate(t, dev, queue, 8, 10) pk_attn_gate(t, dev, queue, 0, 300) //! past one 256-key stride of the staged scores diff --git a/modules/dasLLAMA/tests/test_tts_pocket.das b/modules/dasLLAMA/tests/test_tts_pocket.das index 05badd9fd2..e388eab1be 100644 --- a/modules/dasLLAMA/tests/test_tts_pocket.das +++ b/modules/dasLLAMA/tests/test_tts_pocket.das @@ -11,6 +11,7 @@ require dasllama/dasllama_tts require dasllama/dasllama_audio_io require dasllama/dasllama_tts_types require dasllama/dasllama_tts_blocks +require dasllama/dasllama_tts_slab // the prompt's cache admission and residency record, model-free require dasllama/dasllama_math // with_dasllama_jobque_ require dasllama/dasllama // moe_gpu_drop_model: the LLM tier's model drop between two syntheses require daslib/jobque_boost @@ -42,6 +43,8 @@ require _model_tier let GPU_CODEC_BAR = 1e-5 //! the codec seat against the CPU chain on the f32 lane: reads 9e-7 on the exact stamps (M5 Max, -jit) let GPU_FRAME_BAR = 2e-5 //! the frames seat against the CPU chain on the f32 lane, teacher-forced: reads 2e-6 on the latents (M5 Max, -jit) let GPU_CODEC_SERVED_BAR = 2e-1 //! the codec seat on the served planes: reads 1.3e-2 / 7.7e-3 (q8 / kq) on the M5 Max and 9.2e-2 / 1.5e-1 on the pod's x64 CPU chain, all of it the CPU's Q8_0 activation blocks of the first layer's rows (the tower reads the same quants dequantized and feeds f32: it sits at 3e-6 of the CPU f32 chain on either box) +let GPU_PROMPT_BAR = 1e-5 //! the prompt seat's K/V rows against the CPU chain on the f32 lane (`test_pocket_prompt_gpu`, `-jit`, no overrides): reads 5e-7 / 1.6e-6 (keys / values) on the M5 Max under its tune sidecar, 5e-7 / 1.5e-6 on the RTX PRO 4500 +let GPU_PROMPT_SERVED_BAR = 5e-2 //! the prompt seat's K/V rows on the served planes against the file's own CPU chain (the same cell): reads 3e-3 / 1.2e-2 on the q8 file, 7e-3 / 2.8e-2 on the kq file (M5 Max); 3e-3 / 1.2e-2 and 6e-3 / 2.3e-2 on the RTX PRO 4500 let GPU_FRAME_SERVED_BAR = 1e-1 //! the frames seat on the served planes, teacher-forced: reads 1.9e-2 / 3.7e-2 on the q8 file, 3.8e-2 / 7.1e-2 on the kq file (alba / caro_davy) on the M5 Max let GPU_FRAME_FREE_BAR = 5e-3 //! the frames seat free-running on the f32 lane, its own noise from the same seed, every frame fed its own output: reads 1.5e-4 and 1.1e-4 on the two stage cases on the M5 Max, 7e-5 and 1.8e-3 on the pod (lj028's 150 frames compound the f16 GEMM feed's rounding) @@ -629,7 +632,7 @@ def private frames_seat_leg(t : T?; m : PocketModel; var k : ForcedCase; tag : s return <- r if (fg != fc) t |> success(same_rng(sc.rng, rng_cpu), "{tag}: the generator ends where the CPU loop's does") let batch = pocket_frame_batch() - t |> equal(e1 - e0, (fc + batch) / batch + 1l, "{tag}: one tower encode a batch of {batch}, and the codec's") + t |> equal(e1 - e0, (fc + batch) / batch + 2l, "{tag}: one tower encode a batch of {batch}, the prompt's and the codec's") var inscope gpu <- frame_rows(m, sc, fg) let dl = rel_l2(gpu.lat, r.cpu.lat) let dc = rel_l2(gpu.conds, r.cpu.conds) @@ -658,6 +661,54 @@ def private frames_seat_leg(t : T?; m : PocketModel; var k : ForcedCase; tag : s return <- r } +//! The prompt seat's two model-free rules on a bare voice state: the cache admission's four refusals, and the +//! residency record - spent by one take, and held to the voice, the slot build, the fill, the row count and the +//! embedding rows it was set under. +[test] +def test_pocket_prompt_rules(t : T?) { + t |> run("the prompt admission refuses an empty prompt, a cache count off the layers, rows past the capacity and a capacity past the keys") @(t : T?) { + var inscope m = PocketModel() + m.backbone.layers |> resize(2) + m.backbone.d = 8l + m.backbone.heads = 2l + var inscope vs = PocketVoiceState(len = 5l) + vs.caches |> resize(2) + for (kc in vs.caches) { + kv_cache_init(kc, 2l, 4l, 20l) + } + t |> success(tts_pk_prompt_admit(m, vs, 3l, 64l), "a three-row prompt over a 20-row cache of two layers is admitted") + t |> success(!tts_pk_prompt_admit(m, vs, 0l, 64l), "an empty prompt is refused") + t |> success(!tts_pk_prompt_admit(m, vs, 16l, 64l), "voice + text past the capacity is refused") + t |> success(tts_pk_prompt_admit(m, vs, 15l, 20l), "a capacity at the keys' cap is admitted") + t |> success(!tts_pk_prompt_admit(m, vs, 3l, 16l), "a capacity past the keys' cap is refused") + vs.caches |> resize(1) + t |> success(!tts_pk_prompt_admit(m, vs, 3l, 64l), "a cache count off the layers is refused") + } + t |> run("the residency record is spent by one take and held to what it was set under") @(t : T?) { + var inscope vs = PocketVoiceState(len = 5l) + vs.caches |> resize(1) + kv_cache_init(vs.caches[0], 2l, 4l, 20l) + var inscope rows <- [for (i in range(3 * 8)); float(i) * 0.25] + var rec : TtsPkPromptDev + t |> equal(tts_pk_prompt_dev_take(rec, vs, 7ul, rows, 3l, 8l), 0l, "an empty record holds no rows") + tts_pk_prompt_dev_set(rec, vs, 7ul, rows, 3l, 8l) + t |> equal(tts_pk_prompt_dev_take(rec, vs, 7ul, rows, 3l, 8l), 3l, "the record set under this voice, build, prompt and rows hands them back") + t |> equal(tts_pk_prompt_dev_take(rec, vs, 7ul, rows, 3l, 8l), 0l, "one take spends it") + tts_pk_prompt_dev_set(rec, vs, 7ul, rows, 3l, 8l) + t |> equal(tts_pk_prompt_dev_take(rec, vs, 8ul, rows, 3l, 8l), 0l, "another slot build misses it") + tts_pk_prompt_dev_set(rec, vs, 7ul, rows, 3l, 8l) + t |> equal(tts_pk_prompt_dev_take(rec, vs, 7ul, rows, 2l, 8l), 0l, "another row count misses it") + tts_pk_prompt_dev_set(rec, vs, 7ul, rows, 3l, 8l) + var inscope other := rows + other[1]++ + t |> equal(tts_pk_prompt_dev_take(rec, vs, 7ul, other, 3l, 8l), 0l, "other embedding rows of the same length miss it") + tts_pk_prompt_dev_set(rec, vs, 7ul, rows, 3l, 8l) + vs.len = 6l + t |> equal(tts_pk_prompt_dev_take(rec, vs, 7ul, rows, 3l, 8l), 0l, "another fill misses it") + t |> equal(tts_pk_prompt_dev_take(rec, vs, 7ul, rows, 3l, 8l), 0l, "a miss spends it too") + } +} + [test] def test_pocket_seat_stats(t : T?) { t |> run("a seat name no stage carries panics, so a misspelt key cannot read as zero counters") @(t : T?) { @@ -668,9 +719,10 @@ def test_pocket_seat_stats(t : T?) { } t |> run("the record names its seats in dispatch order, and each name reads its counters") @(t : T?) { var inscope seats <- pocket_gpu_seats() - t |> equal(length(seats), 2, "two seats") + t |> equal(length(seats), 3, "three seats") t |> equal(seats[0], "codec", "the codec first") t |> equal(seats[1], "frames", "the frame loop second") + t |> equal(seats[2], "prompt", "the prompt third") for (s in seats) { let st = pocket_gpu_stats(s) t |> success(st.served <= st.calls, "{s}: served never past called") @@ -678,6 +730,98 @@ def test_pocket_seat_stats(t : T?) { } } +typedef PromptRows = tuple; v : array> + +//! every layer's K and V rows [vs.len, vs.len + n) as the prompt left them, each layer's rows after the one before +def private prompt_kv(vs : PocketVoiceState; n : int64) : PromptRows { + var r : PromptRows + let d = vs.caches[0].heads * vs.caches[0].dh + let per = n * d + r.k |> resize(per * long_length(vs.caches)) + r.v |> resize(per * long_length(vs.caches)) + unsafe { + for (li, kc in iter_range(vs.caches), vs.caches) { + kv_cache_read_rows(kc, vs.len, vs.len + n, 0l, addr(r.k[int64(li) * per]), addr(r.v[int64(li) * per])) + } + } + return <- r +} + +//! The prompt seat against the CPU chain over one case's text: every layer's K and V rows within `bar` with the bar's +//! one-element control and the text rotated by one token as the compare's control, one hook call served, one encode, +//! the knob-off leg bit-equal with its decline recorded; false where no device serves (the skip registered). +[cold_path, unused_argument(t, m, k, tag, bar, sc)] // off-Apple the gated body is empty +def private prompt_seat_leg(t : T?; m : PocketModel; var k : ForcedCase; tag : string; bar : float; var sc : PocketScratch) : bool { + static_if (typeinfo builtin_module_exists(das_metal) || typeinfo builtin_module_exists(vulkan)) { + let n = long_length(k.ids) + set_serving_gpu_tower(false) + pocket_prompt_rows(m, k.ids, k.vs, sc) + var inscope cpu <- prompt_kv(k.vs, n) + set_serving_gpu_tower(true) + let h0 = pocket_gpu_stats("prompt") + let e0 = gpu_encodes() + pocket_prompt_rows(m, k.ids, k.vs, sc) + let h1 = pocket_gpu_stats("prompt") + let e1 = gpu_encodes() + return false if (h1.served - h0.served != 1 && !tower_served(t, tag, "{gpu_declines()}")) + t |> equal(h1.calls - h0.calls, 1, "{tag}: one hook call") + t |> equal(e1 - e0, 1l, "{tag}: one tower encode") + var inscope gpu <- prompt_kv(k.vs, n) + let dk = rel_l2(gpu.k, cpu.k) + let dv = rel_l2(gpu.v, cpu.v) + to_log(LOG_INFO, "pocket {tag}: rel-l2 keys {dk} values {dv} against the CPU chain over {n} rows x {length(k.vs.caches)} layers (bar {bar})\n") + t |> success(dk <= bar && dv <= bar, "{tag}: the keys and values within {bar}") + var inscope w <- one_wrong(gpu.v, cpu.v, bar) + t |> success(rel_l2(w, cpu.v) > bar, "{tag}: one value moved by twice the bar reds it (the bar's control)") + knob_off_leg(t, tag, 1l, $ => pocket_gpu_stats("prompt")) $ { + pocket_prompt_rows(m, k.ids, k.vs, sc) + var inscope off <- prompt_kv(k.vs, n) + return rel_l2(off.k, cpu.k) == 0.0 && rel_l2(off.v, cpu.v) == 0.0 + } + set_serving_gpu_tower(true) + var inscope ids_rot <- [for (i in range64(n)); k.ids[(i + 1l) % n]] + pocket_prompt_rows(m, ids_rot, k.vs, sc) + var inscope gx <- prompt_kv(k.vs, n) + t |> success(rel_l2(gx.k, cpu.k) > bar && rel_l2(gx.v, cpu.v) > bar, "{tag}: the text rotated by one token reds the keys' and the values' bars (the compare's control)") + return true + } else { + return false + } +} + +//! The prompt seat on the f32 lane and on the served planes of the q8 and kq files, as `prompt_seat_leg` holds it. +[test] +def test_pocket_prompt_gpu(t : T?) { + let path = gguf_path() + if (!model_available(t, path)) { + return + } + var inscope man <- stocked_manifest(t) + return if (empty(man.cases)) + static_if (typeinfo builtin_module_exists(das_metal) || typeinfo builtin_module_exists(vulkan)) { + pocket_lane_cell(PocketLane.f32) { + var inscope m <- load_pocket(path) + var inscope sc = PocketScratch() + for (c in man.cases) { + continue if (!c.stages || c.voice != "alba") + var inscope k <- forced_case(m, c, man, sc) + return if (!prompt_seat_leg(t, m, k, "{c.id} {c.voice} prompt", GPU_PROMPT_BAR, sc)) + break + } + } + pocket_served_files(t) $(spath : string; m : PocketModel; var sc : PocketScratch) : bool { + for (c in man.cases) { + continue if (!c.stages || c.voice != "alba") + var inscope k <- forced_case(m, c, man, sc) + return prompt_seat_leg(t, m, k, "{c.id} {c.voice} {base_name(spath)} prompt served", GPU_PROMPT_SERVED_BAR, sc) + } + return true + } + } else { + t |> skip("no GPU module in this build") + } +} + //! The frames seat on the Metal tower against the CPU chain on the f32 lane, teacher-forced on the //! oracle's noise and frames: the latents, the conditioning rows and the EOS logits within //! GPU_FRAME_BAR with the bar's one-element control and the x3-scaled noise as the compare's @@ -704,7 +848,7 @@ def test_pocket_frames_metal(t : T?) { var inscope leg <- frames_seat_leg(t, m, k, tag, GPU_FRAME_BAR, true, sc) return if (leg.fc < 0l) let fc = leg.fc - knob_off_leg(t, tag, 2l, $ => pocket_gpu_stats("frames")) $ { + knob_off_leg(t, tag, 3l, $ => pocket_gpu_stats("frames")) $ { let fo = forced_frames(m, k, k.x0, k.x1, sc) var inscope off <- frame_rows(m, sc, fo) return fo == fc && rel_l2(off.lat, leg.cpu.lat) == 0.0 @@ -737,9 +881,9 @@ def test_pocket_frames_metal(t : T?) { let eb1 = gpu_encodes() var inscope b3 <- frame_rows(m, sc, f3) t |> equal(f3, fc, "{tag}: batches of three make the frame count") - let batches3 = (fc + 2l) / 3l + 1l // the frames in threes and the codec's encode; the batch the EOS frame lands in may split the tail into one more + let batches3 = (fc + 2l) / 3l + 2l // the frames in threes, the prompt's and the codec's encodes; the batch the EOS frame lands in may split the tail into one more let extra = eb1 - eb0 - batches3 - t |> success(extra >= 0l && extra <= 1l, "{tag}: one tower encode a batch of three, and the codec's ({eb1 - eb0} encodes over {fc} frames, {batches3} or one more)") + t |> success(extra >= 0l && extra <= 1l, "{tag}: one tower encode a batch of three, the prompt's and the codec's ({eb1 - eb0} encodes over {fc} frames, {batches3} or one more)") t |> success(f3 == fc && rel_l2(b3.lat, leg.cpu.lat) <= GPU_FRAME_BAR, "{tag}: batches of three land within {GPU_FRAME_BAR}") set_pocket_frame_batch(0l) t |> equal(pocket_frame_batch(), 1l, "{tag}: a batch below one clamps to one") @@ -826,7 +970,7 @@ def test_pocket_frames_after_model_drop(t : T?) { } //! The served synthesis across the tower knob on the served lane: every chunk's codec and frame -//! loop on the tower, one codec encode a chunk, the knob-off leg declining both seats per chunk, +//! loop on the tower, one codec encode a chunk, the knob-off leg declining every seat per chunk, //! both legs speaking. [test] def test_pocket_synthesis_metal(t : T?) { diff --git a/modules/dasLLAMA/tests/test_vulkan_tower_kernels.das b/modules/dasLLAMA/tests/test_vulkan_tower_kernels.das index 7497e32dc7..e3182ad1d8 100644 --- a/modules/dasLLAMA/tests/test_vulkan_tower_kernels.das +++ b/modules/dasLLAMA/tests/test_vulkan_tower_kernels.das @@ -590,7 +590,7 @@ def test_vkt_tower_conformer(t0 : T?) { // nolint:STYLE038 — one flat gate o var s_attn = cp |> cls_set(fixed_array(qp, kp, vp, aop, rp, wp), @@set_tower_g4a_attn64_cls) var pa_silu = TowerBiasArgs(d = uint(d), nelem = uint(nd), boff = 0u, act = BIAS_ACT_SILU) var pa_convmod = TowerConvModArgs(d = uint(d), nelem = uint(nd)) - var pa_dw = DwConvRowsArgs(c = uint(d), t_in = uint(rows), t_out = uint(rows), woff = uint(hs)) + var pa_dw = GkDwConvRowsArgs(c = uint(d), t_in = uint(rows), t_out = uint(rows), woff = uint(hs)) let qscale = (1.0 / sqrt(float(hs))) / log(2.0) let kscale = log(1.0 + exp(1.0)) / log(2.0) var pa_attn = TowerG4aAttnArgs(d = uint(d), npos = uint(rows), pdsoff = 0u) @@ -711,7 +711,7 @@ def test_vkt_tower_fastconformer(t0 : T?) { let cpl = cp |> out_(nd) let wp = cp |> in_(wn) var s_dw = cp |> cls_set(fixed_array(gp, cpl, wp), @@set_tower_cn_dw_cls) - var pa_dw = DwConvRowsArgs(c = uint(d), k = uint(CN_CONV_KERNEL), t_in = uint(rows), t_out = uint(rows), woff = uint(taps_off), bnwoff = uint(bnw_off), + var pa_dw = GkDwConvRowsArgs(c = uint(d), k = uint(CN_CONV_KERNEL), t_in = uint(rows), t_out = uint(rows), woff = uint(taps_off), bnwoff = uint(bnw_off), bnboff = uint(bnb_off)) cp |> cell_record() $(raw : VkCommandBuffer; var h : VkHaz) { enc_tower_cn_dw_cls(raw, h, s_dw, pa_dw, int64((nd + 255) / 256)) @@ -787,7 +787,7 @@ def test_vkt_tower_ln_rows(t0 : T?) { // nolint:STYLE038 — one flat gate ove let cosp = cp |> in_(cos_t) let sinp = cp |> in_(sin_t) var s_rope = cp |> cls_set(fixed_array(qkvp, cosp, sinp), @@set_tower_rope_tab_cls) - var pa_rope = RopeTabArgs(rows = uint(rows), qs = uint(3 * d_compact), col = uint(d_compact), d = uint(ROWS_HEADS * LN_HS), dh = uint(LN_HS), + var pa_rope = GkRopeTabArgs(rows = uint(rows), qs = uint(3 * d_compact), col = uint(d_compact), d = uint(ROWS_HEADS * LN_HS), dh = uint(LN_HS), coff = 0u, soff = 0u) var s_ln = cp |> cls_set(fixed_array(xp, xbp, wp), @@set_tower_ln_cls) var s_acts <- [for (ap in act_p); cp |> cls_set(fixed_array(ap, wp), @@set_tower_bias_act_cls)] diff --git a/modules/dasLLAMA/tests/test_vulkan_tts_conv_kernels.das b/modules/dasLLAMA/tests/test_vulkan_tts_conv_kernels.das index 13d374aa84..f9f06189fe 100644 --- a/modules/dasLLAMA/tests/test_vulkan_tts_conv_kernels.das +++ b/modules/dasLLAMA/tests/test_vulkan_tts_conv_kernels.das @@ -387,7 +387,7 @@ def private pool_dw_arm(t : T?) { let yp = cp |> out_(ny) let wp = cp |> in_(wn) var s = cp |> cls_set(fixed_array(xp, yp, wp), @@set_tts_pool_dw_cls) - var pa = DwConvRowsArgs(c = uint(PD_C), k = uint(PD_K), stride = 2u, pad_l = 1u, dil = 1u, t_in = uint(PD_TIN), t_out = uint(tout), + var pa = GkDwConvRowsArgs(c = uint(PD_C), k = uint(PD_K), stride = 2u, pad_l = 1u, dil = 1u, t_in = uint(PD_TIN), t_out = uint(tout), woff = uint(PD_WOFF), boff = uint(PD_WOFF + PD_C * PD_K + PD_GAP)) cp |> cell_record() $(raw : VkCommandBuffer; var h : VkHaz) { enc_tts_pool_dw_cls(raw, h, s, pa, int64(ny)) diff --git a/modules/dasLLAMA/tests/test_vulkan_tts_kernels.das b/modules/dasLLAMA/tests/test_vulkan_tts_kernels.das index 9dfb64305b..841016bb6e 100644 --- a/modules/dasLLAMA/tests/test_vulkan_tts_kernels.das +++ b/modules/dasLLAMA/tests/test_vulkan_tts_kernels.das @@ -375,7 +375,7 @@ def test_vkt_sigsum(t0 : T?) { let lp = cp |> in_(logits) let rp = cp |> out_(SG_T) var s = cp |> cls_set(fixed_array(lp, rp), @@set_tts_sigsum_cls) - var pa = TtsSigSumArgs(t = uint(SG_T), nb = uint(SG_NB), ls = uint(SG_LS), speed = SG_SPEED) + var pa = GkSigSumArgs(t = uint(SG_T), nb = uint(SG_NB), ls = uint(SG_LS), speed = SG_SPEED) cp |> cell_record() $(raw : VkCommandBuffer; var h : VkHaz) { enc_tts_sigsum_cls(raw, h, s, pa, int64(SG_T)) } @@ -485,9 +485,8 @@ let private AD_T = 45 let private AD_C = 72 let private AD_SD = 16 //! the style vector's width let private AD_EPS = 1e-5 -let private AD_XOFF = 5 //! each plane's element base: the prefix before it is never written -let private AD_YOFF = 9 -let private AD_HOFF = 7 +let private AD_HOFF = 7 //! the style rows' element base: the prefix before it is never read +let private AD_SLACK = 5 //! elements past the run, left at the sentinel let private AD_GWOFF = 3 //! the slab's norm scale row, then its shift row and Snake's alpha row, POISON between //! an AdaIN over AD_C channels from an AD_SD-wide style: fc w [2c][sd] and b, the norm's scale and shift rows @@ -550,26 +549,23 @@ def test_vkt_adain(t0 : T?) { var f <- adain_fixture() with_cell_planes() $(var cp : CellPlanes) { let x0p = cp |> in_(f.x) - let xp = cp |> in_(f.x, AD_XOFF) - let xip = cp |> io_(f.x, AD_XOFF) + let xp = cp |> in_(f.x) + let xip = cp |> io_(f.x, 0, AD_SLACK) let hp = cp |> in_(f.style_fc, AD_HOFF) let wp = cp |> in_(f.slab) let stp = cp |> out_(colstats_floats(AD_T, c)) - let ylp = cp |> out_(n, AD_YOFF) - let ysp = cp |> out_(n, AD_YOFF) + let ylp = cp |> out_(n, 0, AD_SLACK) + let ysp = cp |> out_(n, 0, AD_SLACK) var s_l = cp |> cls_set(fixed_array(xp, ylp, stp, hp, wp), @@set_tts_adain_leaky_cls) var s_s = cp |> cls_set(fixed_array(xp, ysp, stp, hp, wp), @@set_tts_adain_snake_cls) var s_i = cp |> cls_set(fixed_array(xip, xip, stp, hp, wp), @@set_tts_adain_leaky_cls) - static_assert(AD_EPS == TTS_ADAIN_EPS, "the CPU oracle's eps is the stamps' literal") - var pa = TtsAdainArgs(c = uint(c), t = uint(AD_T), xoff = uint(AD_XOFF), yoff = uint(AD_YOFF), hoff = uint(AD_HOFF), - gwoff = uint(AD_GWOFF), gboff = uint(f.gboff), aoff = uint(f.aoff)) - var pi = pa - pi.yoff = uint(AD_XOFF) + static_assert(AD_EPS == GK_ADAIN_EPS, "the CPU oracle's eps is the stamps' literal") + var pa = GkAdainArgs(c = uint(c), t = uint(AD_T), hoff = uint(AD_HOFF), gwoff = uint(AD_GWOFF), gboff = uint(f.gboff), aoff = uint(f.aoff)) cp |> cell_record() $(raw : VkCommandBuffer; var h : VkHaz) { colstats_enc(raw, h, cp, x0p, stp, AD_T, c) enc_tts_adain_leaky_cls(raw, h, s_l, pa, int64(n)) enc_tts_adain_snake_cls(raw, h, s_s, pa, int64(n)) - enc_tts_adain_leaky_cls(raw, h, s_i, pi, int64(n)) + enc_tts_adain_leaky_cls(raw, h, s_i, pa, int64(n)) } let got_l <- cp |> plane_run(ylp) let got_i <- cp |> plane_run(xip) @@ -578,9 +574,9 @@ def test_vkt_adain(t0 : T?) { check_rows(t, "tts_adain_leaky in place vs the CPU adain_rows_into and leaky_relu", got_i, f.want_leaky, BAR_APPROX, AD_C) t |> success(got_i[0] != f.x[0], "tts_adain_leaky in place control: the input row is overwritten") t |> equal(mismatch_exact(got_i, got_l), 0, "tts_adain_leaky: in place == out of place, bit for bit") - t |> success(outside_run(cp |> plane_whole(ylp), AD_YOFF, n, NAN_SENTINEL) && outside_run(cp |> plane_whole(ysp), AD_YOFF, n, NAN_SENTINEL), - "tts_adain: the prefix before yoff keeps its sentinel") - t |> success(outside_run(cp |> plane_whole(xip), AD_XOFF, n, POISON), "tts_adain in place: the prefix before xoff is untouched") + t |> success(outside_run(cp |> plane_whole(ylp), 0, n, NAN_SENTINEL) && outside_run(cp |> plane_whole(ysp), 0, n, NAN_SENTINEL), + "tts_adain: the elements past the run keep their sentinel") + t |> success(outside_run(cp |> plane_whole(xip), 0, n, POISON), "tts_adain in place: the elements past the run are untouched") } delete f } else { @@ -590,9 +586,6 @@ def test_vkt_adain(t0 : T?) { } let private AS_N = 1030 //! four whole workgroups and six more -let private AS_AOFF = 3 -let private AS_BOFF = 7 -let private AS_OOFF = 11 let private AS_SLACK = 5 //! elements past the run, left at the sentinel [test] @@ -608,26 +601,24 @@ def test_vkt_add_scale(t0 : T?) { let a <- row_fill(n, 41, 7, 0.05, 2.1) let b <- row_fill(n, 17, 3, 0.04, 1.9) with_cell_planes() $(var cp : CellPlanes) { - let ap = cp |> in_(a, AS_AOFF) - let aip = cp |> io_(a, AS_AOFF) - let bp = cp |> in_(b, AS_BOFF) - let op = cp |> out_(n, AS_OOFF, AS_SLACK) + let ap = cp |> in_(a, 0) + let aip = cp |> io_(a, 0, AS_SLACK) + let bp = cp |> in_(b, 0) + let op = cp |> out_(n, 0, AS_SLACK) var s_o = cp |> cls_set(fixed_array(ap, bp, op), @@set_tts_add_scale_cls) var s_i = cp |> cls_set(fixed_array(aip, bp, aip), @@set_tts_add_scale_cls) - var po = TtsAddScaleArgs(nelem = uint(n), scale = scale, aoff = uint(AS_AOFF), boff = uint(AS_BOFF), ooff = uint(AS_OOFF)) - var pi = po - pi.ooff = uint(AS_AOFF) + var po = GkAddScaleArgs(nelem = uint(n), scale = scale) cp |> cell_record() $(raw : VkCommandBuffer; var h : VkHaz) { enc_tts_add_scale_cls(raw, h, s_o, po, int64(n)) - enc_tts_add_scale_cls(raw, h, s_i, pi, int64(n)) + enc_tts_add_scale_cls(raw, h, s_i, po, int64(n)) } let got_i <- cp |> plane_run(aip) let want <- [for (i in range(n)); (a[i] + b[i]) * scale] check_rows(t, "tts_add_scale vs (a + b) * scale", cp |> plane_run(op), want, BAR_EXACT, FLAT_ROWS) check_rows(t, "tts_add_scale in place vs (a + b) * scale", got_i, want, BAR_EXACT, FLAT_ROWS) t |> success(got_i[0] != a[0], "tts_add_scale in place control: the input row is overwritten") - t |> success(outside_run(cp |> plane_whole(op), AS_OOFF, n, NAN_SENTINEL), "tts_add_scale: the elements outside the run keep their sentinel") - t |> success(outside_run(cp |> plane_whole(aip), AS_AOFF, n, POISON), "tts_add_scale in place: the prefix before aoff is untouched") + t |> success(outside_run(cp |> plane_whole(op), 0, n, NAN_SENTINEL), "tts_add_scale: the elements outside the run keep their sentinel") + t |> success(outside_run(cp |> plane_whole(aip), 0, n, POISON), "tts_add_scale in place: the elements past the run are untouched") } } else { t |> skip("dasVulkan not present") diff --git a/modules/dasLLAMA/tests/test_vulkan_tts_pocket_kernels.das b/modules/dasLLAMA/tests/test_vulkan_tts_pocket_kernels.das index 040300d973..4e20687efe 100644 --- a/modules/dasLLAMA/tests/test_vulkan_tts_pocket_kernels.das +++ b/modules/dasLLAMA/tests/test_vulkan_tts_pocket_kernels.das @@ -63,7 +63,7 @@ def test_vkt_pk_rows(t0 : T?) { let ep = cp |> out_(ndst) var s_p = cp |> cls_set(fixed_array(sp, dp), @@set_tts_pk_rows_cls) var s_e = cp |> cls_set(fixed_array(sp, ep), @@set_tts_pk_rows_elu_cls) - var pa = TtsPkRowsArgs(rows = uint(PR_ROWS), width = uint(PR_WIDTH), ss = uint(PR_SS), sc = uint(PR_SC), sr0 = uint(PR_SR0), + var pa = GkPkRowsArgs(rows = uint(PR_ROWS), width = uint(PR_WIDTH), ss = uint(PR_SS), sc = uint(PR_SC), sr0 = uint(PR_SR0), ds = uint(PR_DS), dc = uint(PR_DC), dr0 = uint(PR_DR0)) cp |> cell_record() $(raw : VkCommandBuffer; var h : VkHaz) { enc_tts_pk_rows_cls(raw, h, s_p, pa, int64(nwin)) @@ -142,10 +142,6 @@ let private PA_CAP = 300 let private PA_QS = 264 //! q's row stride: the query span sits at PA_QCOL among other columns let private PA_QCOL = 8 let private PA_PROWS = 48 //! the device caches' rows: the 45 appended, then three poisoned rows no query may see -let private PA_QOFF = 5 -let private PA_KOFF = 7 -let private PA_VOFF = 11 -let private PA_OOFF = 13 let private PA_DEC_POS = 44 //! the decode step's position: the last appended key let private PA_DEC_CTX = 8 let private PA_WIN_CTX = 16 @@ -207,26 +203,26 @@ def private pk_attn_device(kc : TtsKvCache; qfull : array) : PkAttnOut { var inscope rows <- pk_cache_rows(kc, -1l) var inscope rows_x <- pk_cache_rows(kc, int64(PA_DEC_POS - PA_DEC_CTX)) with_cell_planes() $(var cp : CellPlanes) { - let qp = cp |> in_(qfull, PA_QOFF, SLACK) - let kp = cp |> in_(rows.k, PA_KOFF, SLACK) - let vp = cp |> in_(rows.v, PA_VOFF, SLACK) - let kxp = cp |> in_(rows_x.k, PA_KOFF, SLACK) - let vxp = cp |> in_(rows_x.v, PA_VOFF, SLACK) - let oa = cp |> out_(PA_T * d, PA_OOFF) - let ob = cp |> out_(d, PA_OOFF) - let ox = cp |> out_(d, PA_OOFF) - let oc = cp |> out_(PA_T * d, PA_OOFF) + let qp = cp |> in_(qfull, 0, SLACK) + var inscope qdec <- copy_span(qfull, PA_DEC_POS * PA_QS, PA_QS) + let qdp = cp |> in_(qdec, 0, SLACK) + let kp = cp |> in_(rows.k, 0, SLACK) + let vp = cp |> in_(rows.v, 0, SLACK) + let kxp = cp |> in_(rows_x.k, 0, SLACK) + let vxp = cp |> in_(rows_x.v, 0, SLACK) + let oa = cp |> out_(PA_T * d) + let ob = cp |> out_(d) + let ox = cp |> out_(d) + let oc = cp |> out_(PA_T * d) var s_a = cp |> cls_set(fixed_array(qp, kp, vp, oa), @@set_tts_pk_attn_cls) - var s_b = cp |> cls_set(fixed_array(qp, kp, vp, ob), @@set_tts_pk_attn_cls) - var s_x = cp |> cls_set(fixed_array(qp, kxp, vxp, ox), @@set_tts_pk_attn_cls) + var s_b = cp |> cls_set(fixed_array(qdp, kp, vp, ob), @@set_tts_pk_attn_cls) + var s_x = cp |> cls_set(fixed_array(qdp, kxp, vxp, ox), @@set_tts_pk_attn_cls) var s_c = cp |> cls_set(fixed_array(qp, kp, vp, oc), @@set_tts_pk_attn_cls) - var pa = TtsPkAttnArgs(t = uint(PA_T), qs = uint(PA_QS), qcol = uint(PA_QCOL), d = uint(d), pos0 = 0u, ctx = 0u, - qoff = uint(PA_QOFF), koff = uint(PA_KOFF), voff = uint(PA_VOFF), ooff = uint(PA_OOFF)) + var pa = GkPkAttnArgs(t = uint(PA_T), qs = uint(PA_QS), qcol = uint(PA_QCOL), d = uint(d), pos0 = 0u, ctx = 0u) var pb = pa pb.t = 1u pb.pos0 = uint(PA_DEC_POS) pb.ctx = uint(PA_DEC_CTX) - pb.qoff = uint(PA_QOFF + PA_DEC_POS * PA_QS) var pc = pa pc.ctx = uint(PA_WIN_CTX) let heads = int64(PA_HEADS) @@ -271,7 +267,7 @@ def test_vkt_pk_attn(t0 : T?) { } } t |> success(moved > 0, "tts_pk_attn control: the 16-key window moves the prompt's late rows off the every-key rows ({moved} apart)") - t |> success(outside_run(got.a_whole, PA_OOFF, PA_T * d, NAN_SENTINEL), "tts_pk_attn: the prefix before ooff keeps its sentinel") + t |> success(outside_run(got.a_whole, 0, PA_T * d, NAN_SENTINEL), "tts_pk_attn: the elements past the rows keep their sentinel") } else { t |> skip("dasVulkan not present") } @@ -479,7 +475,7 @@ def test_vkt_pk_rope(t0 : T?) { let xp = cp |> io_(x, 0, SLACK) let wp = cp |> in_(slab) var s = cp |> cls_set(fixed_array(xp, wp, wp), @@set_tts_pk_rope_cls) - var pa = RopeTabArgs(rows = uint(PO_T), qs = uint(PO_QS), col = uint(PO_D), d = uint(PO_D), dh = uint(PO_DH), coff = uint(coff), soff = uint(soff)) + var pa = GkRopeTabArgs(rows = uint(PO_T), qs = uint(PO_QS), col = uint(PO_D), d = uint(PO_D), dh = uint(PO_DH), coff = uint(coff), soff = uint(soff)) cp |> cell_record() $(raw : VkCommandBuffer; var h : VkHaz) { enc_tts_pk_rope_cls(raw, h, s, pa, int64(nk / 2)) } diff --git a/modules/dasLLAMA/tests/test_vulkan_tts_source_kernels.das b/modules/dasLLAMA/tests/test_vulkan_tts_source_kernels.das index 18adc54436..e071336b2b 100644 --- a/modules/dasLLAMA/tests/test_vulkan_tts_source_kernels.das +++ b/modules/dasLLAMA/tests/test_vulkan_tts_source_kernels.das @@ -85,8 +85,8 @@ def test_vkt_reflect1(t0 : T?) { let xp = cp |> in_(x) let yp = cp |> out_(total, 0, SLACK) var s = cp |> cls_set(fixed_array(xp, yp), @@set_tts_pk_rows_cls) - var shift = TtsPkRowsArgs(rows = uint(RF_T), width = uint(RF_C), ss = uint(RF_C), sc = 0u, sr0 = 0u, ds = uint(RF_C), dc = 0u, dr0 = 1u) - var head = TtsPkRowsArgs(rows = 1u, width = uint(RF_C), ss = uint(RF_C), sc = 0u, sr0 = 1u, ds = uint(RF_C), dc = 0u, dr0 = 0u) + var shift = GkPkRowsArgs(rows = uint(RF_T), width = uint(RF_C), ss = uint(RF_C), sc = 0u, sr0 = 0u, ds = uint(RF_C), dc = 0u, dr0 = 1u) + var head = GkPkRowsArgs(rows = 1u, width = uint(RF_C), ss = uint(RF_C), sc = 0u, sr0 = 1u, ds = uint(RF_C), dc = 0u, dr0 = 0u) cp |> cell_record() $(raw : VkCommandBuffer; var h : VkHaz) { enc_tts_pk_rows_cls(raw, h, s, shift, int64(RF_T * RF_C)) enc_tts_pk_rows_cls(raw, h, s, head, int64(RF_C)) @@ -147,7 +147,7 @@ def private src_run(f : TtsSrcFixture; torch : bool; f0 : array; noise : let lowp = cp |> out_(nlow) let lowcp = cp |> out_(nlow) let harp = cp |> out_(f.n, 0, SLACK) - var pa = TtsSourceArgs(n = uint(f.n), nlow = uint(f.nlow), nh = uint(TTS_SRC_NH), up = uint(f.up), sr = TTS_SRC_SR, sine_amp = 0.1, + var pa = GkSourceArgs(n = uint(f.n), nlow = uint(f.nlow), nh = uint(TTS_SRC_NH), up = uint(f.up), sr = TTS_SRC_SR, sine_amp = 0.1, noise_std = 0.003, voiced_thr = 10.0, seed = seed, woff = uint(SRC_WOFF), boff = uint(src_boff())) var s_nz = cp |> cls_set(fixed_array(nzp), @@set_tts_src_noise_cls) let low_ix = fixed_array(f0p, nup, lowp) @@ -197,7 +197,7 @@ def private long_cumsum_arm(t : T?; tag : string; torch : bool; up : int) { let cp_out = cp |> out_(nlong) let ix = fixed_array(lp, cp_out) var s = torch ? cp |> cls_set(ix, @@set_tts_src_cumsum_torch_cls) : cp |> cls_set(ix, @@set_tts_src_cumsum_onnx_cls) - var pa = TtsSourceArgs(n = uint(nlong * up), nlow = uint(nlong), nh = 1u, up = uint(up)) + var pa = GkSourceArgs(n = uint(nlong * up), nlow = uint(nlong), nh = 1u, up = uint(up)) cp |> cell_record() $(raw : VkCommandBuffer; var h : VkHaz) { if (torch) { enc_tts_src_cumsum_torch_cls(raw, h, s, pa, 1l) @@ -360,7 +360,7 @@ def private stft_arm(t : T?; reflect : bool) { let spp = cp |> out_(nspec, 0, SLACK) var s = reflect ? cp |> cls_set(fixed_array(hp, wp, sp), @@set_tts_stft_reflect_cls) : cp |> cls_set(fixed_array(hp, wp, sp), @@set_tts_stft_edge_cls) var sq = reflect ? cp |> cls_set(fixed_array(hpp, wp, spp), @@set_tts_stft_reflect_cls) : cp |> cls_set(fixed_array(hpp, wp, spp), @@set_tts_stft_edge_cls) - var pa = TtsStftArgs(n = uint(ST_N), frames = uint(frames), bins = uint(ST_BINS), k = uint(ST_K), hop = uint(ST_HOP), pad = uint(ST_PAD), + var pa = GkStftArgs(n = uint(ST_N), frames = uint(frames), bins = uint(ST_BINS), k = uint(ST_K), hop = uint(ST_HOP), pad = uint(ST_PAD), eps = 1e-9, reoff = uint(ST_REOFF), imoff = uint(imoff)) let total = int64(frames * ST_BINS) cp |> cell_record() $(raw : VkCommandBuffer; var h : VkHaz) { @@ -440,7 +440,7 @@ def private istft_arm(t : T?; envelope : bool) { let opq = cp |> out_(n, 0, SLACK) var s = cp |> cls_set(fixed_array(yq, wq, oq), @@set_tts_istft_cls) var sq = cp |> cls_set(fixed_array(ypq, wq, opq), @@set_tts_istft_cls) - var pa = TtsIstftArgs(tt = uint(IS_TT), stride = uint(IS_STRIDE), bins = uint(IS_BINS), k = uint(IS_K), hop = uint(IS_HOP), + var pa = GkIstftArgs(tt = uint(IS_TT), stride = uint(IS_STRIDE), bins = uint(IS_BINS), k = uint(IS_K), hop = uint(IS_HOP), pad = uint(IS_PAD), envelope = envelope ? 1u : 0u, n = uint(n), reoff = uint(IS_REOFF), imoff = uint(imoff), wnoff = uint(wnoff)) cp |> cell_record() $(raw : VkCommandBuffer; var h : VkHaz) { enc_tts_istft_cls(raw, h, s, pa, int64(n))