Skip to main content

entrenar/autograd/cuda_backward/
structured.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::{CudaStream, GpuBuffer, LaunchConfig};
8#[cfg(feature = "cuda")]
9use trueno_gpu::kernels::backward::{
10    BatchedRmsNormBackwardKernel, BatchedSoftmaxBackwardKernel, LayerNormBackwardKernel,
11    RmsNormGammaReduceKernel, SoftmaxBackwardKernel,
12};
13#[cfg(feature = "cuda")]
14use trueno_gpu::kernels::BatchedVectorizedRmsNormKernel;
15#[cfg(feature = "cuda")]
16use trueno_gpu::kernels::Kernel;
17
18use super::super::cuda_tensor::{CudaTensorError, Result};
19#[cfg(feature = "cuda")]
20use super::cache::KERNEL_CACHE;
21#[cfg(feature = "cuda")]
22use provable_contracts_macros::requires;
23
24/// Softmax backward pass on GPU
25///
26/// Computes: grad_input = softmax_output * (grad_output - sum(grad_output * softmax_output))
27#[cfg(feature = "cuda")]
28// Contract: backward-pass-v1 / softmax_backward
29#[requires(batch_size > 0 && seq_len > 0)]
30pub fn softmax_backward(
31    softmax_output: &GpuBuffer<f32>,
32    grad_output: &GpuBuffer<f32>,
33    grad_input: &mut GpuBuffer<f32>,
34    batch_size: u32,
35    seq_len: u32,
36    stream: &CudaStream,
37) -> Result<()> {
38    let cache = KERNEL_CACHE.get().ok_or(CudaTensorError::DeviceNotInitialized)?;
39    let mut cache = cache.lock().map_err(|_err| {
40        CudaTensorError::KernelError("Failed to acquire kernel cache lock".to_string())
41    })?;
42
43    let key = format!("softmax_backward_{batch_size}_{seq_len}");
44    let module = match cache.get_cached(&key) {
45        Some(m) => m,
46        None => {
47            let kernel = SoftmaxBackwardKernel::new(batch_size, seq_len);
48            let ptx = kernel.emit_ptx_for_target(cache.sm_target());
49            cache.get_or_compile(&key, &ptx)?
50        }
51    };
52
53    // Softmax backward uses warp-parallel reduction.
54    // FALSIFY-CUDA-NF4-TRAIN-LOSS-PARITY-001: launch a FULL 32-lane warp —
55    // the kernel's shfl.sync reductions use membermask 0xFFFFFFFF, which is
56    // undefined if any named lane is inactive (seq_len < 32 launched a
57    // partial warp). Lanes >= seq_len use predicated loads (identity 0.0),
58    // so a full warp is correct for every seq_len.
59    let config = LaunchConfig {
60        grid: (batch_size, 1, 1),
61        block: (32, 1, 1), // FULL warp — see above
62        shared_mem: 0,
63    };
64
65    let output_ptr = softmax_output.as_ptr();
66    let grad_out_ptr = grad_output.as_ptr();
67    let grad_in_ptr = grad_input.as_ptr();
68
69    let mut args: [*mut std::ffi::c_void; 5] = [
70        &output_ptr as *const _ as *mut _,
71        &grad_out_ptr as *const _ as *mut _,
72        &grad_in_ptr as *const _ as *mut _,
73        &batch_size as *const _ as *mut _,
74        &seq_len as *const _ as *mut _,
75    ];
76
77    // SAFETY: Kernel launch requires FFI. All buffers are valid GPU allocations with
78    // matching sizes, and the kernel parameters match the expected PTX signature.
79    unsafe {
80        stream.launch_kernel(module, "softmax_backward", &config, &mut args).map_err(|e| {
81            CudaTensorError::KernelError(format!("Softmax backward launch failed: {e:?}"))
82        })?;
83    }
84
85    Ok(())
86}
87
88/// Batched softmax backward pass on GPU (handles row_size > 32)
89///
90/// Computes: grad_input[r][i] = y[r][i] * (grad_output[r][i] - Σⱼ grad_output[r][j] * y[r][j])
91///
92/// Uses stride-loop + warp-shuffle reduction (one warp per row, one block per row).
93///
94/// # Contract (C-BSMAX-BACK-002)
95///
96/// - **Precondition**: softmax_output contains valid softmax output, all buffers have at least
97///   total_rows * row_size elements, row_size > 0, total_rows > 0, KERNEL_CACHE initialized
98/// - **Postcondition**: grad_input[r][i] = y[r][i] * (∂L/∂y[r][i] - dot(∂L/∂y[r], y[r]))
99/// - **Invariant**: Zero CPU-side data transfers; in-place safe (grad_input may alias grad_output)
100#[cfg(feature = "cuda")]
101pub fn batched_softmax_backward(
102    softmax_output: &GpuBuffer<f32>,
103    grad_output: &GpuBuffer<f32>,
104    grad_input: &mut GpuBuffer<f32>,
105    total_rows: u32,
106    row_size: u32,
107    stream: &CudaStream,
108) -> Result<()> {
109    let cache = KERNEL_CACHE.get().ok_or(CudaTensorError::DeviceNotInitialized)?;
110    let mut cache = cache.lock().map_err(|_err| {
111        CudaTensorError::KernelError("Failed to acquire kernel cache lock".to_string())
112    })?;
113
114    // Contract: dimension-independent-kernels-v1.yaml
115    // Note: BatchedSoftmaxBackwardKernel not yet dimension-independent in trueno,
116    // but using generic key prepares for the fix.
117    let key = "batched_softmax_backward";
118    let module = match cache.get_cached(key) {
119        Some(m) => m,
120        None => {
121            let kernel = BatchedSoftmaxBackwardKernel::new(total_rows, row_size);
122            let ptx = kernel.emit_ptx_for_target(cache.sm_target());
123            cache.get_or_compile(key, &ptx)?
124        }
125    };
126
127    // One FULL warp (32 threads) per row, one block per row.
128    // FALSIFY-CUDA-NF4-TRAIN-LOSS-PARITY-001: was `32.min(row_size)` — a
129    // partial warp makes the kernel's 0xFFFFFFFF shfl.sync reductions
130    // undefined (garbage row reductions for seq < 32). Guarded loops carry
131    // identity 0.0 on idle lanes, so a full warp is correct.
132    let config = LaunchConfig { grid: (total_rows, 1, 1), block: (32, 1, 1), shared_mem: 0 };
133
134    let output_ptr = softmax_output.as_ptr();
135    let grad_out_ptr = grad_output.as_ptr();
136    let grad_in_ptr = grad_input.as_ptr();
137
138    let mut args: [*mut std::ffi::c_void; 5] = [
139        &output_ptr as *const _ as *mut _,
140        &grad_out_ptr as *const _ as *mut _,
141        &grad_in_ptr as *const _ as *mut _,
142        &total_rows as *const _ as *mut _,
143        &row_size as *const _ as *mut _,
144    ];
145
146    // SAFETY: Kernel launch requires FFI. All buffers are valid GPU allocations with
147    // matching sizes, and the kernel parameters match the expected PTX signature.
148    unsafe {
149        stream.launch_kernel(module, "batched_softmax_backward", &config, &mut args).map_err(
150            |e| {
151                CudaTensorError::KernelError(format!(
152                    "Batched softmax backward launch failed: {e:?}"
153                ))
154            },
155        )?;
156    }
157
158    Ok(())
159}
160
161/// RMSNorm backward pass on GPU
162///
163/// Computes gradients for input (and placeholder for gamma parameters).
164/// Uses stride-loop kernel that supports arbitrary hidden_size (no warp-only limit).
165///
166/// # Contract (C-RMSBACK-WRAP-001)
167///
168/// - **Precondition**: input contains original forward input, gamma has hidden_size elements,
169///   all buffers allocated with at least batch_size * hidden_size elements
170/// - **Postcondition**: grad_input contains ∂L/∂x per the RMSNorm backward formula;
171///   `grad_gamma[i]` contains `Σ_r (∂L/∂y[r][i] · x[r][i] / rms[r])` summed in
172///   fixed iteration order over rows (FALSIFY-GPUTRAIN-006).
173/// - **Invariant**: Uses batched stride-loop kernel + deterministic per-row partial
174///   reduction; no hidden_size upper limit; bit-exactly reproducible across two
175///   cuda:0 seed=0 runs (no atomicAdd in the gamma accumulation path).
176#[cfg(feature = "cuda")]
177pub fn rms_norm_backward(
178    input: &GpuBuffer<f32>,
179    gamma: &GpuBuffer<f32>,
180    grad_output: &GpuBuffer<f32>,
181    grad_input: &mut GpuBuffer<f32>,
182    grad_gamma: &mut GpuBuffer<f32>,
183    batch_size: u32,
184    hidden_size: u32,
185    eps: f32,
186    stream: &CudaStream,
187) -> Result<()> {
188    let cache = KERNEL_CACHE.get().ok_or(CudaTensorError::DeviceNotInitialized)?;
189    let mut cache = cache.lock().map_err(|_err| {
190        CudaTensorError::KernelError("Failed to acquire kernel cache lock".to_string())
191    })?;
192
193    // FALSIFY-GPUTRAIN-006: allocate per-row partial buffer
194    // `[batch_size × hidden_size]` for the deterministic two-stage reduction. Each
195    // backward block writes EXCLUSIVELY to `grad_gamma_partial[block_idx]`, then the
196    // companion `RmsNormGammaReduceKernel` sums rows in fixed order
197    // (`r = 0, 1, …, batch_size - 1`) into the final `grad_gamma[hidden_size]`.
198    // No atomicAdd is involved — the result is bit-exact across cuda:0 seed=0 reruns.
199    let partial_elem_count = (batch_size as usize) * (hidden_size as usize);
200    let ctx = cache.ctx().clone();
201    let grad_gamma_partial: GpuBuffer<f32> =
202        GpuBuffer::new(&ctx, partial_elem_count).map_err(|e| {
203            CudaTensorError::KernelError(format!(
204                "RMSNorm backward: grad_gamma_partial alloc failed ({batch_size}×{hidden_size}): {e:?}"
205            ))
206        })?;
207
208    // ── Stage 1: per-row partial backward kernel ────────────────────────
209    // Contract: dimension-independent-kernels-v1.yaml (FALSIFY-DIM-001)
210    let key = "batched_rms_norm_backward";
211    let module = match cache.get_cached(key) {
212        Some(m) => m,
213        None => {
214            let kernel = BatchedRmsNormBackwardKernel::new(batch_size, hidden_size, eps);
215            let ptx = kernel.emit_ptx_for_target(cache.sm_target());
216            cache.get_or_compile(key, &ptx)?
217        }
218    };
219
220    // One warp (32 threads) per row, one block per row
221    let config = LaunchConfig {
222        grid: (batch_size, 1, 1),
223        block: (32.min(hidden_size), 1, 1),
224        shared_mem: 0,
225    };
226
227    let input_ptr = input.as_ptr();
228    let gamma_ptr = gamma.as_ptr();
229    let grad_out_ptr = grad_output.as_ptr();
230    let grad_in_ptr = grad_input.as_ptr();
231    // FALSIFY-GPUTRAIN-006: pass the per-row partial buffer (NOT the final
232    // grad_gamma) so the backward kernel writes per-row slots without atomics.
233    let grad_gamma_partial_ptr = grad_gamma_partial.as_ptr();
234
235    let mut args: [*mut std::ffi::c_void; 8] = [
236        &input_ptr as *const _ as *mut _,
237        &gamma_ptr as *const _ as *mut _,
238        &grad_out_ptr as *const _ as *mut _,
239        &grad_in_ptr as *const _ as *mut _,
240        &grad_gamma_partial_ptr as *const _ as *mut _,
241        &batch_size as *const _ as *mut _,
242        &hidden_size as *const _ as *mut _,
243        &eps as *const _ as *mut _,
244    ];
245
246    // SAFETY: Kernel launch requires FFI. All buffers are valid GPU allocations with
247    // matching sizes, and the kernel parameters match the expected PTX signature.
248    unsafe {
249        stream.launch_kernel(module, "batched_rms_norm_backward", &config, &mut args).map_err(
250            |e| CudaTensorError::KernelError(format!("RMSNorm backward launch failed: {e:?}")),
251        )?;
252    }
253
254    // ── Stage 2: deterministic fixed-order cross-row reduction ──────────
255    let reduce_key = "rms_norm_gamma_reduce";
256    let reduce_module = match cache.get_cached(reduce_key) {
257        Some(m) => m,
258        None => {
259            let kernel = RmsNormGammaReduceKernel::new(batch_size, hidden_size);
260            let ptx = kernel.emit_ptx_for_target(cache.sm_target());
261            cache.get_or_compile(reduce_key, &ptx)?
262        }
263    };
264
265    let reduce_config = LaunchConfig {
266        grid: (hidden_size.div_ceil(RmsNormGammaReduceKernel::BLOCK_SIZE), 1, 1),
267        block: (RmsNormGammaReduceKernel::BLOCK_SIZE, 1, 1),
268        shared_mem: 0,
269    };
270
271    let final_grad_gamma_ptr = grad_gamma.as_ptr();
272
273    let mut reduce_args: [*mut std::ffi::c_void; 4] = [
274        &grad_gamma_partial_ptr as *const _ as *mut _,
275        &final_grad_gamma_ptr as *const _ as *mut _,
276        &batch_size as *const _ as *mut _,
277        &hidden_size as *const _ as *mut _,
278    ];
279
280    // SAFETY: Same FFI invariants as Stage 1. Both buffers are valid GPU
281    // allocations sized batch_size*hidden_size and hidden_size respectively.
282    unsafe {
283        stream
284            .launch_kernel(reduce_module, "rms_norm_gamma_reduce", &reduce_config, &mut reduce_args)
285            .map_err(|e| {
286                CudaTensorError::KernelError(format!("RMSNorm gamma-reduce launch failed: {e:?}"))
287            })?;
288    }
289
290    // grad_gamma_partial drops here; cudaFree is implicit via GpuBuffer Drop.
291    drop(grad_gamma_partial);
292    Ok(())
293}
294
295/// RMSNorm forward pass on GPU (KAIZEN-066).
296///
297/// Computes: output = input * rsqrt(mean(input^2) + eps) * gamma
298///
299/// Uses BatchedVectorizedRmsNormKernel — 8 warps per block, processes
300/// seq_len rows in parallel via Grid.y.
301///
302/// # Contract (C-GPUNORM-001)
303///
304/// - **Precondition**: input has batch_size * hidden_size elements, gamma has hidden_size elements
305/// - **Postcondition**: output contains RMSNorm(input) * gamma
306/// - **Invariant**: Same numerical result as CPU norm.forward_batched (within fp32 precision)
307#[cfg(feature = "cuda")]
308pub fn rms_norm_forward(
309    input: &GpuBuffer<f32>,
310    gamma: &GpuBuffer<f32>,
311    output: &mut GpuBuffer<f32>,
312    batch_size: u32,
313    hidden_size: u32,
314    stream: &CudaStream,
315) -> Result<()> {
316    let cache = KERNEL_CACHE.get().ok_or(CudaTensorError::DeviceNotInitialized)?;
317    let mut cache = cache.lock().map_err(|_err| {
318        CudaTensorError::KernelError("Failed to acquire kernel cache lock".to_string())
319    })?;
320
321    let key = format!("batched_rmsnorm_fwd_{hidden_size}");
322    let module = match cache.get_cached(&key) {
323        Some(m) => m,
324        None => {
325            let kernel = BatchedVectorizedRmsNormKernel::new(hidden_size, batch_size);
326            let ptx = kernel.emit_ptx_for_target(cache.sm_target());
327            cache.get_or_compile(&key, &ptx)?
328        }
329    };
330
331    // Grid: (1, batch_size, 1) — one block per row, each block processes one row
332    // Block: (256, 1, 1) — 8 warps per block for parallel reduction
333    let config = LaunchConfig {
334        grid: (1, batch_size, 1),
335        block: (256, 1, 1),
336        shared_mem: 8 * 4, // 8 warp partial sums
337    };
338
339    let input_ptr = input.as_ptr();
340    let output_ptr = output.as_ptr();
341    let gamma_ptr = gamma.as_ptr();
342
343    let mut args: [*mut std::ffi::c_void; 3] = [
344        &input_ptr as *const _ as *mut _,
345        &output_ptr as *const _ as *mut _,
346        &gamma_ptr as *const _ as *mut _,
347    ];
348
349    // SAFETY: Kernel launch requires FFI. input has batch_size * hidden_size elements,
350    // output has batch_size * hidden_size elements, gamma has hidden_size elements.
351    // Parameters match PTX signature (u64 input_ptr, u64 output_ptr, u64 gamma_ptr).
352    unsafe {
353        stream.launch_kernel(module, "batched_rmsnorm_vectorized", &config, &mut args).map_err(
354            |e| CudaTensorError::KernelError(format!("RMSNorm forward launch failed: {e:?}")),
355        )?;
356    }
357
358    Ok(())
359}
360
361/// LayerNorm backward pass on GPU
362///
363/// Computes gradients for input, gamma, and beta parameters
364#[cfg(feature = "cuda")]
365pub fn layer_norm_backward(
366    input: &GpuBuffer<f32>,
367    gamma: &GpuBuffer<f32>,
368    grad_output: &GpuBuffer<f32>,
369    grad_input: &mut GpuBuffer<f32>,
370    grad_gamma: &mut GpuBuffer<f32>,
371    grad_beta: &mut GpuBuffer<f32>,
372    batch_size: u32,
373    hidden_size: u32,
374    stream: &CudaStream,
375) -> Result<()> {
376    let cache = KERNEL_CACHE.get().ok_or(CudaTensorError::DeviceNotInitialized)?;
377    let mut cache = cache.lock().map_err(|_err| {
378        CudaTensorError::KernelError("Failed to acquire kernel cache lock".to_string())
379    })?;
380
381    let key = format!("layer_norm_backward_{batch_size}_{hidden_size}");
382    let module = match cache.get_cached(&key) {
383        Some(m) => m,
384        None => {
385            let kernel = LayerNormBackwardKernel::new(batch_size, hidden_size);
386            let ptx = kernel.emit_ptx_for_target(cache.sm_target());
387            cache.get_or_compile(&key, &ptx)?
388        }
389    };
390
391    let config = LaunchConfig {
392        grid: (batch_size, 1, 1),
393        block: (256.min(hidden_size), 1, 1),
394        shared_mem: 0,
395    };
396
397    let input_ptr = input.as_ptr();
398    let gamma_ptr = gamma.as_ptr();
399    let grad_out_ptr = grad_output.as_ptr();
400    let grad_in_ptr = grad_input.as_ptr();
401    let grad_gamma_ptr = grad_gamma.as_ptr();
402    let grad_beta_ptr = grad_beta.as_ptr();
403
404    let mut args: [*mut std::ffi::c_void; 8] = [
405        &input_ptr as *const _ as *mut _,
406        &gamma_ptr as *const _ as *mut _,
407        &grad_out_ptr as *const _ as *mut _,
408        &grad_in_ptr as *const _ as *mut _,
409        &grad_gamma_ptr as *const _ as *mut _,
410        &grad_beta_ptr as *const _ as *mut _,
411        &batch_size as *const _ as *mut _,
412        &hidden_size as *const _ as *mut _,
413    ];
414
415    // SAFETY: Kernel launch requires FFI. All buffers are valid GPU allocations with
416    // matching sizes, and the kernel parameters match the expected PTX signature.
417    unsafe {
418        stream.launch_kernel(module, "layer_norm_backward", &config, &mut args).map_err(|e| {
419            CudaTensorError::KernelError(format!("LayerNorm backward launch failed: {e:?}"))
420        })?;
421    }
422
423    Ok(())
424}