Skip to main content

entrenar/autograd/cuda_forward/
matmul.rs

1#![allow(unsafe_code)]
2#![allow(trivial_casts)]
3#![allow(clippy::borrow_as_ptr)]
4#![allow(clippy::ref_as_ptr)]
5
6#[cfg(feature = "cuda")]
7use trueno_gpu::driver::{CublasHandle, CudaStream, GemmOp, GpuBuffer, LaunchConfig};
8#[cfg(feature = "cuda")]
9use trueno_gpu::kernels::{
10    Batched4DGemmKernel, FusedSwigluKernel, GemmKernel, Kernel, Nf4GemmKernel,
11    Nf4GemmTransposeKernel, Nf4TensorCoreGemmKernel,
12};
13
14use crate::autograd::cuda_tensor::{CudaTensorError, Result};
15
16#[cfg(feature = "cuda")]
17use super::cache::FORWARD_KERNEL_CACHE;
18
19/// Bind a cuBLAS handle to the caller's stream before dispatching a GEMM.
20///
21/// FALSIFY-CUDA-NF4-FORWARD-NAN-001 (cuda-nf4-forward-stream-ordering-v1):
22/// a fresh cuBLAS handle issues its GEMMs on the legacy DEFAULT stream, and a
23/// `CU_STREAM_NON_BLOCKING` trainer stream never implicitly synchronizes with
24/// it — so an unbound GEMM races the PTX producer/consumer kernels around it
25/// (rmsnorm→QKV, softmax→scores@V, final-norm→lm_head), yielding NaN losses
26/// on the FIRST QLoRA training step. Binding per call (instead of bind-once
27/// at trainer construction) also keeps the process-global handle off
28/// DESTROYED streams after a trainer drops. `cublasSetStream_v2` is a handle
29/// field write (~100ns) — negligible next to any GEMM. Callers hold the
30/// kernel-cache mutex, so bind+launch is atomic across threads.
31#[cfg(feature = "cuda")]
32pub(crate) fn bind_cublas_stream(cublas: &CublasHandle, stream: &CudaStream) -> Result<()> {
33    cublas
34        .set_stream(stream)
35        .map_err(|e| CudaTensorError::KernelError(format!("cuBLAS set_stream failed: {e:?}")))
36}
37
38/// Fused SwiGLU forward pass on GPU (ENT-150)
39///
40/// Computes: output = SiLU(gate) * up
41/// Fuses two operations into one kernel for better memory bandwidth.
42#[cfg(feature = "cuda")]
43pub fn fused_swiglu_forward(
44    gate: &GpuBuffer<f32>,
45    up: &GpuBuffer<f32>,
46    output: &mut GpuBuffer<f32>,
47    n: u32,
48    stream: &CudaStream,
49) -> Result<()> {
50    let cache = FORWARD_KERNEL_CACHE.get().ok_or(CudaTensorError::DeviceNotInitialized)?;
51    let mut cache = cache.lock().map_err(|_err| {
52        CudaTensorError::KernelError("Failed to acquire kernel cache lock".to_string())
53    })?;
54
55    let key = "fused_swiglu_forward".to_string(); // PTX is n-independent (trueno#184)
56    let module = match cache.get_cached(&key) {
57        Some(m) => m,
58        None => {
59            let kernel = FusedSwigluKernel::new(n);
60            let ptx = kernel.emit_ptx_for_target(cache.sm_target());
61            cache.get_or_compile(&key, &ptx)?
62        }
63    };
64
65    let config = LaunchConfig { grid: (n.div_ceil(256), 1, 1), block: (256, 1, 1), shared_mem: 0 };
66
67    let gate_ptr = gate.as_ptr();
68    let up_ptr = up.as_ptr();
69    let output_ptr = output.as_ptr();
70
71    let mut args: [*mut std::ffi::c_void; 4] = [
72        &gate_ptr as *const _ as *mut _,
73        &up_ptr as *const _ as *mut _,
74        &output_ptr as *const _ as *mut _,
75        &n as *const _ as *mut _,
76    ];
77
78    // SAFETY: Kernel launch requires FFI. All buffers are valid GPU allocations with
79    // matching sizes, and the kernel parameters match the expected PTX signature.
80    unsafe {
81        stream.launch_kernel(module, "fused_swiglu", &config, &mut args).map_err(|e| {
82            CudaTensorError::KernelError(format!("Fused SwiGLU forward launch failed: {e:?}"))
83        })?;
84    }
85
86    Ok(())
87}
88
89/// GEMM forward pass on GPU
90///
91/// Computes: C = A @ B where A is MxK, B is KxN, C is MxN
92///
93/// Dispatches to cuBLAS tensor core TF32 when available (ALB-075), falling back to PTX
94/// naive GEMM. Backward GEMMs use CUBLAS_DEFAULT_MATH (SIMD) per ALB-076/trueno#170.
95#[cfg(feature = "cuda")]
96pub fn gemm_forward(
97    a: &GpuBuffer<f32>,
98    b: &GpuBuffer<f32>,
99    c: &mut GpuBuffer<f32>,
100    m: u32,
101    k: u32,
102    n: u32,
103    stream: &CudaStream,
104) -> Result<()> {
105    let cache = FORWARD_KERNEL_CACHE.get().ok_or(CudaTensorError::DeviceNotInitialized)?;
106    let mut cache = cache.lock().map_err(|_err| {
107        CudaTensorError::KernelError("Failed to acquire kernel cache lock".to_string())
108    })?;
109    if let Some(cublas) = cache.cublas() {
110        bind_cublas_stream(cublas, stream)?;
111        return cublas_gemm_forward(cublas, a, b, c, m, k, n);
112    }
113
114    // PTX fallback
115    let key = format!("gemm_forward_{m}_{k}_{n}");
116    let module = match cache.get_cached(&key) {
117        Some(m) => m,
118        None => {
119            let kernel = GemmKernel::naive(m, n, k);
120            let ptx = kernel.emit_ptx_for_target(cache.sm_target());
121            cache.get_or_compile(&key, &ptx)?
122        }
123    };
124
125    // Use 16x16 thread blocks for GEMM
126    // Kernel: col = ctaid.x * 16 + tid.x, row = ctaid.y * 16 + tid.y
127    // So grid.x = ceil(N/16) for columns, grid.y = ceil(M/16) for rows
128    let config = LaunchConfig {
129        grid: (n.div_ceil(16), m.div_ceil(16), 1),
130        block: (16, 16, 1),
131        shared_mem: 0,
132    };
133
134    let a_ptr = a.as_ptr();
135    let b_ptr = b.as_ptr();
136    let c_ptr = c.as_ptr();
137
138    // PTX kernel signature: (a_ptr, b_ptr, c_ptr, m, n, k)
139    // CRITICAL: must match param declaration order in GemmKernel::build_naive()
140    let mut args: [*mut std::ffi::c_void; 6] = [
141        &a_ptr as *const _ as *mut _,
142        &b_ptr as *const _ as *mut _,
143        &c_ptr as *const _ as *mut _,
144        &m as *const _ as *mut _,
145        &n as *const _ as *mut _,
146        &k as *const _ as *mut _,
147    ];
148
149    // SAFETY: Kernel launch requires FFI. All buffers are valid GPU allocations with
150    // matching sizes, and the kernel parameters match the expected PTX signature.
151    unsafe {
152        stream.launch_kernel(module, "gemm_naive", &config, &mut args).map_err(|e| {
153            CudaTensorError::KernelError(format!("GEMM forward launch failed: {e:?}"))
154        })?;
155    }
156
157    Ok(())
158}
159
160/// GEMM with transposed B: C[M,N] = A[M,K] @ B[N,K]^T
161/// B is stored row-major [N,K]. entrenar#318: GPU lm_head with embed_original.
162#[cfg(feature = "cuda")]
163pub fn gemm_forward_bt(
164    a: &GpuBuffer<f32>,
165    b: &GpuBuffer<f32>,
166    c: &mut GpuBuffer<f32>,
167    m: u32,
168    k: u32,
169    n: u32,
170    stream: &CudaStream,
171) -> Result<()> {
172    let cache = FORWARD_KERNEL_CACHE.get().ok_or(CudaTensorError::DeviceNotInitialized)?;
173    let cache = cache.lock().map_err(|_| CudaTensorError::KernelError("cache lock".to_string()))?;
174    if let Some(cublas) = cache.cublas() {
175        bind_cublas_stream(cublas, stream)?;
176        return cublas_gemm_forward_bt(cublas, a, b, c, m, k, n);
177    }
178    Err(CudaTensorError::KernelError("gemm_forward_bt requires cuBLAS".to_string()))
179}
180
181#[cfg(feature = "cuda")]
182fn cublas_gemm_forward_bt(
183    cublas: &CublasHandle,
184    a: &GpuBuffer<f32>,
185    b: &GpuBuffer<f32>,
186    c: &mut GpuBuffer<f32>,
187    m: u32,
188    k: u32,
189    n: u32,
190) -> Result<()> {
191    // Row-major C[M,N] = A[M,K] @ B[N,K]^T
192    // Column-major: C^T[N,M] = Trans(B_col[K,N])[N,K] @ A_col[K,M]
193    cublas
194        .gemm_f32(
195            GemmOp::Trans,   // B transposed
196            GemmOp::NoTrans, // A not transposed
197            n as i32,
198            m as i32,
199            k as i32,
200            1.0,
201            b.as_ptr(),
202            k as i32, // ldb = K (B is [K,N] in col-major, transposed to [N,K])
203            a.as_ptr(),
204            k as i32, // lda = K (A is [K,M] in col-major)
205            0.0,
206            c.as_ptr(),
207            n as i32, // ldc = N
208        )
209        .map_err(|e| CudaTensorError::KernelError(format!("cuBLAS GEMM BT failed: {e:?}")))
210}
211
212/// cuBLAS GEMM forward: C[M,N] = A[M,K] @ B[K,N] (row-major via B^T@A^T identity)
213#[cfg(feature = "cuda")]
214fn cublas_gemm_forward(
215    cublas: &CublasHandle,
216    a: &GpuBuffer<f32>,
217    b: &GpuBuffer<f32>,
218    c: &mut GpuBuffer<f32>,
219    m: u32,
220    k: u32,
221    n: u32,
222) -> Result<()> {
223    cublas
224        .gemm_f32(
225            GemmOp::NoTrans,
226            GemmOp::NoTrans,
227            n as i32,
228            m as i32,
229            k as i32,
230            1.0,
231            b.as_ptr(),
232            n as i32,
233            a.as_ptr(),
234            k as i32,
235            0.0,
236            c.as_ptr(),
237            n as i32,
238        )
239        .map_err(|e| CudaTensorError::KernelError(format!("cuBLAS GEMM forward failed: {e:?}")))
240}
241
242/// cuBLAS backward A: grad_A[M,K] = grad_C[M,N] @ B[K,N]^T
243#[cfg(feature = "cuda")]
244pub(crate) fn cublas_gemm_backward_a(
245    cublas: &CublasHandle,
246    grad_output: &GpuBuffer<f32>,
247    b: &GpuBuffer<f32>,
248    grad_a: &mut GpuBuffer<f32>,
249    m: u32,
250    k: u32,
251    n: u32,
252) -> Result<()> {
253    cublas
254        .gemm_f32(
255            GemmOp::Trans,
256            GemmOp::NoTrans,
257            k as i32,
258            m as i32,
259            n as i32,
260            1.0,
261            b.as_ptr(),
262            n as i32,
263            grad_output.as_ptr(),
264            n as i32,
265            0.0,
266            grad_a.as_ptr(),
267            k as i32,
268        )
269        .map_err(|e| CudaTensorError::KernelError(format!("cuBLAS GEMM backward_a failed: {e:?}")))
270}
271
272/// cuBLAS backward A with accumulation: grad_A += grad_C @ B^T (PMAT-484)
273///
274/// Same as cublas_gemm_backward_a but uses beta=1.0 to ACCUMULATE into grad_a
275/// instead of overwriting. Enables fused Gate+Up backward without a separate
276/// cuda_add_inplace call.
277#[cfg(feature = "cuda")]
278pub(crate) fn cublas_gemm_backward_a_accumulate(
279    cublas: &CublasHandle,
280    grad_output: &GpuBuffer<f32>,
281    b: &GpuBuffer<f32>,
282    grad_a: &mut GpuBuffer<f32>,
283    m: u32,
284    k: u32,
285    n: u32,
286) -> Result<()> {
287    cublas
288        .gemm_f32(
289            GemmOp::Trans,
290            GemmOp::NoTrans,
291            k as i32,
292            m as i32,
293            n as i32,
294            1.0,
295            b.as_ptr(),
296            n as i32,
297            grad_output.as_ptr(),
298            n as i32,
299            1.0, // ACCUMULATE: C = 1.0 * A @ B + 1.0 * C
300            grad_a.as_ptr(),
301            k as i32,
302        )
303        .map_err(|e| {
304            CudaTensorError::KernelError(format!("cuBLAS GEMM backward_a accumulate failed: {e:?}"))
305        })
306}
307
308/// cuBLAS backward B: grad_B[K,N] = A[M,K]^T @ grad_C[M,N]
309#[cfg(feature = "cuda")]
310pub(crate) fn cublas_gemm_backward_b(
311    cublas: &CublasHandle,
312    a: &GpuBuffer<f32>,
313    grad_output: &GpuBuffer<f32>,
314    grad_b: &mut GpuBuffer<f32>,
315    m: u32,
316    k: u32,
317    n: u32,
318) -> Result<()> {
319    cublas
320        .gemm_f32(
321            GemmOp::NoTrans,
322            GemmOp::Trans,
323            n as i32,
324            k as i32,
325            m as i32,
326            1.0,
327            grad_output.as_ptr(),
328            n as i32,
329            a.as_ptr(),
330            k as i32,
331            0.0,
332            grad_b.as_ptr(),
333            n as i32,
334        )
335        .map_err(|e| CudaTensorError::KernelError(format!("cuBLAS GEMM backward_b failed: {e:?}")))
336}
337
338/// Batched 4D GEMM forward pass on GPU for multi-head attention
339///
340/// Computes: C[b,h] = A[b,h] @ B[b,h] for each batch b and head h
341/// Pattern: [batch, heads, m, k] @ [batch, heads, k, n] -> [batch, heads, m, n]
342///
343/// # Contract (C-B4DGEMM-001)
344///
345/// - **Precondition**: a.len() >= batch * heads * m * k, b.len() >= batch * heads * k * n,
346///   c.len() >= batch * heads * m * n
347/// - **Postcondition**: C[b,h] = A[b,h] @ B[b,h] for all (b,h) in [0,batch)×[0,heads)
348/// - **Invariant**: Zero CPU-side data transfers
349#[cfg(feature = "cuda")]
350pub fn batched_4d_gemm_forward(
351    a: &GpuBuffer<f32>,
352    b: &GpuBuffer<f32>,
353    c: &mut GpuBuffer<f32>,
354    batch: u32,
355    heads: u32,
356    m: u32,
357    n: u32,
358    k: u32,
359    stream: &CudaStream,
360) -> Result<()> {
361    let cache = FORWARD_KERNEL_CACHE.get().ok_or(CudaTensorError::DeviceNotInitialized)?;
362    let mut cache = cache.lock().map_err(|_err| {
363        CudaTensorError::KernelError("Failed to acquire kernel cache lock".to_string())
364    })?;
365
366    // ALB-075 Phase 4: cuBLAS strided batched GEMM for attention (16x faster than PTX)
367    if let Some(cublas) = cache.cublas() {
368        bind_cublas_stream(cublas, stream)?;
369        let batch_count = (batch * heads) as i32;
370        let stride_a = i64::from(m) * i64::from(k);
371        let stride_b = i64::from(k) * i64::from(n);
372        let stride_c = i64::from(m) * i64::from(n);
373        return cublas
374            .gemm_f32_strided_batched_row_major(
375                m as i32,
376                n as i32,
377                k as i32,
378                1.0,
379                a.as_ptr(),
380                stride_a,
381                b.as_ptr(),
382                stride_b,
383                0.0,
384                c.as_ptr(),
385                stride_c,
386                batch_count,
387            )
388            .map_err(|e| {
389                CudaTensorError::KernelError(format!("cuBLAS batched 4D GEMM failed: {e:?}"))
390            });
391    }
392
393    let kernel = Batched4DGemmKernel::new(batch, heads, m, n, k);
394    let tile_size = kernel.config.tile_size;
395
396    let key = format!("batched_4d_gemm_{batch}_{heads}_{m}_{n}_{k}");
397    let module = match cache.get_cached(&key) {
398        Some(m) => m,
399        None => {
400            let ptx = kernel.emit_ptx_for_target(cache.sm_target());
401            cache.get_or_compile(&key, &ptx)?
402        }
403    };
404
405    // Grid: ((m+tile-1)/tile, (n+tile-1)/tile, batch * heads)
406    // Block: (tile_size, tile_size, 1)
407    // Shared memory: tile_size * tile_size * 4 * 2 bytes (tiles for A and B)
408    let config = LaunchConfig {
409        grid: (n.div_ceil(tile_size), m.div_ceil(tile_size), batch * heads),
410        block: (tile_size, tile_size, 1),
411        shared_mem: tile_size * tile_size * 4 * 2,
412    };
413
414    let a_ptr = a.as_ptr();
415    let b_ptr = b.as_ptr();
416    let c_ptr = c.as_ptr();
417
418    // PTX kernel signature: (a_ptr, b_ptr, c_ptr, batch, heads, m, n, k)
419    let mut args: [*mut std::ffi::c_void; 8] = [
420        &a_ptr as *const _ as *mut _,
421        &b_ptr as *const _ as *mut _,
422        &c_ptr as *const _ as *mut _,
423        &batch as *const _ as *mut _,
424        &heads as *const _ as *mut _,
425        &m as *const _ as *mut _,
426        &n as *const _ as *mut _,
427        &k as *const _ as *mut _,
428    ];
429
430    // SAFETY: Kernel launch requires FFI. All buffers are valid GPU allocations with
431    // matching sizes, and the kernel parameters match the expected PTX signature.
432    unsafe {
433        stream.launch_kernel(module, "batched_4d_gemm", &config, &mut args).map_err(|e| {
434            CudaTensorError::KernelError(format!("Batched 4D GEMM forward launch failed: {e:?}"))
435        })?;
436    }
437
438    Ok(())
439}
440
441/// NF4 quantized GEMM forward pass on GPU (trueno#108).
442///
443/// Computes: C = A @ dequant(B_nf4) where:
444/// - A is MxK (f32 activations)
445/// - B_nf4 is packed 4-bit NF4 weights (u8)
446/// - B_scales is per-block f32 scale factors
447/// - C is MxN (f32 output)
448///
449/// The kernel fuses dequantization with matmul: no intermediate fp32 weight buffer needed.
450///
451/// # Contract: C-NF4-003 (GEMM Numerical Parity)
452///
453/// `nf4_gemm(A, Q) ≈ naive_gemm(A, dequantize(Q))` within 1e-3 per-element.
454#[cfg(feature = "cuda")]
455pub fn gemm_nf4_forward(
456    a: &GpuBuffer<f32>,
457    b_nf4: &GpuBuffer<u8>,
458    b_scales: &GpuBuffer<f32>,
459    c: &mut GpuBuffer<f32>,
460    m: u32,
461    k: u32,
462    n: u32,
463    stream: &CudaStream,
464) -> Result<()> {
465    let cache = FORWARD_KERNEL_CACHE.get().ok_or(CudaTensorError::DeviceNotInitialized)?;
466    let mut cache = cache.lock().map_err(|_err| {
467        CudaTensorError::KernelError("Failed to acquire kernel cache lock".to_string())
468    })?;
469
470    let kernel = Nf4GemmKernel::new(m, n, k);
471    let tile_size = kernel.tile_size;
472
473    // Cache key excludes M (seq_len) — PTX is shape-independent (m/n/k are
474    // runtime params, only tile_size is baked in). Including M causes cache misses
475    // when actual seq_len differs from max_seq_len used during pre-warming,
476    // triggering on-demand JIT that fails on Blackwell (trueno#184).
477    let key = format!("nf4_gemm_forward_{k}_{n}");
478    let module = match cache.get_cached(&key) {
479        Some(m) => m,
480        None => {
481            let ptx = kernel.emit_ptx_for_target(cache.sm_target());
482            cache.get_or_compile(&key, &ptx)?
483        }
484    };
485
486    // Use tile_size × tile_size thread blocks (same as Q4K GEMM)
487    let config = LaunchConfig {
488        grid: (n.div_ceil(tile_size), m.div_ceil(tile_size), 1),
489        block: (tile_size * tile_size, 1, 1),
490        shared_mem: 16 * 4, // NF4 codebook LUT (16 × f32)
491    };
492
493    let a_ptr = a.as_ptr();
494    let b_nf4_ptr = b_nf4.as_ptr();
495    let b_scales_ptr = b_scales.as_ptr();
496    let c_ptr = c.as_ptr();
497
498    // PTX kernel signature: (a_ptr, b_nf4_ptr, b_scales_ptr, c_ptr, m, n, k)
499    // CRITICAL: must match param declaration order in Nf4GemmKernel::build_ptx()
500    let mut args: [*mut std::ffi::c_void; 7] = [
501        &a_ptr as *const _ as *mut _,
502        &b_nf4_ptr as *const _ as *mut _,
503        &b_scales_ptr as *const _ as *mut _,
504        &c_ptr as *const _ as *mut _,
505        &m as *const _ as *mut _,
506        &n as *const _ as *mut _,
507        &k as *const _ as *mut _,
508    ];
509
510    // SAFETY: Kernel launch requires FFI. All buffers are valid GPU allocations with
511    // matching sizes, and the kernel parameters match the expected PTX signature.
512    unsafe {
513        stream.launch_kernel(module, "nf4_gemm_fused", &config, &mut args).map_err(|e| {
514            CudaTensorError::KernelError(format!("NF4 GEMM forward launch failed: {e:?}"))
515        })?;
516    }
517
518    Ok(())
519}
520
521/// PMAT-481: NF4 tensor core GEMM — WMMA 16×16×16 with inline NF4 dequant in SHMEM.
522///
523/// Dequantizes NF4 blocks to FP16 in shared memory, uses tensor cores for matmul.
524/// Expected 5-40x compute improvement over naive tiled NF4 GEMM.
525///
526/// Contract: nf4-tensor-core-gemm-v1.yaml (F-NF4-TC-001, F-NF4-TC-002)
527#[cfg(feature = "cuda")]
528pub fn gemm_nf4_tc_forward(
529    a: &GpuBuffer<f32>,
530    b_nf4: &GpuBuffer<u8>,
531    b_scales: &GpuBuffer<f32>,
532    c: &mut GpuBuffer<f32>,
533    m: u32,
534    k: u32,
535    n: u32,
536    stream: &CudaStream,
537) -> Result<()> {
538    let cache = FORWARD_KERNEL_CACHE.get().ok_or(CudaTensorError::DeviceNotInitialized)?;
539    let mut cache = cache.lock().map_err(|_err| {
540        CudaTensorError::KernelError("Failed to acquire kernel cache lock".to_string())
541    })?;
542
543    let kernel = Nf4TensorCoreGemmKernel::new(m, n, k);
544
545    let key = format!("nf4_tc_gemm_forward_{k}_{n}");
546    let module = match cache.get_cached(&key) {
547        Some(m) => m,
548        None => {
549            let ptx = kernel.emit_ptx_for_target(cache.sm_target());
550            cache.get_or_compile(&key, &ptx)?
551        }
552    };
553
554    // WMMA: 1 warp (32 threads) per 16×16 tile
555    let config = LaunchConfig {
556        grid: (n.div_ceil(16), m.div_ceil(16), 1),
557        block: (32, 1, 1),
558        shared_mem: 16 * 16 * 2 * 2, // A[16×16] + B[16×16] in FP16
559    };
560
561    let a_ptr = a.as_ptr();
562    let b_nf4_ptr = b_nf4.as_ptr();
563    let b_scales_ptr = b_scales.as_ptr();
564    let c_ptr = c.as_ptr();
565
566    // Kernel signature: (a_ptr, scales_ptr, data_ptr, c_ptr, m, n, k)
567    let mut args: [*mut std::ffi::c_void; 7] = [
568        &a_ptr as *const _ as *mut _,
569        &b_scales_ptr as *const _ as *mut _,
570        &b_nf4_ptr as *const _ as *mut _,
571        &c_ptr as *const _ as *mut _,
572        &m as *const _ as *mut _,
573        &n as *const _ as *mut _,
574        &k as *const _ as *mut _,
575    ];
576
577    // SAFETY: launches a CUDA kernel via the driver API. The argument pointer array, grid/block config, and module/function name match the kernel's signature, and every referenced device buffer is allocated, correctly sized, and lives until the stream-ordered launch completes.
578    unsafe {
579        stream.launch_kernel(module, "nf4_tensor_core_gemm", &config, &mut args).map_err(|e| {
580            CudaTensorError::KernelError(format!(
581                "NF4 tensor core GEMM forward launch failed: {e:?}"
582            ))
583        })?;
584    }
585
586    Ok(())
587}
588
589/// PMAT-475: Fused NF4 Gate+Up GEMM — computes both projections with shared input load.
590///
591/// Eliminates one full input activation read from DRAM per call.
592/// Savings: M × K × 4 bytes/call (12 MB/layer for Qwen 1.5B batch=4 seq=512).
593///
594/// `gate[M×N] = A[M×K] @ dequant(W_gate_nf4)` and
595/// `up[M×N]   = A[M×K] @ dequant(W_up_nf4)` in one kernel launch.
596pub fn gemm_nf4_gate_up_forward(
597    a: &GpuBuffer<f32>,
598    wg_nf4: &GpuBuffer<u8>,
599    wg_scales: &GpuBuffer<f32>,
600    wu_nf4: &GpuBuffer<u8>,
601    wu_scales: &GpuBuffer<f32>,
602    gate: &mut GpuBuffer<f32>,
603    up: &mut GpuBuffer<f32>,
604    m: u32,
605    k: u32,
606    n: u32,
607    stream: &CudaStream,
608) -> Result<()> {
609    use trueno_gpu::kernels::FusedNf4GateUpGemmKernel;
610
611    let cache = FORWARD_KERNEL_CACHE.get().ok_or(CudaTensorError::DeviceNotInitialized)?;
612    let mut cache = cache.lock().map_err(|_err| {
613        CudaTensorError::KernelError("Failed to acquire kernel cache lock".to_string())
614    })?;
615
616    let kernel = FusedNf4GateUpGemmKernel::new(m, n, k);
617    let tile = kernel.tile_size;
618    let key = format!("fused_nf4_gate_up_{k}_{n}");
619    let module = match cache.get_cached(&key) {
620        Some(m) => m,
621        None => {
622            let ptx = kernel.emit_ptx_for_target(cache.sm_target());
623            cache.get_or_compile(&key, &ptx)?
624        }
625    };
626
627    let config = LaunchConfig {
628        grid: (n.div_ceil(tile), m.div_ceil(tile), 1),
629        block: (tile * tile, 1, 1),
630        shared_mem: 16 * 4,
631    };
632
633    let a_ptr = a.as_ptr();
634    let gate_ptr = gate.as_ptr();
635    let up_ptr = up.as_ptr();
636    let wg_nf4_ptr = wg_nf4.as_ptr();
637    let wg_scales_ptr = wg_scales.as_ptr();
638    let wu_nf4_ptr = wu_nf4.as_ptr();
639    let wu_scales_ptr = wu_scales.as_ptr();
640
641    let mut args: [*mut std::ffi::c_void; 10] = [
642        &gate_ptr as *const _ as *mut _,
643        &up_ptr as *const _ as *mut _,
644        &a_ptr as *const _ as *mut _,
645        &wg_scales_ptr as *const _ as *mut _,
646        &wg_nf4_ptr as *const _ as *mut _,
647        &wu_scales_ptr as *const _ as *mut _,
648        &wu_nf4_ptr as *const _ as *mut _,
649        &m as *const _ as *mut _,
650        &n as *const _ as *mut _,
651        &k as *const _ as *mut _,
652    ];
653
654    // SAFETY: launches a CUDA kernel via the driver API. The argument pointer array, grid/block config, and module/function name match the kernel's signature, and every referenced device buffer is allocated, correctly sized, and lives until the stream-ordered launch completes.
655    unsafe {
656        stream.launch_kernel(module, "fused_nf4_gate_up_gemm", &config, &mut args).map_err(
657            |e| CudaTensorError::KernelError(format!("Fused NF4 gate+up launch: {e:?}")),
658        )?;
659    }
660
661    Ok(())
662}
663
664/// BF16-precision GEMM forward pass on GPU (R-002: BF16 mixed precision).
665///
666/// Computes: C = A @ B where A is MxK, B is KxN, C is MxN
667/// Both inputs are f32 (FP32 master weights), but compute is done at BF16
668/// precision: each operand is truncated to BF16 (7-bit mantissa) before
669/// multiply, with FP32 accumulation. Output is FP32.
670///
671/// This implements the standard mixed-precision pattern:
672/// - FP32 storage (master weights stay in full precision)
673/// - BF16 compute (reduced precision multiply for bandwidth savings)
674/// - FP32 accumulation (no loss in reduction precision)
675///
676/// # Contract (C-BF16GEMM-001)
677///
678/// - `C[i,j] = Σ_k trunc_bf16(A[i,k]) * trunc_bf16(B[k,j])` accumulated in f32
679/// - `trunc_bf16(x)` = f32::from_bits(x.to_bits() & 0xFFFF0000)
680/// - Output matches CPU BF16 reference within f32 accumulation tolerance
681#[cfg(feature = "cuda")]
682pub fn gemm_forward_bf16(
683    a: &GpuBuffer<f32>,
684    b: &GpuBuffer<f32>,
685    c: &mut GpuBuffer<f32>,
686    m: u32,
687    k: u32,
688    n: u32,
689    stream: &CudaStream,
690) -> Result<()> {
691    let cache = FORWARD_KERNEL_CACHE.get().ok_or(CudaTensorError::DeviceNotInitialized)?;
692    let mut cache = cache.lock().map_err(|_err| {
693        CudaTensorError::KernelError("Failed to acquire kernel cache lock".to_string())
694    })?;
695
696    let key = format!("gemm_bf16_compute_{m}_{k}_{n}");
697    let module = match cache.get_cached(&key) {
698        Some(m) => m,
699        None => {
700            let ptx = build_gemm_bf16_compute_ptx(cache.sm_target());
701            cache.get_or_compile(&key, &ptx)?
702        }
703    };
704
705    let config = LaunchConfig {
706        grid: (n.div_ceil(16), m.div_ceil(16), 1),
707        block: (16, 16, 1),
708        shared_mem: 0,
709    };
710
711    let a_ptr = a.as_ptr();
712    let b_ptr = b.as_ptr();
713    let c_ptr = c.as_ptr();
714
715    // PTX kernel signature: (a_ptr, b_ptr, c_ptr, m, n, k)
716    // CRITICAL: must match param declaration order in build_gemm_bf16_compute_ptx()
717    let mut args: [*mut std::ffi::c_void; 6] = [
718        &a_ptr as *const _ as *mut _,
719        &b_ptr as *const _ as *mut _,
720        &c_ptr as *const _ as *mut _,
721        &m as *const _ as *mut _,
722        &n as *const _ as *mut _,
723        &k as *const _ as *mut _,
724    ];
725
726    // SAFETY: Kernel launch requires FFI. All buffers are valid GPU allocations with
727    // matching sizes, and the kernel parameters match the expected PTX signature.
728    unsafe {
729        stream.launch_kernel(module, "gemm_bf16_compute", &config, &mut args).map_err(|e| {
730            CudaTensorError::KernelError(format!("BF16 GEMM forward launch failed: {e:?}"))
731        })?;
732    }
733
734    Ok(())
735}
736
737/// Build PTX for BF16-precision GEMM kernel.
738///
739/// Naive GEMM with inline BF16 truncation: loads f32, truncates to bf16 precision
740/// (AND 0xFFFF0000), multiplies as f32, accumulates in f32. This matches the
741/// precision characteristics of hardware BF16 tensor cores (BF16 multiply, f32 accum,
742/// safe because forward GEMMs are NoTrans/NoTrans — unaffected by ALB-076/trueno#170).
743#[cfg(feature = "cuda")]
744fn build_gemm_bf16_compute_ptx(sm_target: &str) -> String {
745    format!(
746        r".version 7.0
747.target {sm_target}
748.address_size 64
749
750.visible .entry gemm_bf16_compute(
751    .param .u64 a_ptr,
752    .param .u64 b_ptr,
753    .param .u64 c_ptr,
754    .param .u32 M,
755    .param .u32 N,
756    .param .u32 K
757) {{
758    .reg .u32 %r<20>;
759    .reg .u64 %rd<8>;
760    .reg .f32 %f<4>;
761    .reg .pred %p<4>;
762
763    // col = ctaid.x * 16 + tid.x
764    mov.u32 %r0, %ctaid.x;
765    mov.u32 %r1, %ntid.x;
766    mov.u32 %r2, %tid.x;
767    mad.lo.u32 %r3, %r0, %r1, %r2;
768
769    // row = ctaid.y * 16 + tid.y
770    mov.u32 %r4, %ctaid.y;
771    mov.u32 %r5, %ntid.y;
772    mov.u32 %r6, %tid.y;
773    mad.lo.u32 %r7, %r4, %r5, %r6;
774
775    // Load params
776    ld.param.u64 %rd0, [a_ptr];
777    ld.param.u64 %rd1, [b_ptr];
778    ld.param.u64 %rd2, [c_ptr];
779    ld.param.u32 %r8, [M];
780    ld.param.u32 %r9, [N];
781    ld.param.u32 %r10, [K];
782
783    // Bounds check: row < M && col < N
784    setp.ge.u32 %p0, %r7, %r8;
785    setp.ge.u32 %p1, %r3, %r9;
786    or.pred %p2, %p0, %p1;
787    @%p2 bra exit;
788
789    // acc = 0.0f
790    mov.f32 %f0, 0f00000000;
791
792    // Loop: for i = 0; i < K; i++
793    mov.u32 %r11, 0;
794loop_start:
795    setp.ge.u32 %p3, %r11, %r10;
796    @%p3 bra loop_end;
797
798    // Load A[row, i] as u32 bits, truncate to bf16 precision
799    mul.lo.u32 %r12, %r7, %r10;
800    add.u32 %r12, %r12, %r11;
801    mul.wide.u32 %rd3, %r12, 4;
802    add.u64 %rd3, %rd0, %rd3;
803    ld.global.u32 %r13, [%rd3];
804    and.b32 %r13, %r13, 0xFFFF0000;
805    mov.b32 %f1, %r13;
806
807    // Load B[i, col] as u32 bits, truncate to bf16 precision
808    mul.lo.u32 %r14, %r11, %r9;
809    add.u32 %r14, %r14, %r3;
810    mul.wide.u32 %rd4, %r14, 4;
811    add.u64 %rd4, %rd1, %rd4;
812    ld.global.u32 %r15, [%rd4];
813    and.b32 %r15, %r15, 0xFFFF0000;
814    mov.b32 %f2, %r15;
815
816    // acc += a_bf16 * b_bf16 (FMA in f32 accumulator)
817    fma.rn.f32 %f0, %f1, %f2, %f0;
818
819    add.u32 %r11, %r11, 1;
820    bra loop_start;
821
822loop_end:
823    // Store C[row, col]
824    mul.lo.u32 %r16, %r7, %r9;
825    add.u32 %r16, %r16, %r3;
826    mul.wide.u32 %rd5, %r16, 4;
827    add.u64 %rd5, %rd2, %rd5;
828    st.global.f32 [%rd5], %f0;
829
830exit:
831    ret;
832}}
833"
834    )
835}
836
837/// cuBLAS GEMM for NF4 QLoRA forward with pre-dequantized fp32 weights (ENT-287).
838///
839/// Computes: `C[M,N] = A[M,K] @ W[N,K]^T` where W is stored row-major `[N,K]`
840/// (HuggingFace convention: `[out_features, in_features]`).
841///
842/// # Weight Layout Derivation
843///
844/// W is row-major `[N,K]`: element `(i,j)` at offset `i*K + j`.
845/// In column-major this is `[K,N]` with leading dimension K.
846///
847/// We want `C = A @ W^T`. Expanding in row-major: `C[M,N] = A[M,K] @ W^T[K,N]`.
848///
849/// Column-major equivalent: `C_cm[N,M] = (W^T)_cm[N,K] @ A_cm[K,M]`.
850/// Since W_cm is `[K,N]`, applying cuBLAS Trans gives `[N,K]` with `lda = K`.
851/// A_cm is `[K,M]` with `ldb = K`. C_cm is `[N,M]` with `ldc = N`.
852///
853/// cuBLAS call: `(Trans, NoTrans, N, M, K, W_ptr, K, A_ptr, K, C_ptr, N)`.
854///
855/// # Arguments
856///
857/// * `a` - Input activations `[M, K]` row-major (f32)
858/// * `w` - Weight matrix `[N, K]` row-major = `[K, N]` col-major (f32, original fp32 weights)
859/// * `c` - Output `[M, N]` row-major (f32)
860/// * `m` - Rows of A (seq_len)
861/// * `k` - Columns of A / columns of W (input dimension)
862/// * `n` - Rows of W (output dimension)
863///
864/// # Contract (C-NF4CUBLAS-001)
865///
866/// `gemm_nf4_dequant_cublas(A, W) = A @ W^T` within f32 precision.
867#[cfg(feature = "cuda")]
868pub fn gemm_nf4_dequant_cublas(
869    a: &GpuBuffer<f32>,
870    w: &GpuBuffer<f32>,
871    c: &mut GpuBuffer<f32>,
872    m: u32,
873    k: u32,
874    n: u32,
875    stream: &CudaStream,
876) -> Result<()> {
877    let cache = FORWARD_KERNEL_CACHE.get().ok_or(CudaTensorError::DeviceNotInitialized)?;
878    let cache = cache.lock().map_err(|_err| {
879        CudaTensorError::KernelError("Failed to acquire kernel cache lock".to_string())
880    })?;
881
882    let cublas = cache.cublas().ok_or_else(|| {
883        CudaTensorError::KernelError("cuBLAS not available for NF4 dequant GEMM".to_string())
884    })?;
885    bind_cublas_stream(cublas, stream)?;
886
887    // C[M,N] = A[M,K] @ W[N,K]^T
888    // col-major: C_cm[N,M] = W_cm_transposed[N,K] @ A_cm[K,M]
889    // W_cm is [K,N] with lda=K. Trans on it gives [N,K].
890    // A_cm is [K,M] with lda=K.
891    // C_cm is [N,M] with ldc=N.
892    cublas
893        .gemm_f32(
894            GemmOp::Trans,   // W_cm[K,N] transposed → [N,K]
895            GemmOp::NoTrans, // A_cm[K,M]
896            n as i32,        // rows of op(W) = N
897            m as i32,        // cols of op(A) = M
898            k as i32,        // shared dim = K
899            1.0,
900            w.as_ptr(), // W: row-major [N,K] = col-major [K,N], lda=K
901            k as i32,   // lda = K (leading dim of W_cm[K,N])
902            a.as_ptr(), // A: row-major [M,K] = col-major [K,M], lda=K
903            k as i32,   // ldb = K (leading dim of A_cm[K,M])
904            0.0,
905            c.as_ptr(), // C: row-major [M,N] = col-major [N,M], ldc=N
906            n as i32,   // ldc = N
907        )
908        .map_err(|e| {
909            CudaTensorError::KernelError(format!("cuBLAS NF4 dequant forward failed: {e:?}"))
910        })
911}
912
913/// cuBLAS GEMM for NF4 QLoRA backward: grad_input (ENT-287).
914///
915/// Computes: `grad_input[M,K] = grad_output[M,N] @ W[N,K]` where W is row-major `[N,K]`.
916///
917/// This is standard GEMM `C = A @ B` where `B = W[N,K]`.
918///
919/// Derivation:
920/// Row-major: `C[M,K] = A[M,N] @ B[N,K]`
921/// col-major: `C_cm[K,M] = B_cm[K,N] @ A_cm[N,M]`
922/// - B = W row-major `[N,K]` = col-major `[K,N]` with `lda = K`
923/// - A = grad_out row-major `[M,N]` = col-major `[N,M]` with `ldb = N`
924/// - C = grad_in row-major `[M,K]` = col-major `[K,M]` with `ldc = K`
925/// So: `cublas(NoTrans, NoTrans, K, M, N, W_ptr, K, grad_out_ptr, N, grad_in_ptr, K)`
926///
927/// # Arguments
928///
929/// * `grad_output` - Upstream gradient `[M, N]` (f32)
930/// * `w` - Weight matrix `[N, K]` row-major (f32, pre-dequantized)
931/// * `grad_input` - Output gradient `[M, K]` (f32)
932/// * `m` - Rows (seq_len)
933/// * `k` - Output columns (input dimension)
934/// * `n` - Shared dimension (output dimension)
935///
936/// # Contract (C-NF4CUBLAS-002)
937///
938/// `gemm_nf4_backward_a_cublas(grad, W) = grad @ W` within f32 precision.
939#[cfg(feature = "cuda")]
940pub fn gemm_nf4_backward_a_cublas(
941    grad_output: &GpuBuffer<f32>,
942    w: &GpuBuffer<f32>,
943    grad_input: &mut GpuBuffer<f32>,
944    m: u32,
945    k: u32,
946    n: u32,
947    stream: &CudaStream,
948) -> Result<()> {
949    let cache = FORWARD_KERNEL_CACHE.get().ok_or(CudaTensorError::DeviceNotInitialized)?;
950    let cache = cache.lock().map_err(|_err| {
951        CudaTensorError::KernelError("Failed to acquire kernel cache lock".to_string())
952    })?;
953
954    let cublas = cache.cublas().ok_or_else(|| {
955        CudaTensorError::KernelError("cuBLAS not available for NF4 backward GEMM".to_string())
956    })?;
957    bind_cublas_stream(cublas, stream)?;
958
959    // grad_in[M,K] = grad_out[M,N] @ W[N,K]
960    // col-major: C_cm[K,M] = W_cm[K,N] @ A_cm[N,M]
961    cublas
962        .gemm_f32(
963            GemmOp::NoTrans, // W_cm[K,N] as-is
964            GemmOp::NoTrans, // grad_out_cm[N,M] as-is
965            k as i32,        // rows of W_cm = K
966            m as i32,        // cols of grad_out_cm = M
967            n as i32,        // shared dim = N
968            1.0,
969            w.as_ptr(),           // W: row-major [N,K] = col-major [K,N], lda=K
970            k as i32,             // lda = K
971            grad_output.as_ptr(), // grad_out: row-major [M,N] = col-major [N,M], ldb=N
972            n as i32,             // ldb = N
973            0.0,
974            grad_input.as_ptr(), // grad_in: row-major [M,K] = col-major [K,M], ldc=K
975            k as i32,            // ldc = K
976        )
977        .map_err(|e| CudaTensorError::KernelError(format!("cuBLAS NF4 backward_a failed: {e:?}")))
978}
979
980/// NF4 transposed GEMM for backward pass (ENT-153: QLoRA backward).
981///
982/// Computes: `grad_input[M×K] = grad_output[M×N] @ dequant(W_nf4[K×N])^T`
983///
984/// This is the gradient-flow kernel: given upstream gradient and frozen NF4 weights,
985/// computes the input gradient without materializing fp32 weights.
986///
987/// # Arguments
988///
989/// * `grad_output` - Upstream gradient `[M × N]` (f32)
990/// * `w_nf4` - Frozen NF4-packed weights for `W[K × N]` (u8)
991/// * `w_scales` - Per-block scales for `W[K × N]` (f32)
992/// * `grad_input` - Output gradient `[M × K]` (f32)
993/// * `m` - Rows of grad_output (seq_len)
994/// * `n` - Columns of W (reduction dimension)
995/// * `k` - Rows of W (output columns = input dimension)
996///
997/// # Contract: C-NF4T-001 (Transposed GEMM Parity)
998///
999/// `gemm_nf4_backward_a(grad, W_nf4) ≈ gemm(grad, dequant(W)^T)` within 1e-3.
1000#[cfg(feature = "cuda")]
1001pub fn gemm_nf4_backward_a(
1002    grad_output: &GpuBuffer<f32>,
1003    w_nf4: &GpuBuffer<u8>,
1004    w_scales: &GpuBuffer<f32>,
1005    grad_input: &mut GpuBuffer<f32>,
1006    m: u32,
1007    n: u32,
1008    k: u32,
1009    stream: &CudaStream,
1010) -> Result<()> {
1011    let cache = FORWARD_KERNEL_CACHE.get().ok_or(CudaTensorError::DeviceNotInitialized)?;
1012    let mut cache = cache.lock().map_err(|_err| {
1013        CudaTensorError::KernelError("Failed to acquire kernel cache lock".to_string())
1014    })?;
1015
1016    let kernel = Nf4GemmTransposeKernel::new(m, n, k);
1017    let tile_size = kernel.tile_size;
1018
1019    // Cache key excludes M (seq_len) — PTX is shape-independent (trueno#184).
1020    let key = format!("nf4_gemm_transpose_{n}_{k}");
1021    let module = match cache.get_cached(&key) {
1022        Some(m) => m,
1023        None => {
1024            let ptx = kernel.emit_ptx_for_target(cache.sm_target());
1025            cache.get_or_compile(&key, &ptx)?
1026        }
1027    };
1028
1029    // Output is [M × K], tiled with tile_size
1030    let config = LaunchConfig {
1031        grid: (k.div_ceil(tile_size), m.div_ceil(tile_size), 1),
1032        block: (tile_size * tile_size, 1, 1),
1033        shared_mem: 16 * 4, // NF4 codebook LUT
1034    };
1035
1036    let a_ptr = grad_output.as_ptr();
1037    let b_nf4_ptr = w_nf4.as_ptr();
1038    let b_scales_ptr = w_scales.as_ptr();
1039    let c_ptr = grad_input.as_ptr();
1040
1041    let mut args: [*mut std::ffi::c_void; 7] = [
1042        &a_ptr as *const _ as *mut _,
1043        &b_nf4_ptr as *const _ as *mut _,
1044        &b_scales_ptr as *const _ as *mut _,
1045        &c_ptr as *const _ as *mut _,
1046        &m as *const _ as *mut _,
1047        &n as *const _ as *mut _,
1048        &k as *const _ as *mut _,
1049    ];
1050
1051    // SAFETY: Kernel launch requires FFI. All buffers are valid GPU allocations.
1052    unsafe {
1053        stream.launch_kernel(module, "nf4_gemm_transpose", &config, &mut args).map_err(|e| {
1054            CudaTensorError::KernelError(format!("NF4 GEMM transpose launch failed: {e:?}"))
1055        })?;
1056    }
1057
1058    Ok(())
1059}
1060
1061/// PMAT-481: NF4 tensor core backward GEMM — WMMA 16×16×16 with inline NF4 dequant.
1062///
1063/// Computes `grad_input[M×K] = grad_output[M×N] @ dequant(B_nf4[K×N])^T`
1064///
1065/// Eliminates separate dequant kernel + generic cuBLAS GEMM per backward projection.
1066/// Uses trueno `Nf4TensorCoreGemmBackwardAKernel` (WMMA, shared memory NF4 dequant).
1067///
1068/// Contract: nf4-backward-tensor-core-gemm-v1.yaml
1069#[cfg(feature = "cuda")]
1070pub fn gemm_nf4_tc_backward_a(
1071    grad_output: &GpuBuffer<f32>,
1072    w_nf4: &GpuBuffer<u8>,
1073    w_scales: &GpuBuffer<f32>,
1074    grad_input: &mut GpuBuffer<f32>,
1075    m: u32,
1076    n: u32,
1077    k: u32,
1078    stream: &CudaStream,
1079) -> Result<()> {
1080    use trueno_gpu::kernels::backward::Nf4TensorCoreGemmBackwardAKernel;
1081
1082    let cache = FORWARD_KERNEL_CACHE.get().ok_or(CudaTensorError::DeviceNotInitialized)?;
1083    let mut cache = cache.lock().map_err(|_err| {
1084        CudaTensorError::KernelError("Failed to acquire kernel cache lock".to_string())
1085    })?;
1086
1087    let kernel = Nf4TensorCoreGemmBackwardAKernel::new(m, n, k);
1088
1089    // Cache key: backward TC is shape-independent for (n, k) pair
1090    let key = format!("nf4_tc_gemm_backward_a_{n}_{k}");
1091    let module = match cache.get_cached(&key) {
1092        Some(m) => m,
1093        None => {
1094            let ptx = kernel.emit_ptx_for_target(cache.sm_target());
1095            cache.get_or_compile(&key, &ptx)?
1096        }
1097    };
1098
1099    // WMMA backward: Grid = (ceil(K/16), ceil(M/16)), Block = 32 threads (1 warp)
1100    let config = LaunchConfig {
1101        grid: (k.div_ceil(16), m.div_ceil(16), 1),
1102        block: (32, 1, 1),
1103        shared_mem: 16 * 16 * 2 * 2, // grad_out[16×16] + B^T[16×16] in FP16
1104    };
1105
1106    let grad_out_ptr = grad_output.as_ptr();
1107    let scales_ptr = w_scales.as_ptr();
1108    let data_ptr = w_nf4.as_ptr();
1109    let grad_a_ptr = grad_input.as_ptr();
1110
1111    // Kernel signature: (grad_out_ptr, scales_ptr, data_ptr, grad_a_ptr, m, n, k)
1112    let mut args: [*mut std::ffi::c_void; 7] = [
1113        &grad_out_ptr as *const _ as *mut _,
1114        &scales_ptr as *const _ as *mut _,
1115        &data_ptr as *const _ as *mut _,
1116        &grad_a_ptr as *const _ as *mut _,
1117        &m as *const _ as *mut _,
1118        &n as *const _ as *mut _,
1119        &k as *const _ as *mut _,
1120    ];
1121
1122    // SAFETY: launches a CUDA kernel via the driver API. The argument pointer array, grid/block config, and module/function name match the kernel's signature, and every referenced device buffer is allocated, correctly sized, and lives until the stream-ordered launch completes.
1123    unsafe {
1124        stream
1125            .launch_kernel(module, "nf4_tensor_core_gemm_backward_a", &config, &mut args)
1126            .map_err(|e| {
1127                CudaTensorError::KernelError(format!(
1128                    "NF4 tensor core GEMM backward_a launch failed: {e:?}"
1129                ))
1130            })?;
1131    }
1132
1133    Ok(())
1134}