entrenar/autograd/cuda_backward/
structured.rs1#![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#[cfg(feature = "cuda")]
28#[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 let config = LaunchConfig {
60 grid: (batch_size, 1, 1),
61 block: (32, 1, 1), 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 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#[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 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 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 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#[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 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 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 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 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 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 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 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 drop(grad_gamma_partial);
292 Ok(())
293}
294
295#[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 let config = LaunchConfig {
334 grid: (1, batch_size, 1),
335 block: (256, 1, 1),
336 shared_mem: 8 * 4, };
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 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#[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 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}