1#[cfg(any(feature = "gpu", feature = "gpu-wasm"))]
9use super::super::runtime;
10use super::super::shaders;
11use super::GpuDevice;
12
13impl GpuDevice {
14 #[cfg(all(feature = "gpu", not(target_arch = "wasm32")))]
21 pub fn silu_backward(
22 &self,
23 input: &[f32],
24 grad_output: &[f32],
25 grad_input: &mut [f32],
26 ) -> Result<(), String> {
27 runtime::block_on(self.silu_backward_async(input, grad_output, grad_input))
28 }
29
30 pub async fn silu_backward_async(
32 &self,
33 input: &[f32],
34 grad_output: &[f32],
35 grad_input: &mut [f32],
36 ) -> Result<(), String> {
37 let n = input.len();
38 if grad_output.len() != n || grad_input.len() != n {
39 return Err(format!(
40 "SiLU backward: length mismatch: input={}, grad_output={}, grad_input={}",
41 n,
42 grad_output.len(),
43 grad_input.len()
44 ));
45 }
46
47 self.execute_backward_elementwise(
48 "SiLU Backward",
49 shaders::backward::SILU_BACKWARD_SHADER,
50 input,
51 grad_output,
52 grad_input,
53 n as u32,
54 )
55 .await
56 }
57
58 async fn execute_backward_elementwise(
62 &self,
63 op_name: &str,
64 shader_source: &str,
65 input: &[f32],
66 grad_output: &[f32],
67 grad_input: &mut [f32],
68 n: u32,
69 ) -> Result<(), String> {
70 use wgpu;
71
72 let shader = self.device.create_shader_module(wgpu::ShaderModuleDescriptor {
73 label: Some(&format!("{op_name} Shader")),
74 source: wgpu::ShaderSource::Wgsl(shader_source.into()),
75 });
76
77 let input_buf = self.create_storage_buffer(&format!("{op_name} input"), input, true);
79 let grad_out_buf =
80 self.create_storage_buffer(&format!("{op_name} grad_output"), grad_output, true);
81 let grad_in_buf = self.create_rw_storage_buffer(
82 &format!("{op_name} grad_input"),
83 (grad_input.len() * 4) as u64,
84 );
85
86 let uniform_data: [u32; 4] = [n, 0, 0, 0];
88 let uniform_buf = self.create_uniform_buffer(
89 &format!("{op_name} uniform"),
90 bytemuck::cast_slice(&uniform_data),
91 );
92
93 let bgl = self.device.create_bind_group_layout(&wgpu::BindGroupLayoutDescriptor {
95 label: Some(&format!("{op_name} BGL")),
96 entries: &[
97 storage_entry(0, true),
98 storage_entry(1, true),
99 storage_entry(2, false),
100 uniform_entry(3),
101 ],
102 });
103
104 let bg = self.device.create_bind_group(&wgpu::BindGroupDescriptor {
105 label: Some(&format!("{op_name} BG")),
106 layout: &bgl,
107 entries: &[
108 wgpu::BindGroupEntry { binding: 0, resource: input_buf.as_entire_binding() },
109 wgpu::BindGroupEntry { binding: 1, resource: grad_out_buf.as_entire_binding() },
110 wgpu::BindGroupEntry { binding: 2, resource: grad_in_buf.as_entire_binding() },
111 wgpu::BindGroupEntry { binding: 3, resource: uniform_buf.as_entire_binding() },
112 ],
113 });
114
115 let pipeline_layout = self.device.create_pipeline_layout(&wgpu::PipelineLayoutDescriptor {
116 label: Some(&format!("{op_name} PL")),
117 bind_group_layouts: &[&bgl],
118 push_constant_ranges: &[],
119 });
120
121 let pipeline = self.device.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
122 label: Some(&format!("{op_name} Pipeline")),
123 layout: Some(&pipeline_layout),
124 module: &shader,
125 entry_point: Some("main"),
126 compilation_options: Default::default(),
127 cache: None,
128 });
129
130 let staging = self.device.create_buffer(&wgpu::BufferDescriptor {
132 label: Some(&format!("{op_name} Staging")),
133 size: (grad_input.len() * 4) as u64,
134 usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
135 mapped_at_creation: false,
136 });
137
138 let mut encoder =
140 self.device.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
141 {
142 let mut pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
143 pass.set_pipeline(&pipeline);
144 pass.set_bind_group(0, &bg, &[]);
145 let total_wg = n.div_ceil(256);
147 pass.dispatch_workgroups(total_wg.min(65535), total_wg.div_ceil(65535), 1);
148 }
149 encoder.copy_buffer_to_buffer(&grad_in_buf, 0, &staging, 0, (grad_input.len() * 4) as u64);
150 self.queue.submit(Some(encoder.finish()));
151
152 let slice = staging.slice(..);
154 let (sender, receiver) = futures_intrusive::channel::shared::oneshot_channel();
155 slice.map_async(wgpu::MapMode::Read, move |r| {
156 sender.send(r).ok();
157 });
158 self.device.poll(wgpu::PollType::Wait { submission_index: None, timeout: None }).ok();
159 receiver
160 .receive()
161 .await
162 .ok_or_else(|| format!("{op_name}: map_async cancelled"))?
163 .map_err(|e| format!("{op_name}: map_async failed: {e}"))?;
164
165 let data = slice.get_mapped_range();
166 grad_input.copy_from_slice(bytemuck::cast_slice(&data));
167 drop(data);
168 staging.unmap();
169
170 Ok(())
171 }
172
173 fn create_storage_buffer(&self, label: &str, data: &[f32], read_only: bool) -> wgpu::Buffer {
176 let buf = self.device.create_buffer(&wgpu::BufferDescriptor {
177 label: Some(label),
178 size: (data.len() * 4) as u64,
179 usage: wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_DST,
180 mapped_at_creation: false,
181 });
182 self.queue.write_buffer(&buf, 0, bytemuck::cast_slice(data));
183 let _ = read_only; buf
185 }
186
187 fn create_rw_storage_buffer(&self, label: &str, size: u64) -> wgpu::Buffer {
188 self.device.create_buffer(&wgpu::BufferDescriptor {
189 label: Some(label),
190 size,
191 usage: wgpu::BufferUsages::STORAGE
192 | wgpu::BufferUsages::COPY_SRC
193 | wgpu::BufferUsages::COPY_DST,
194 mapped_at_creation: false,
195 })
196 }
197
198 fn create_uniform_buffer(&self, label: &str, data: &[u8]) -> wgpu::Buffer {
199 let buf = self.device.create_buffer(&wgpu::BufferDescriptor {
200 label: Some(label),
201 size: data.len() as u64,
202 usage: wgpu::BufferUsages::UNIFORM | wgpu::BufferUsages::COPY_DST,
203 mapped_at_creation: false,
204 });
205 self.queue.write_buffer(&buf, 0, data);
206 buf
207 }
208
209 #[cfg(all(feature = "gpu", not(target_arch = "wasm32")))]
216 pub fn gemm_backward_a(
217 &self,
218 grad_c: &[f32],
219 b: &[f32],
220 grad_a: &mut [f32],
221 m: u32,
222 k: u32,
223 n: u32,
224 ) -> Result<(), String> {
225 runtime::block_on(self.gemm_backward_a_async(grad_c, b, grad_a, m, k, n))
226 }
227
228 pub async fn gemm_backward_a_async(
230 &self,
231 grad_c: &[f32],
232 b: &[f32],
233 grad_a: &mut [f32],
234 m: u32,
235 k: u32,
236 n: u32,
237 ) -> Result<(), String> {
238 self.execute_backward_gemm(
239 "GEMM Backward A",
240 shaders::backward::GEMM_BACKWARD_A_SHADER,
241 grad_c,
242 b,
243 grad_a,
244 m,
245 k,
246 n,
247 )
248 .await
249 }
250
251 #[cfg(all(feature = "gpu", not(target_arch = "wasm32")))]
253 pub fn gemm_backward_b(
254 &self,
255 a: &[f32],
256 grad_c: &[f32],
257 grad_b: &mut [f32],
258 m: u32,
259 k: u32,
260 n: u32,
261 ) -> Result<(), String> {
262 runtime::block_on(self.gemm_backward_b_async(a, grad_c, grad_b, m, k, n))
263 }
264
265 pub async fn gemm_backward_b_async(
267 &self,
268 a: &[f32],
269 grad_c: &[f32],
270 grad_b: &mut [f32],
271 m: u32,
272 k: u32,
273 n: u32,
274 ) -> Result<(), String> {
275 self.execute_backward_gemm(
276 "GEMM Backward B",
277 shaders::backward::GEMM_BACKWARD_B_SHADER,
278 a,
279 grad_c,
280 grad_b,
281 m,
282 k,
283 n,
284 )
285 .await
286 }
287
288 async fn execute_backward_gemm(
292 &self,
293 op_name: &str,
294 shader_source: &str,
295 buf_a: &[f32],
296 buf_b: &[f32],
297 output: &mut [f32],
298 m: u32,
299 k: u32,
300 n: u32,
301 ) -> Result<(), String> {
302 use wgpu;
303
304 let shader = self.device.create_shader_module(wgpu::ShaderModuleDescriptor {
305 label: Some(&format!("{op_name} Shader")),
306 source: wgpu::ShaderSource::Wgsl(shader_source.into()),
307 });
308
309 let a_buf = self.create_storage_buffer(&format!("{op_name} A"), buf_a, true);
310 let b_buf = self.create_storage_buffer(&format!("{op_name} B"), buf_b, true);
311 let out_buf =
312 self.create_rw_storage_buffer(&format!("{op_name} Output"), (output.len() * 4) as u64);
313
314 let dims: [u32; 4] = [m, k, n, 0];
316 let uniform_buf =
317 self.create_uniform_buffer(&format!("{op_name} Dims"), bytemuck::cast_slice(&dims));
318
319 let bgl = self.device.create_bind_group_layout(&wgpu::BindGroupLayoutDescriptor {
320 label: None,
321 entries: &[
322 storage_entry(0, true),
323 storage_entry(1, true),
324 storage_entry(2, false),
325 uniform_entry(3),
326 ],
327 });
328
329 let bg = self.device.create_bind_group(&wgpu::BindGroupDescriptor {
330 label: None,
331 layout: &bgl,
332 entries: &[
333 wgpu::BindGroupEntry { binding: 0, resource: a_buf.as_entire_binding() },
334 wgpu::BindGroupEntry { binding: 1, resource: b_buf.as_entire_binding() },
335 wgpu::BindGroupEntry { binding: 2, resource: out_buf.as_entire_binding() },
336 wgpu::BindGroupEntry { binding: 3, resource: uniform_buf.as_entire_binding() },
337 ],
338 });
339
340 let pl = self.device.create_pipeline_layout(&wgpu::PipelineLayoutDescriptor {
341 label: None,
342 bind_group_layouts: &[&bgl],
343 push_constant_ranges: &[],
344 });
345
346 let pipeline = self.device.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
347 label: Some(&format!("{op_name} Pipeline")),
348 layout: Some(&pl),
349 module: &shader,
350 entry_point: Some("main"),
351 compilation_options: Default::default(),
352 cache: None,
353 });
354
355 let staging = self.device.create_buffer(&wgpu::BufferDescriptor {
356 label: None,
357 size: (output.len() * 4) as u64,
358 usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
359 mapped_at_creation: false,
360 });
361
362 let mut encoder =
363 self.device.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
364 {
365 let mut pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
366 pass.set_pipeline(&pipeline);
367 pass.set_bind_group(0, &bg, &[]);
368
369 let out_rows = if op_name.contains("A") { m } else { k };
373 let out_cols = if op_name.contains("A") { k } else { n };
374 pass.dispatch_workgroups(out_rows.div_ceil(16), out_cols.div_ceil(16), 1);
375 }
376 encoder.copy_buffer_to_buffer(&out_buf, 0, &staging, 0, (output.len() * 4) as u64);
377 self.queue.submit(Some(encoder.finish()));
378
379 let slice = staging.slice(..);
380 let (sender, receiver) = futures_intrusive::channel::shared::oneshot_channel();
381 slice.map_async(wgpu::MapMode::Read, move |r| {
382 sender.send(r).ok();
383 });
384 self.device.poll(wgpu::PollType::Wait { submission_index: None, timeout: None }).ok();
385 receiver
386 .receive()
387 .await
388 .ok_or_else(|| format!("{op_name}: map cancelled"))?
389 .map_err(|e| format!("{op_name}: map failed: {e}"))?;
390
391 let data = slice.get_mapped_range();
392 output.copy_from_slice(bytemuck::cast_slice(&data));
393 drop(data);
394 staging.unmap();
395
396 Ok(())
397 }
398
399 #[cfg(all(feature = "gpu", not(target_arch = "wasm32")))]
403 pub fn rope_backward(
404 &self,
405 grad_output: &[f32],
406 grad_input: &mut [f32],
407 num_heads: u32,
408 head_dim: u32,
409 seq_len: u32,
410 theta: f32,
411 ) -> Result<(), String> {
412 runtime::block_on(self.rope_backward_async(
413 grad_output,
414 grad_input,
415 num_heads,
416 head_dim,
417 seq_len,
418 theta,
419 ))
420 }
421
422 pub async fn rope_backward_async(
424 &self,
425 grad_output: &[f32],
426 grad_input: &mut [f32],
427 num_heads: u32,
428 head_dim: u32,
429 seq_len: u32,
430 theta: f32,
431 ) -> Result<(), String> {
432 use wgpu;
433
434 let n = grad_output.len();
435 let total_pairs = num_heads * seq_len * (head_dim / 2);
436
437 let shader = self.device.create_shader_module(wgpu::ShaderModuleDescriptor {
438 label: Some("RoPE Backward Shader"),
439 source: wgpu::ShaderSource::Wgsl(shaders::backward::ROPE_BACKWARD_SHADER.into()),
440 });
441
442 let go_buf = self.create_storage_buffer("rope_bwd grad_out", grad_output, true);
443 let gi_buf = self.create_rw_storage_buffer("rope_bwd grad_in", (n * 4) as u64);
444
445 let params: [u32; 4] = [num_heads, head_dim, seq_len, theta.log2().to_bits()];
447 let uniform_buf =
448 self.create_uniform_buffer("rope_bwd params", bytemuck::cast_slice(¶ms));
449
450 let bgl = self.device.create_bind_group_layout(&wgpu::BindGroupLayoutDescriptor {
451 label: None,
452 entries: &[storage_entry(0, true), storage_entry(1, false), uniform_entry(2)],
453 });
454 let bg = self.device.create_bind_group(&wgpu::BindGroupDescriptor {
455 label: None,
456 layout: &bgl,
457 entries: &[
458 wgpu::BindGroupEntry { binding: 0, resource: go_buf.as_entire_binding() },
459 wgpu::BindGroupEntry { binding: 1, resource: gi_buf.as_entire_binding() },
460 wgpu::BindGroupEntry { binding: 2, resource: uniform_buf.as_entire_binding() },
461 ],
462 });
463
464 let pl = self.device.create_pipeline_layout(&wgpu::PipelineLayoutDescriptor {
465 label: None,
466 bind_group_layouts: &[&bgl],
467 push_constant_ranges: &[],
468 });
469 let pipeline = self.device.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
470 label: Some("RoPE Backward"),
471 layout: Some(&pl),
472 module: &shader,
473 entry_point: Some("main"),
474 compilation_options: Default::default(),
475 cache: None,
476 });
477
478 let staging = self.device.create_buffer(&wgpu::BufferDescriptor {
479 label: None,
480 size: (n * 4) as u64,
481 usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
482 mapped_at_creation: false,
483 });
484
485 let mut encoder =
486 self.device.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
487 {
488 let mut pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
489 pass.set_pipeline(&pipeline);
490 pass.set_bind_group(0, &bg, &[]);
491 let total_wg = total_pairs.div_ceil(256);
492 pass.dispatch_workgroups(total_wg.min(65535), total_wg.div_ceil(65535), 1);
493 }
494 encoder.copy_buffer_to_buffer(&gi_buf, 0, &staging, 0, (n * 4) as u64);
495 self.queue.submit(Some(encoder.finish()));
496
497 let slice = staging.slice(..);
498 let (sender, receiver) = futures_intrusive::channel::shared::oneshot_channel();
499 slice.map_async(wgpu::MapMode::Read, move |r| {
500 sender.send(r).ok();
501 });
502 self.device.poll(wgpu::PollType::Wait { submission_index: None, timeout: None }).ok();
503 receiver
504 .receive()
505 .await
506 .ok_or("RoPE backward: cancelled".to_string())?
507 .map_err(|e| format!("RoPE backward: {e}"))?;
508 let data = slice.get_mapped_range();
509 grad_input.copy_from_slice(bytemuck::cast_slice(&data));
510 drop(data);
511 staging.unmap();
512 Ok(())
513 }
514
515 #[cfg(all(feature = "gpu", not(target_arch = "wasm32")))]
517 pub fn adamw_step(
518 &self,
519 params: &mut [f32],
520 grads: &[f32],
521 m: &mut [f32],
522 v: &mut [f32],
523 lr: f32,
524 beta1: f32,
525 beta2: f32,
526 eps: f32,
527 weight_decay: f32,
528 step: u32,
529 ) -> Result<(), String> {
530 runtime::block_on(self.adamw_step_async(
531 params,
532 grads,
533 m,
534 v,
535 lr,
536 beta1,
537 beta2,
538 eps,
539 weight_decay,
540 step,
541 ))
542 }
543
544 pub async fn adamw_step_async(
546 &self,
547 params: &mut [f32],
548 grads: &[f32],
549 m: &mut [f32],
550 v: &mut [f32],
551 lr: f32,
552 beta1: f32,
553 beta2: f32,
554 eps: f32,
555 weight_decay: f32,
556 step: u32,
557 ) -> Result<(), String> {
558 use wgpu;
559
560 let n = params.len() as u32;
561 let bc1 = 1.0 - beta1.powi(step as i32);
562 let bc2 = 1.0 - beta2.powi(step as i32);
563
564 let shader = self.device.create_shader_module(wgpu::ShaderModuleDescriptor {
565 label: Some("AdamW Step"),
566 source: wgpu::ShaderSource::Wgsl(shaders::backward::ADAMW_STEP_SHADER.into()),
567 });
568
569 let params_buf = self.device.create_buffer(&wgpu::BufferDescriptor {
571 label: Some("adamw params"),
572 size: (params.len() * 4) as u64,
573 usage: wgpu::BufferUsages::STORAGE
574 | wgpu::BufferUsages::COPY_DST
575 | wgpu::BufferUsages::COPY_SRC,
576 mapped_at_creation: false,
577 });
578 self.queue.write_buffer(¶ms_buf, 0, bytemuck::cast_slice(params));
579
580 let grads_buf = self.create_storage_buffer("adamw grads", grads, true);
581
582 let m_buf = self.device.create_buffer(&wgpu::BufferDescriptor {
583 label: Some("adamw m"),
584 size: (m.len() * 4) as u64,
585 usage: wgpu::BufferUsages::STORAGE
586 | wgpu::BufferUsages::COPY_DST
587 | wgpu::BufferUsages::COPY_SRC,
588 mapped_at_creation: false,
589 });
590 self.queue.write_buffer(&m_buf, 0, bytemuck::cast_slice(m));
591
592 let v_buf = self.device.create_buffer(&wgpu::BufferDescriptor {
593 label: Some("adamw v"),
594 size: (v.len() * 4) as u64,
595 usage: wgpu::BufferUsages::STORAGE
596 | wgpu::BufferUsages::COPY_DST
597 | wgpu::BufferUsages::COPY_SRC,
598 mapped_at_creation: false,
599 });
600 self.queue.write_buffer(&v_buf, 0, bytemuck::cast_slice(v));
601
602 let hp: [u32; 8] = [
605 n,
606 lr.to_bits(),
607 beta1.to_bits(),
608 beta2.to_bits(),
609 eps.to_bits(),
610 weight_decay.to_bits(),
611 bc1.to_bits(),
612 bc2.to_bits(),
613 ];
614 let uniform_buf = self.create_uniform_buffer("adamw hp", bytemuck::cast_slice(&hp));
615
616 let bgl = self.device.create_bind_group_layout(&wgpu::BindGroupLayoutDescriptor {
617 label: None,
618 entries: &[
619 storage_entry(0, false), storage_entry(1, true), storage_entry(2, false), storage_entry(3, false), uniform_entry(4),
624 ],
625 });
626 let bg = self.device.create_bind_group(&wgpu::BindGroupDescriptor {
627 label: None,
628 layout: &bgl,
629 entries: &[
630 wgpu::BindGroupEntry { binding: 0, resource: params_buf.as_entire_binding() },
631 wgpu::BindGroupEntry { binding: 1, resource: grads_buf.as_entire_binding() },
632 wgpu::BindGroupEntry { binding: 2, resource: m_buf.as_entire_binding() },
633 wgpu::BindGroupEntry { binding: 3, resource: v_buf.as_entire_binding() },
634 wgpu::BindGroupEntry { binding: 4, resource: uniform_buf.as_entire_binding() },
635 ],
636 });
637
638 let pl = self.device.create_pipeline_layout(&wgpu::PipelineLayoutDescriptor {
639 label: None,
640 bind_group_layouts: &[&bgl],
641 push_constant_ranges: &[],
642 });
643 let pipeline = self.device.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
644 label: Some("AdamW"),
645 layout: Some(&pl),
646 module: &shader,
647 entry_point: Some("main"),
648 compilation_options: Default::default(),
649 cache: None,
650 });
651
652 let params_staging = self.device.create_buffer(&wgpu::BufferDescriptor {
654 label: None,
655 size: (params.len() * 4) as u64,
656 usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
657 mapped_at_creation: false,
658 });
659 let m_staging = self.device.create_buffer(&wgpu::BufferDescriptor {
660 label: None,
661 size: (m.len() * 4) as u64,
662 usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
663 mapped_at_creation: false,
664 });
665 let v_staging = self.device.create_buffer(&wgpu::BufferDescriptor {
666 label: None,
667 size: (v.len() * 4) as u64,
668 usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
669 mapped_at_creation: false,
670 });
671
672 let mut encoder =
673 self.device.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
674 {
675 let mut pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
676 pass.set_pipeline(&pipeline);
677 pass.set_bind_group(0, &bg, &[]);
678 let total_wg = n.div_ceil(256);
680 pass.dispatch_workgroups(total_wg.min(65535), total_wg.div_ceil(65535), 1);
681 }
682 encoder.copy_buffer_to_buffer(
683 ¶ms_buf,
684 0,
685 ¶ms_staging,
686 0,
687 (params.len() * 4) as u64,
688 );
689 encoder.copy_buffer_to_buffer(&m_buf, 0, &m_staging, 0, (m.len() * 4) as u64);
690 encoder.copy_buffer_to_buffer(&v_buf, 0, &v_staging, 0, (v.len() * 4) as u64);
691 self.queue.submit(Some(encoder.finish()));
692
693 let read_buf = |staging: &wgpu::Buffer, out: &mut [f32]| -> Result<(), String> {
695 let slice = staging.slice(..);
696 let (tx, rx) = std::sync::mpsc::channel();
697 slice.map_async(wgpu::MapMode::Read, move |r| {
698 tx.send(r).ok();
699 });
700 self.device.poll(wgpu::PollType::Wait { submission_index: None, timeout: None }).ok();
701 rx.recv()
702 .map_err(|e| format!("AdamW readback: {e}"))?
703 .map_err(|e| format!("AdamW map: {e}"))?;
704 let data = slice.get_mapped_range();
705 out.copy_from_slice(bytemuck::cast_slice(&data));
706 drop(data);
707 staging.unmap();
708 Ok(())
709 };
710 read_buf(¶ms_staging, params)?;
711 read_buf(&m_staging, m)?;
712 read_buf(&v_staging, v)?;
713
714 Ok(())
715 }
716
717 #[cfg(all(feature = "gpu", not(target_arch = "wasm32")))]
724 pub fn rmsnorm_backward(
725 &self,
726 input: &[f32],
727 gamma: &[f32],
728 grad_output: &[f32],
729 grad_input: &mut [f32],
730 grad_gamma: &mut [f32],
731 num_rows: u32,
732 hidden_dim: u32,
733 eps: f32,
734 ) -> Result<(), String> {
735 runtime::block_on(self.rmsnorm_backward_async(
736 input,
737 gamma,
738 grad_output,
739 grad_input,
740 grad_gamma,
741 num_rows,
742 hidden_dim,
743 eps,
744 ))
745 }
746
747 pub async fn rmsnorm_backward_async(
749 &self,
750 input: &[f32],
751 gamma: &[f32],
752 grad_output: &[f32],
753 grad_input: &mut [f32],
754 grad_gamma: &mut [f32],
755 num_rows: u32,
756 hidden_dim: u32,
757 eps: f32,
758 ) -> Result<(), String> {
759 use wgpu;
760
761 let shader = self.device.create_shader_module(wgpu::ShaderModuleDescriptor {
762 label: Some("RMSNorm Backward"),
763 source: wgpu::ShaderSource::Wgsl(shaders::backward::RMSNORM_BACKWARD_SHADER.into()),
764 });
765
766 let input_buf = self.create_storage_buffer("rms_bwd input", input, true);
767 let gamma_buf = self.create_storage_buffer("rms_bwd gamma", gamma, true);
768 let grad_out_buf = self.create_storage_buffer("rms_bwd grad_out", grad_output, true);
769 let grad_in_buf =
770 self.create_rw_storage_buffer("rms_bwd grad_in", (grad_input.len() * 4) as u64);
771
772 let grad_gamma_buf = self.device.create_buffer(&wgpu::BufferDescriptor {
774 label: Some("rms_bwd grad_gamma"),
775 size: (hidden_dim as usize * 4) as u64,
776 usage: wgpu::BufferUsages::STORAGE
777 | wgpu::BufferUsages::COPY_DST
778 | wgpu::BufferUsages::COPY_SRC,
779 mapped_at_creation: false,
780 });
781 let zeros = vec![0u8; hidden_dim as usize * 4];
783 self.queue.write_buffer(&grad_gamma_buf, 0, &zeros);
784
785 let params: [u32; 4] = [num_rows, hidden_dim, eps.to_bits(), 0];
787 let uniform_buf =
788 self.create_uniform_buffer("rms_bwd params", bytemuck::cast_slice(¶ms));
789
790 let bgl = self.device.create_bind_group_layout(&wgpu::BindGroupLayoutDescriptor {
791 label: None,
792 entries: &[
793 storage_entry(0, true), storage_entry(1, true), storage_entry(2, true), storage_entry(3, false), storage_entry(4, false), uniform_entry(5),
799 ],
800 });
801 let bg = self.device.create_bind_group(&wgpu::BindGroupDescriptor {
802 label: None,
803 layout: &bgl,
804 entries: &[
805 wgpu::BindGroupEntry { binding: 0, resource: input_buf.as_entire_binding() },
806 wgpu::BindGroupEntry { binding: 1, resource: gamma_buf.as_entire_binding() },
807 wgpu::BindGroupEntry { binding: 2, resource: grad_out_buf.as_entire_binding() },
808 wgpu::BindGroupEntry { binding: 3, resource: grad_in_buf.as_entire_binding() },
809 wgpu::BindGroupEntry { binding: 4, resource: grad_gamma_buf.as_entire_binding() },
810 wgpu::BindGroupEntry { binding: 5, resource: uniform_buf.as_entire_binding() },
811 ],
812 });
813
814 let pl = self.device.create_pipeline_layout(&wgpu::PipelineLayoutDescriptor {
815 label: None,
816 bind_group_layouts: &[&bgl],
817 push_constant_ranges: &[],
818 });
819 let pipeline = self.device.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
820 label: Some("RMSNorm Backward"),
821 layout: Some(&pl),
822 module: &shader,
823 entry_point: Some("main"),
824 compilation_options: Default::default(),
825 cache: None,
826 });
827
828 let gi_staging = self.device.create_buffer(&wgpu::BufferDescriptor {
830 label: None,
831 size: (grad_input.len() * 4) as u64,
832 usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
833 mapped_at_creation: false,
834 });
835 let gg_staging = self.device.create_buffer(&wgpu::BufferDescriptor {
836 label: None,
837 size: (hidden_dim as usize * 4) as u64,
838 usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
839 mapped_at_creation: false,
840 });
841
842 let mut encoder =
843 self.device.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
844 {
845 let mut pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
846 pass.set_pipeline(&pipeline);
847 pass.set_bind_group(0, &bg, &[]);
848 pass.dispatch_workgroups(num_rows, 1, 1);
850 }
851 encoder.copy_buffer_to_buffer(
852 &grad_in_buf,
853 0,
854 &gi_staging,
855 0,
856 (grad_input.len() * 4) as u64,
857 );
858 encoder.copy_buffer_to_buffer(
859 &grad_gamma_buf,
860 0,
861 &gg_staging,
862 0,
863 (hidden_dim as usize * 4) as u64,
864 );
865 self.queue.submit(Some(encoder.finish()));
866
867 {
869 let slice = gi_staging.slice(..);
870 let (tx, rx) = std::sync::mpsc::channel();
871 slice.map_async(wgpu::MapMode::Read, move |r| {
872 tx.send(r).ok();
873 });
874 self.device.poll(wgpu::PollType::Wait { submission_index: None, timeout: None }).ok();
875 rx.recv()
876 .map_err(|e| format!("RMSNorm bwd gi: {e}"))?
877 .map_err(|e| format!("RMSNorm bwd gi map: {e}"))?;
878 let data = slice.get_mapped_range();
879 grad_input.copy_from_slice(bytemuck::cast_slice(&data));
880 drop(data);
881 gi_staging.unmap();
882 }
883 {
885 let slice = gg_staging.slice(..);
886 let (tx, rx) = std::sync::mpsc::channel();
887 slice.map_async(wgpu::MapMode::Read, move |r| {
888 tx.send(r).ok();
889 });
890 self.device.poll(wgpu::PollType::Wait { submission_index: None, timeout: None }).ok();
891 rx.recv()
892 .map_err(|e| format!("RMSNorm bwd gg: {e}"))?
893 .map_err(|e| format!("RMSNorm bwd gg map: {e}"))?;
894 let data = slice.get_mapped_range();
895 let raw: &[u32] = bytemuck::cast_slice(&data);
897 for (i, &bits) in raw.iter().enumerate() {
898 grad_gamma[i] = f32::from_bits(bits);
899 }
900 drop(data);
901 gg_staging.unmap();
902 }
903
904 Ok(())
905 }
906
907 #[cfg(all(feature = "gpu", not(target_arch = "wasm32")))]
913 pub fn nf4_dequant(
914 &self,
915 packed: &[u32],
916 scales: &[f32],
917 output: &mut [f32],
918 n: u32,
919 block_size: u32,
920 ) -> Result<(), String> {
921 runtime::block_on(self.nf4_dequant_async(packed, scales, output, n, block_size))
922 }
923
924 pub async fn nf4_dequant_async(
926 &self,
927 packed: &[u32],
928 scales: &[f32],
929 output: &mut [f32],
930 n: u32,
931 block_size: u32,
932 ) -> Result<(), String> {
933 use wgpu;
934
935 let shader = self.device.create_shader_module(wgpu::ShaderModuleDescriptor {
936 label: Some("NF4 Dequant"),
937 source: wgpu::ShaderSource::Wgsl(shaders::backward::NF4_DEQUANT_SHADER.into()),
938 });
939
940 let packed_buf = self.device.create_buffer(&wgpu::BufferDescriptor {
941 label: Some("nf4 packed"),
942 size: (packed.len() * 4) as u64,
943 usage: wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_DST,
944 mapped_at_creation: false,
945 });
946 self.queue.write_buffer(&packed_buf, 0, bytemuck::cast_slice(packed));
947
948 let scales_buf = self.create_storage_buffer("nf4 scales", scales, true);
949 let output_buf = self.create_rw_storage_buffer("nf4 output", (output.len() * 4) as u64);
950
951 let params: [u32; 4] = [n, block_size, 0, 0];
952 let uniform_buf = self.create_uniform_buffer("nf4 params", bytemuck::cast_slice(¶ms));
953
954 let bgl = self.device.create_bind_group_layout(&wgpu::BindGroupLayoutDescriptor {
955 label: None,
956 entries: &[
957 storage_entry(0, true), storage_entry(1, true), storage_entry(2, false), uniform_entry(3),
961 ],
962 });
963 let bg = self.device.create_bind_group(&wgpu::BindGroupDescriptor {
964 label: None,
965 layout: &bgl,
966 entries: &[
967 wgpu::BindGroupEntry { binding: 0, resource: packed_buf.as_entire_binding() },
968 wgpu::BindGroupEntry { binding: 1, resource: scales_buf.as_entire_binding() },
969 wgpu::BindGroupEntry { binding: 2, resource: output_buf.as_entire_binding() },
970 wgpu::BindGroupEntry { binding: 3, resource: uniform_buf.as_entire_binding() },
971 ],
972 });
973
974 let pl = self.device.create_pipeline_layout(&wgpu::PipelineLayoutDescriptor {
975 label: None,
976 bind_group_layouts: &[&bgl],
977 push_constant_ranges: &[],
978 });
979 let pipeline = self.device.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
980 label: Some("NF4 Dequant"),
981 layout: Some(&pl),
982 module: &shader,
983 entry_point: Some("main"),
984 compilation_options: Default::default(),
985 cache: None,
986 });
987
988 let staging = self.device.create_buffer(&wgpu::BufferDescriptor {
989 label: None,
990 size: (output.len() * 4) as u64,
991 usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
992 mapped_at_creation: false,
993 });
994
995 let mut encoder =
996 self.device.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
997 {
998 let mut pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
999 pass.set_pipeline(&pipeline);
1000 pass.set_bind_group(0, &bg, &[]);
1001 let total_wg = n.div_ceil(256);
1005 let x = total_wg.min(65535);
1006 let y = total_wg.div_ceil(65535);
1007 pass.dispatch_workgroups(x, y, 1);
1008 }
1009 encoder.copy_buffer_to_buffer(&output_buf, 0, &staging, 0, (output.len() * 4) as u64);
1010 self.queue.submit(Some(encoder.finish()));
1011
1012 let slice = staging.slice(..);
1013 let (sender, receiver) = futures_intrusive::channel::shared::oneshot_channel();
1014 slice.map_async(wgpu::MapMode::Read, move |r| {
1015 sender.send(r).ok();
1016 });
1017 self.device.poll(wgpu::PollType::Wait { submission_index: None, timeout: None }).ok();
1018 receiver
1019 .receive()
1020 .await
1021 .ok_or("NF4 dequant: cancelled".to_string())?
1022 .map_err(|e| format!("NF4 dequant: {e}"))?;
1023 let data = slice.get_mapped_range();
1024 output.copy_from_slice(bytemuck::cast_slice(&data));
1025 drop(data);
1026 staging.unmap();
1027
1028 Ok(())
1029 }
1030}
1031
1032fn storage_entry(binding: u32, read_only: bool) -> wgpu::BindGroupLayoutEntry {
1033 wgpu::BindGroupLayoutEntry {
1034 binding,
1035 visibility: wgpu::ShaderStages::COMPUTE,
1036 ty: wgpu::BindingType::Buffer {
1037 ty: wgpu::BufferBindingType::Storage { read_only },
1038 has_dynamic_offset: false,
1039 min_binding_size: None,
1040 },
1041 count: None,
1042 }
1043}
1044
1045fn uniform_entry(binding: u32) -> wgpu::BindGroupLayoutEntry {
1046 wgpu::BindGroupLayoutEntry {
1047 binding,
1048 visibility: wgpu::ShaderStages::COMPUTE,
1049 ty: wgpu::BindingType::Buffer {
1050 ty: wgpu::BufferBindingType::Uniform,
1051 has_dynamic_offset: false,
1052 min_binding_size: None,
1053 },
1054 count: None,
1055 }
1056}
1057
1058#[cfg(all(test, feature = "gpu"))]
1059mod tests {
1060 use super::*;
1061
1062 fn device_or_skip() -> Option<GpuDevice> {
1067 match GpuDevice::new() {
1068 Ok(device) => Some(device),
1069 Err(err) => {
1070 println!("SKIP: no GPU adapter on this host ({err}); nothing here is asserted");
1071 None
1072 }
1073 }
1074 }
1075
1076 fn silu_backward_cpu(input: &[f32], grad_output: &[f32]) -> Vec<f32> {
1078 input
1079 .iter()
1080 .zip(grad_output.iter())
1081 .map(|(&x, &dy)| {
1082 let sigmoid = 1.0 / (1.0 + (-x).exp());
1083 let y = x * sigmoid;
1084 let silu_prime = sigmoid * (1.0 + x - y);
1085 dy * silu_prime
1086 })
1087 .collect()
1088 }
1089
1090 #[test]
1092 fn test_falsify_wgpu_001_silu_backward_parity() {
1093 let Some(device) = device_or_skip() else { return };
1094
1095 let input: Vec<f32> = (-50..50).map(|i| i as f32 * 0.1).collect();
1096 let grad_output: Vec<f32> = (0..100).map(|i| (i as f32 - 50.0) * 0.01).collect();
1097 let expected = silu_backward_cpu(&input, &grad_output);
1098
1099 let mut grad_input = vec![0.0f32; 100];
1100 device.silu_backward(&input, &grad_output, &mut grad_input).expect("silu_backward");
1101
1102 let max_diff = grad_input
1103 .iter()
1104 .zip(expected.iter())
1105 .map(|(a, b)| (a - b).abs())
1106 .fold(0.0f32, f32::max);
1107
1108 assert!(
1109 max_diff < 1e-4,
1110 "FALSIFY-WGPU-001: SiLU backward max diff = {max_diff} (threshold: 1e-4)"
1111 );
1112 }
1113
1114 #[test]
1116 fn test_silu_backward_at_zero() {
1117 let Some(device) = device_or_skip() else { return };
1118
1119 let input = vec![0.0f32; 4];
1120 let grad_output = vec![1.0f32; 4];
1121 let mut grad_input = vec![0.0f32; 4];
1122
1123 device.silu_backward(&input, &grad_output, &mut grad_input).expect("silu_backward");
1124
1125 for &g in &grad_input {
1127 assert!((g - 0.5).abs() < 1e-5, "silu'(0) should be 0.5, got {g}");
1128 }
1129 }
1130
1131 #[test]
1133 fn test_silu_backward_length_mismatch() {
1134 let Some(device) = device_or_skip() else { return };
1135
1136 let input = vec![1.0f32; 10];
1137 let grad_output = vec![1.0f32; 5]; let mut grad_input = vec![0.0f32; 10];
1139
1140 let result = device.silu_backward(&input, &grad_output, &mut grad_input);
1141 assert!(result.is_err());
1142 }
1143
1144 fn matmul_cpu(a: &[f32], b: &[f32], m: usize, k: usize, n: usize) -> Vec<f32> {
1146 let mut c = vec![0.0f32; m * n];
1147 for i in 0..m {
1148 for j in 0..n {
1149 let mut sum = 0.0f32;
1150 for p in 0..k {
1151 sum += a[i * k + p] * b[p * n + j];
1152 }
1153 c[i * n + j] = sum;
1154 }
1155 }
1156 c
1157 }
1158
1159 #[test]
1164 fn test_falsify_wgpu_001_gemm_backward_a_parity() {
1165 let Some(device) = device_or_skip() else { return };
1166
1167 let (m, k, n) = (4, 8, 6);
1168
1169 let grad_c: Vec<f32> = (0..m * n).map(|i| (i as f32 - 12.0) * 0.1).collect();
1171 let b: Vec<f32> = (0..k * n).map(|i| (i as f32 - 24.0) * 0.05).collect();
1172
1173 let mut b_t = vec![0.0f32; n * k];
1176 for i in 0..k {
1177 for j in 0..n {
1178 b_t[j * k + i] = b[i * n + j];
1179 }
1180 }
1181 let expected = matmul_cpu(&grad_c, &b_t, m, n, k);
1182
1183 let mut grad_a = vec![0.0f32; m * k];
1184 device
1185 .gemm_backward_a(&grad_c, &b, &mut grad_a, m as u32, k as u32, n as u32)
1186 .expect("gemm_backward_a");
1187
1188 let max_diff =
1189 grad_a.iter().zip(expected.iter()).map(|(a, b)| (a - b).abs()).fold(0.0f32, f32::max);
1190
1191 assert!(
1192 max_diff < 1e-3,
1193 "FALSIFY-WGPU-001: GEMM backward A max diff = {max_diff} (threshold: 1e-3)"
1194 );
1195 }
1196
1197 #[test]
1201 fn test_falsify_wgpu_001_gemm_backward_b_parity() {
1202 let Some(device) = device_or_skip() else { return };
1203
1204 let (m, k, n) = (4, 8, 6);
1205
1206 let a: Vec<f32> = (0..m * k).map(|i| (i as f32 - 16.0) * 0.1).collect();
1207 let grad_c: Vec<f32> = (0..m * n).map(|i| (i as f32 - 12.0) * 0.05).collect();
1208
1209 let mut a_t = vec![0.0f32; k * m];
1211 for i in 0..m {
1212 for j in 0..k {
1213 a_t[j * m + i] = a[i * k + j];
1214 }
1215 }
1216 let expected = matmul_cpu(&a_t, &grad_c, k, m, n);
1217
1218 let mut grad_b = vec![0.0f32; k * n];
1219 device
1220 .gemm_backward_b(&a, &grad_c, &mut grad_b, m as u32, k as u32, n as u32)
1221 .expect("gemm_backward_b");
1222
1223 let max_diff =
1224 grad_b.iter().zip(expected.iter()).map(|(a, b)| (a - b).abs()).fold(0.0f32, f32::max);
1225
1226 assert!(
1227 max_diff < 1e-3,
1228 "FALSIFY-WGPU-001: GEMM backward B max diff = {max_diff} (threshold: 1e-3)"
1229 );
1230 }
1231
1232 #[test]
1234 fn test_falsify_wgpu_001_rope_backward_parity() {
1235 let Some(device) = device_or_skip() else { return };
1236
1237 let (num_heads, head_dim, seq_len) = (2, 4, 3);
1238 let theta = 10000.0f32;
1239 let n = num_heads * head_dim * seq_len;
1240
1241 let grad_output: Vec<f32> = (0..n).map(|i| (i as f32 - 12.0) * 0.1).collect();
1242
1243 let half_dim = head_dim / 2;
1245 let mut expected = vec![0.0f32; n];
1246 for h in 0..num_heads {
1247 for s in 0..seq_len {
1248 for p in 0..half_dim {
1249 let freq_exp = -((2 * p) as f32) / head_dim as f32 * theta.log2();
1250 let inv_freq = 2.0f32.powf(freq_exp);
1251 let angle = s as f32 * inv_freq;
1252 let (sin_a, cos_a) = angle.sin_cos();
1253
1254 let base = h * seq_len * head_dim + s * head_dim;
1255 let even = base + 2 * p;
1256 let odd = base + 2 * p + 1;
1257
1258 let dy_even = grad_output[even];
1259 let dy_odd = grad_output[odd];
1260
1261 expected[even] = dy_even * cos_a + dy_odd * sin_a;
1263 expected[odd] = -dy_even * sin_a + dy_odd * cos_a;
1264 }
1265 }
1266 }
1267
1268 let mut grad_input = vec![0.0f32; n];
1269 device
1270 .rope_backward(
1271 &grad_output,
1272 &mut grad_input,
1273 num_heads as u32,
1274 head_dim as u32,
1275 seq_len as u32,
1276 theta,
1277 )
1278 .expect("rope_backward");
1279
1280 let max_diff = grad_input
1281 .iter()
1282 .zip(expected.iter())
1283 .map(|(a, b)| (a - b).abs())
1284 .fold(0.0f32, f32::max);
1285
1286 assert!(
1287 max_diff < 1e-4,
1288 "FALSIFY-WGPU-001: RoPE backward max diff = {max_diff} (threshold: 1e-4)"
1289 );
1290 }
1291
1292 #[test]
1294 fn test_falsify_wgpu_001_adamw_step_parity() {
1295 let Some(device) = device_or_skip() else { return };
1296
1297 let n = 16;
1298 let mut params: Vec<f32> = (0..n).map(|i| i as f32 * 0.1).collect();
1299 let grads: Vec<f32> = (0..n).map(|i| (i as f32 - 8.0) * 0.01).collect();
1300 let mut m_state = vec![0.0f32; n];
1301 let mut v_state = vec![0.0f32; n];
1302
1303 let lr: f32 = 1e-3;
1304 let beta1: f32 = 0.9;
1305 let beta2: f32 = 0.999;
1306 let eps: f32 = 1e-8;
1307 let wd: f32 = 0.01;
1308 let step = 1u32;
1309
1310 let bc1: f32 = 1.0 - beta1.powi(step as i32);
1312 let bc2: f32 = 1.0 - beta2.powi(step as i32);
1313 let mut cpu_params = params.clone();
1314 let mut cpu_m = m_state.clone();
1315 let mut cpu_v = v_state.clone();
1316 for i in 0..n {
1317 cpu_m[i] = beta1 * cpu_m[i] + (1.0 - beta1) * grads[i];
1318 cpu_v[i] = beta2 * cpu_v[i] + (1.0 - beta2) * grads[i] * grads[i];
1319 let m_hat = cpu_m[i] / bc1;
1320 let v_hat = cpu_v[i] / bc2;
1321 cpu_params[i] -= lr * (m_hat / (v_hat.sqrt() + eps) + wd * cpu_params[i]);
1322 }
1323
1324 device
1325 .adamw_step(
1326 &mut params,
1327 &grads,
1328 &mut m_state,
1329 &mut v_state,
1330 lr as f32,
1331 beta1 as f32,
1332 beta2 as f32,
1333 eps as f32,
1334 wd as f32,
1335 step,
1336 )
1337 .expect("adamw_step");
1338
1339 let max_diff =
1340 params.iter().zip(cpu_params.iter()).map(|(a, b)| (a - b).abs()).fold(0.0f32, f32::max);
1341
1342 assert!(
1343 max_diff < 1e-4,
1344 "FALSIFY-WGPU-001: AdamW step max diff = {max_diff} (threshold: 1e-4)"
1345 );
1346 }
1347
1348 #[test]
1350 fn test_falsify_wgpu_001_rmsnorm_backward_parity() {
1351 let Some(device) = device_or_skip() else { return };
1352
1353 let (num_rows, hidden_dim) = (3, 8);
1354 let eps: f32 = 1e-5;
1355 let n = num_rows * hidden_dim;
1356
1357 let input: Vec<f32> = (0..n).map(|i| (i as f32 - 12.0) * 0.1).collect();
1358 let gamma: Vec<f32> = (0..hidden_dim).map(|i| 1.0 + i as f32 * 0.1).collect();
1359 let grad_output: Vec<f32> = (0..n).map(|i| (i as f32 - 12.0) * 0.05).collect();
1360
1361 let mut cpu_grad_input = vec![0.0f32; n];
1363 let mut cpu_grad_gamma = vec![0.0f32; hidden_dim];
1364 for r in 0..num_rows {
1365 let row = &input[r * hidden_dim..(r + 1) * hidden_dim];
1366 let grow = &grad_output[r * hidden_dim..(r + 1) * hidden_dim];
1367
1368 let sum_x2: f32 = row.iter().map(|x| x * x).sum();
1369 let mean_x2 = sum_x2 / hidden_dim as f32;
1370 let var_eps = mean_x2 + eps;
1371 let rms = var_eps.sqrt();
1372 let inv_rms = 1.0 / rms;
1373
1374 let sum_xgg: f32 = row
1375 .iter()
1376 .zip(grow.iter())
1377 .zip(gamma.iter())
1378 .map(|((&x, &gy), &g)| x * gy * g)
1379 .sum();
1380 let mean_xgg = sum_xgg / hidden_dim as f32;
1381
1382 for i in 0..hidden_dim {
1383 let x = row[i];
1384 let gy = grow[i];
1385 let g = gamma[i];
1386 let gamma_gy = g * gy;
1387 let correction = (x / var_eps) * mean_xgg;
1388 cpu_grad_input[r * hidden_dim + i] = inv_rms * (gamma_gy - correction);
1389 cpu_grad_gamma[i] += gy * x * inv_rms;
1390 }
1391 }
1392
1393 let mut grad_input = vec![0.0f32; n];
1394 let mut grad_gamma = vec![0.0f32; hidden_dim];
1395
1396 device
1397 .rmsnorm_backward(
1398 &input,
1399 &gamma,
1400 &grad_output,
1401 &mut grad_input,
1402 &mut grad_gamma,
1403 num_rows as u32,
1404 hidden_dim as u32,
1405 eps,
1406 )
1407 .expect("rmsnorm_backward");
1408
1409 let gi_max_diff = grad_input
1410 .iter()
1411 .zip(cpu_grad_input.iter())
1412 .map(|(a, b)| (a - b).abs())
1413 .fold(0.0f32, f32::max);
1414
1415 let gg_max_diff = grad_gamma
1416 .iter()
1417 .zip(cpu_grad_gamma.iter())
1418 .map(|(a, b)| (a - b).abs())
1419 .fold(0.0f32, f32::max);
1420
1421 assert!(
1422 gi_max_diff < 1e-3,
1423 "FALSIFY-WGPU-001: RMSNorm grad_input max diff = {gi_max_diff}"
1424 );
1425 assert!(
1426 gg_max_diff < 1e-2,
1427 "FALSIFY-WGPU-001: RMSNorm grad_gamma max diff = {gg_max_diff} (atomic CAS accumulation)"
1428 );
1429 }
1430
1431 #[test]
1433 fn test_falsify_wgpu_003_nf4_dequant_parity() {
1434 let Some(device) = device_or_skip() else { return };
1435
1436 let nf4_lut: [f32; 16] = [
1438 -1.0,
1439 -0.6961928,
1440 -0.5250731,
1441 -0.39491749,
1442 -0.28444138,
1443 -0.18477343,
1444 -0.09105004,
1445 0.0,
1446 0.0795803,
1447 0.1609302,
1448 0.24611230,
1449 0.33791524,
1450 0.44070983,
1451 0.5626170,
1452 0.7229568,
1453 1.0,
1454 ];
1455
1456 let block_size = 4u32; let n = 8u32; let packed: Vec<u32> = vec![0x90F5_1C73_u32];
1471
1472 let scales: Vec<f32> = vec![2.0, 0.5]; let indices = [3, 7, 12, 1, 5, 15, 0, 9];
1474
1475 let mut expected = vec![0.0f32; n as usize];
1477 for i in 0..n as usize {
1478 let scale = scales[i / block_size as usize];
1479 expected[i] = nf4_lut[indices[i]] * scale;
1480 }
1481
1482 let mut output = vec![0.0f32; n as usize];
1483 device.nf4_dequant(&packed, &scales, &mut output, n, block_size).expect("nf4_dequant");
1484
1485 let max_diff =
1486 output.iter().zip(expected.iter()).map(|(a, b)| (a - b).abs()).fold(0.0f32, f32::max);
1487
1488 assert!(
1489 max_diff < 1e-6,
1490 "FALSIFY-WGPU-003: NF4 dequant max diff = {max_diff} (threshold: 1e-6)"
1491 );
1492 }
1493}