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#[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#[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(); 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 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#[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 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 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 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 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#[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 cublas
194 .gemm_f32(
195 GemmOp::Trans, GemmOp::NoTrans, n as i32,
198 m as i32,
199 k as i32,
200 1.0,
201 b.as_ptr(),
202 k as i32, a.as_ptr(),
204 k as i32, 0.0,
206 c.as_ptr(),
207 n as i32, )
209 .map_err(|e| CudaTensorError::KernelError(format!("cuBLAS GEMM BT failed: {e:?}")))
210}
211
212#[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#[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#[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, 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#[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#[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 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 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 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 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#[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 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 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, };
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 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 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#[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 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, };
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 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 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
589pub 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 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#[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 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 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#[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#[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 cublas
893 .gemm_f32(
894 GemmOp::Trans, GemmOp::NoTrans, n as i32, m as i32, k as i32, 1.0,
900 w.as_ptr(), k as i32, a.as_ptr(), k as i32, 0.0,
905 c.as_ptr(), n as i32, )
908 .map_err(|e| {
909 CudaTensorError::KernelError(format!("cuBLAS NF4 dequant forward failed: {e:?}"))
910 })
911}
912
913#[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 cublas
962 .gemm_f32(
963 GemmOp::NoTrans, GemmOp::NoTrans, k as i32, m as i32, n as i32, 1.0,
969 w.as_ptr(), k as i32, grad_output.as_ptr(), n as i32, 0.0,
974 grad_input.as_ptr(), k as i32, )
977 .map_err(|e| CudaTensorError::KernelError(format!("cuBLAS NF4 backward_a failed: {e:?}")))
978}
979
980#[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 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 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, };
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 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#[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 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 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, };
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 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 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}