spark_model/layers/ops/
ssm_gdn_a2.rs

1// SPDX-License-Identifier: AGPL-3.0-only
2
3//! GDN prefill ops — extracted from `ssm_gdn_a.rs` during the ≤500-line split.
4//! All public items remain available at `crate::layers::ops::*` via the
5//! re-export in `ops.rs`.
6
7#![allow(unused_imports)]
8
9use anyhow::Result;
10use spark_runtime::gpu::{DevicePtr, GpuBackend, KernelHandle};
11use spark_runtime::kernel_args::{KernelLaunch, div_ceil};
12
13use crate::layers::moe;
14use crate::weight_map::{DenseWeight, Fp8DenseWeight, Fp8Weight, QuantizedWeight};
15
16use super::*;
17
18/// Gated delta rule prefill (multi-token, sequential SSM update within kernel).
19///
20/// Processes `seq_len` tokens sequentially per (batch, head) pair.
21/// Supports strided access: Q/K/V/gate/beta may have different strides
22/// between tokens (e.g., from conv1d output with interleaved Q|K|V layout).
23///
24/// Kernel: `gated_delta_rule_prefill(h_state, query, key, value,
25///          gate, beta, output, batch_size, seq_len, num_k_heads,
26///          num_v_heads, k_dim, v_dim, qk_stride, v_stride, gb_stride)`
27/// Grid: (num_v_heads, batch, 1)  Block: (128, 1, 1)
28#[allow(clippy::too_many_arguments)]
29pub fn gdn_prefill(
30    gpu: &dyn GpuBackend,
31    kernel: KernelHandle,
32    h_state: DevicePtr,
33    query: DevicePtr,
34    key: DevicePtr,
35    value: DevicePtr,
36    gate: DevicePtr,
37    beta: DevicePtr,
38    output: DevicePtr,
39    batch_size: u32,
40    seq_len: u32,
41    num_k_heads: u32,
42    num_v_heads: u32,
43    k_dim: u32,
44    v_dim: u32,
45    qk_stride: u32,
46    v_stride: u32,
47    gb_stride: u32,
48    stream: u64,
49) -> Result<()> {
50    KernelLaunch::new(gpu, kernel)
51        .grid([num_v_heads, batch_size, 1])
52        .block([128, 1, 1])
53        .shared_mem(4 * k_dim * 4) // double-buffered k[128]+q[128] × 2 buffers × 4 bytes
54        .arg_ptr(h_state)
55        .arg_ptr(query)
56        .arg_ptr(key)
57        .arg_ptr(value)
58        .arg_ptr(gate)
59        .arg_ptr(beta)
60        .arg_ptr(output)
61        .arg_u32(batch_size)
62        .arg_u32(seq_len)
63        .arg_u32(num_k_heads)
64        .arg_u32(num_v_heads)
65        .arg_u32(k_dim)
66        .arg_u32(v_dim)
67        .arg_u32(qk_stride)
68        .arg_u32(v_stride)
69        .arg_u32(gb_stride)
70        .launch(stream)
71}
72
73/// Split-v_dim prefill: 2 CTAs per v-head, 64 threads each.
74///
75/// Kernel: `gated_delta_rule_prefill_split(h_state, query, key, value,
76///          gate, beta, output, batch_size, seq_len, num_k_heads,
77///          num_v_heads, k_dim, v_dim, qk_stride, v_stride, gb_stride)`
78/// Grid: (num_v_heads * 2, batch, 1)  Block: (64, 1, 1)
79#[allow(clippy::too_many_arguments)]
80pub fn gdn_prefill_split(
81    gpu: &dyn GpuBackend,
82    kernel: KernelHandle,
83    h_state: DevicePtr,
84    query: DevicePtr,
85    key: DevicePtr,
86    value: DevicePtr,
87    gate: DevicePtr,
88    beta: DevicePtr,
89    output: DevicePtr,
90    batch_size: u32,
91    seq_len: u32,
92    num_k_heads: u32,
93    num_v_heads: u32,
94    k_dim: u32,
95    v_dim: u32,
96    qk_stride: u32,
97    v_stride: u32,
98    gb_stride: u32,
99    stream: u64,
100) -> Result<()> {
101    KernelLaunch::new(gpu, kernel)
102        .grid([num_v_heads * 2, batch_size, 1])
103        .block([64, 1, 1])
104        .shared_mem(4 * k_dim * 4) // double-buffered k[K_DIM]+q[K_DIM] × 2 buffers × 4 bytes
105        .arg_ptr(h_state)
106        .arg_ptr(query)
107        .arg_ptr(key)
108        .arg_ptr(value)
109        .arg_ptr(gate)
110        .arg_ptr(beta)
111        .arg_ptr(output)
112        .arg_u32(batch_size)
113        .arg_u32(seq_len)
114        .arg_u32(num_k_heads)
115        .arg_u32(num_v_heads)
116        .arg_u32(k_dim)
117        .arg_u32(v_dim)
118        .arg_u32(qk_stride)
119        .arg_u32(v_stride)
120        .arg_u32(gb_stride)
121        .launch(stream)
122}
123
124/// 4-way split prefill: 4 CTAs per v-head, 32 threads each (128 total CTAs).
125///
126/// Kernel: `gated_delta_rule_prefill_split4(h_state, query, key, value,
127///          gate, beta, output, batch_size, seq_len, num_k_heads,
128///          num_v_heads, k_dim, v_dim, qk_stride, v_stride, gb_stride)`
129/// Grid: (num_v_heads * 4, batch, 1)  Block: (32, 1, 1)
130#[allow(clippy::too_many_arguments)]
131pub fn gdn_prefill_split4(
132    gpu: &dyn GpuBackend,
133    kernel: KernelHandle,
134    h_state: DevicePtr,
135    query: DevicePtr,
136    key: DevicePtr,
137    value: DevicePtr,
138    gate: DevicePtr,
139    beta: DevicePtr,
140    output: DevicePtr,
141    batch_size: u32,
142    seq_len: u32,
143    num_k_heads: u32,
144    num_v_heads: u32,
145    k_dim: u32,
146    v_dim: u32,
147    qk_stride: u32,
148    v_stride: u32,
149    gb_stride: u32,
150    stream: u64,
151) -> Result<()> {
152    KernelLaunch::new(gpu, kernel)
153        .grid([num_v_heads * 4, batch_size, 1])
154        .block([32, 1, 1])
155        .shared_mem(4 * k_dim * 4) // double-buffered k[K_DIM]+q[K_DIM] × 2 buffers × 4 bytes
156        .arg_ptr(h_state)
157        .arg_ptr(query)
158        .arg_ptr(key)
159        .arg_ptr(value)
160        .arg_ptr(gate)
161        .arg_ptr(beta)
162        .arg_ptr(output)
163        .arg_u32(batch_size)
164        .arg_u32(seq_len)
165        .arg_u32(num_k_heads)
166        .arg_u32(num_v_heads)
167        .arg_u32(k_dim)
168        .arg_u32(v_dim)
169        .arg_u32(qk_stride)
170        .arg_u32(v_stride)
171        .arg_u32(gb_stride)
172        .launch(stream)
173}
174
175/// Persistent GDN prefill — h_state stays in shared memory for entire sequence.
176///
177/// Same parameters as gdn_prefill_split4 but uses persistent CTAs with
178/// 128 threads and 67 KB shared memory. Each CTA processes ALL tokens for
179/// one v_head, keeping h_state in shared memory (never written to global
180/// until the end). Targets L2 bandwidth (~3 TB/s) instead of LPDDR5X (273 GB/s).
181///
182/// Grid: (num_v_heads, batch, 1)  Block: (128, 1, 1)
183/// Shared: k_dim*v_dim*4 + 4*k_dim*4 bytes
184#[allow(clippy::too_many_arguments)]
185pub fn gdn_prefill_persistent(
186    gpu: &dyn GpuBackend,
187    kernel: KernelHandle,
188    h_state: DevicePtr,
189    query: DevicePtr,
190    key: DevicePtr,
191    value: DevicePtr,
192    gate: DevicePtr,
193    beta: DevicePtr,
194    output: DevicePtr,
195    batch_size: u32,
196    seq_len: u32,
197    num_k_heads: u32,
198    num_v_heads: u32,
199    k_dim: u32,
200    v_dim: u32,
201    qk_stride: u32,
202    v_stride: u32,
203    gb_stride: u32,
204    stream: u64,
205) -> Result<()> {
206    let smem = k_dim * v_dim * 4 + 4 * k_dim * 4; // h_state + double-buffered k/q
207    KernelLaunch::new(gpu, kernel)
208        .grid([num_v_heads, batch_size, 1])
209        .block([128, 1, 1])
210        .shared_mem(smem)
211        .arg_ptr(h_state)
212        .arg_ptr(query)
213        .arg_ptr(key)
214        .arg_ptr(value)
215        .arg_ptr(gate)
216        .arg_ptr(beta)
217        .arg_ptr(output)
218        .arg_u32(batch_size)
219        .arg_u32(seq_len)
220        .arg_u32(num_k_heads)
221        .arg_u32(num_v_heads)
222        .arg_u32(k_dim)
223        .arg_u32(v_dim)
224        .arg_u32(qk_stride)
225        .arg_u32(v_stride)
226        .arg_u32(gb_stride)
227        .launch(stream)
228}
229
230/// Persistent GDN prefill with explicit shared memory size.
231/// Used for WY4-persistent variant which needs more shared memory.
232#[allow(clippy::too_many_arguments)]
233pub fn gdn_prefill_persistent_smem(
234    gpu: &dyn GpuBackend,
235    kernel: KernelHandle,
236    h_state: DevicePtr,
237    query: DevicePtr,
238    key: DevicePtr,
239    value: DevicePtr,
240    gate: DevicePtr,
241    beta: DevicePtr,
242    output: DevicePtr,
243    batch_size: u32,
244    seq_len: u32,
245    num_k_heads: u32,
246    num_v_heads: u32,
247    k_dim: u32,
248    v_dim: u32,
249    qk_stride: u32,
250    v_stride: u32,
251    gb_stride: u32,
252    smem: u32,
253    stream: u64,
254) -> Result<()> {
255    KernelLaunch::new(gpu, kernel)
256        .grid([num_v_heads, batch_size, 1])
257        .block([128, 1, 1])
258        .shared_mem(smem)
259        .arg_ptr(h_state)
260        .arg_ptr(query)
261        .arg_ptr(key)
262        .arg_ptr(value)
263        .arg_ptr(gate)
264        .arg_ptr(beta)
265        .arg_ptr(output)
266        .arg_u32(batch_size)
267        .arg_u32(seq_len)
268        .arg_u32(num_k_heads)
269        .arg_u32(num_v_heads)
270        .arg_u32(k_dim)
271        .arg_u32(v_dim)
272        .arg_u32(qk_stride)
273        .arg_u32(v_stride)
274        .arg_u32(gb_stride)
275        .launch(stream)
276}