1#![allow(dead_code, clippy::many_single_char_names)]
2use std::collections::HashMap;
13
14pub struct LayerActivations {
19 pub attn_norm_out: wgpu::Buffer,
21 pub attn_output: wgpu::Buffer,
23 pub ffn_norm_out: wgpu::Buffer,
25 pub silu_gate_output: wgpu::Buffer,
27 pub rstd_attn: wgpu::Buffer,
29 pub rstd_ffn: wgpu::Buffer,
31 pub softmax_logsumexp: wgpu::Buffer,
33}
34
35pub struct QkvLoRA<'a> {
37 pub q_a: &'a wgpu::Buffer,
38 pub q_b: &'a wgpu::Buffer,
39 pub k_a: &'a wgpu::Buffer,
40 pub k_b: &'a wgpu::Buffer,
41 pub v_a: &'a wgpu::Buffer,
42 pub v_b: &'a wgpu::Buffer,
43 pub rank: u32,
44 pub scale: f32,
45 pub in_dim: u32,
46 pub q_dim: u32,
47 pub kv_dim: u32,
48 pub lora_pipeline: &'a wgpu::ComputePipeline,
49 pub lora_bgl: &'a wgpu::BindGroupLayout,
50}
51
52pub struct WgslForwardPass {
55 device: wgpu::Device,
56 queue: wgpu::Queue,
57
58 matmul_pipeline: wgpu::ComputePipeline,
60 tiled_matmul_pipeline: wgpu::ComputePipeline,
62 gemv_pipeline: wgpu::ComputePipeline,
64 q4k_gemv_pipeline: wgpu::ComputePipeline,
66 attention_pipeline: wgpu::ComputePipeline,
68 attention_bgl: wgpu::BindGroupLayout,
69 rmsnorm_pipeline: wgpu::ComputePipeline,
70 silu_mul_pipeline: wgpu::ComputePipeline,
71 rope_pipeline: wgpu::ComputePipeline,
72 batch_rope_pipeline: wgpu::ComputePipeline,
73 batch_rope_bgl: wgpu::BindGroupLayout,
74 residual_pipeline: wgpu::ComputePipeline,
75
76 matmul_bgl: wgpu::BindGroupLayout,
78 elementwise_bgl: wgpu::BindGroupLayout,
79
80 weight_buffers: HashMap<String, wgpu::Buffer>,
82 q4k_weights: HashMap<String, wgpu::Buffer>,
84 cpu_biases: HashMap<String, Vec<f32>>,
86 kv_cache_k: Vec<wgpu::Buffer>,
88 kv_cache_v: Vec<wgpu::Buffer>,
90
91 hidden_buf: wgpu::Buffer, q_buf: wgpu::Buffer, k_buf: wgpu::Buffer, v_buf: wgpu::Buffer, attn_out_buf: wgpu::Buffer, ffn_gate_buf: wgpu::Buffer, ffn_up_buf: wgpu::Buffer, ffn_silu_buf: wgpu::Buffer, ffn_out_buf: wgpu::Buffer, norm_buf: wgpu::Buffer, staging_buf: wgpu::Buffer, hidden_dim: u32,
107 num_heads: u32,
108 num_kv_heads: u32,
109 head_dim: u32,
110 intermediate_dim: u32,
111 rms_norm_eps: f32,
115}
116
117pub const DEFAULT_RMS_NORM_EPS: f32 = 1e-6;
119
120fn rmsnorm_params(dim: u32, eps: f32) -> [u32; 4] {
122 [dim, eps.to_bits(), 0, 0]
123}
124
125const RMSNORM_SHADER: &str = r#"
128@group(0) @binding(0) var<storage, read> input: array<f32>;
129@group(0) @binding(1) var<storage, read> weight: array<f32>;
130@group(0) @binding(2) var<storage, read_write> output: array<f32>;
131@group(0) @binding(3) var<uniform> params: vec4<u32>; // (dim, bitcast<u32>(eps), 0, 0)
132
133var<workgroup> shared_sum: array<f32, 256>;
134
135@compute @workgroup_size(256)
136fn main(@builtin(local_invocation_id) lid: vec3<u32>,
137 @builtin(workgroup_id) wg_id: vec3<u32>) {
138 let dim = params.x;
139 let row = wg_id.y;
140 let base = row * dim;
141 let tid = lid.x;
142
143 // Compute sum of squares (reduction) for this row
144 var local_sum: f32 = 0.0;
145 var i = tid;
146 while (i < dim) {
147 let val = input[base + i];
148 local_sum += val * val;
149 i += 256u;
150 }
151 shared_sum[tid] = local_sum;
152 workgroupBarrier();
153
154 // Tree reduction
155 var stride = 128u;
156 while (stride > 0u) {
157 if (tid < stride) {
158 shared_sum[tid] += shared_sum[tid + stride];
159 }
160 workgroupBarrier();
161 stride >>= 1u;
162 }
163
164 let eps = bitcast<f32>(params.y);
165 let rms = sqrt(shared_sum[0] / f32(dim) + eps);
166
167 // Normalize and scale
168 i = tid;
169 while (i < dim) {
170 output[base + i] = (input[base + i] / rms) * weight[i];
171 i += 256u;
172 }
173}
174"#;
175
176const SILU_MUL_SHADER: &str = r#"
178@group(0) @binding(0) var<storage, read> gate: array<f32>;
179@group(0) @binding(1) var<storage, read> up: array<f32>;
180@group(0) @binding(2) var<storage, read_write> output: array<f32>;
181@group(0) @binding(3) var<uniform> params: vec4<u32>; // (dim, 0, 0, 0)
182
183@compute @workgroup_size(256)
184fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
185 let idx = gid.x;
186 if (idx >= params.x) { return; }
187 let g = gate[idx];
188 let silu_g = g / (1.0 + exp(-g));
189 output[idx] = silu_g * up[idx];
190}
191"#;
192
193const RESIDUAL_SHADER: &str = r#"
195@group(0) @binding(0) var<storage, read> a: array<f32>;
196@group(0) @binding(1) var<storage, read> b: array<f32>;
197@group(0) @binding(2) var<storage, read_write> output: array<f32>;
198@group(0) @binding(3) var<uniform> params: vec4<u32>; // (dim, 0, 0, 0)
199
200@compute @workgroup_size(256)
201fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
202 let idx = gid.x;
203 if (idx >= params.x) { return; }
204 output[idx] = a[idx] + b[idx];
205}
206"#;
207
208const BATCH_ROPE_SHADER: &str = r#"
212@group(0) @binding(0) var<storage, read_write> qk: array<f32>;
213
214struct RopeParams {
215 seq_len: u32,
216 num_heads: u32,
217 head_dim: u32,
218 _pad: u32,
219}
220
221@group(0) @binding(1) var<uniform> params: RopeParams;
222
223@compute @workgroup_size(256)
224fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
225 let idx = gid.x;
226 let total = params.seq_len * params.num_heads * params.head_dim;
227 if (idx >= total) { return; }
228
229 let head_dim = params.head_dim;
230 let half_hd = head_dim / 2u;
231
232 // Decompose idx into (position, head, pos_in_head)
233 let elements_per_pos = params.num_heads * head_dim;
234 let position = idx / elements_per_pos;
235 let within_pos = idx % elements_per_pos;
236 let head_idx = within_pos / head_dim;
237 let pos_in_head = within_pos % head_dim;
238
239 // Only process the first half of each head (pairs with second half)
240 if (pos_in_head >= half_hd) { return; }
241
242 let theta = pow(1000000.0, -f32(pos_in_head * 2u) / f32(head_dim));
243 let angle = f32(position) * theta;
244 let cos_a = cos(angle);
245 let sin_a = sin(angle);
246
247 let base = position * elements_per_pos + head_idx * head_dim;
248 let i0 = base + pos_in_head;
249 let i1 = i0 + half_hd;
250
251 let x0 = qk[i0];
252 let x1 = qk[i1];
253 qk[i0] = x0 * cos_a - x1 * sin_a;
254 qk[i1] = x0 * sin_a + x1 * cos_a;
255}
256"#;
257
258const ROPE_SHADER: &str = r#"
260@group(0) @binding(0) var<storage, read_write> qk: array<f32>;
261@group(0) @binding(1) var<uniform> params: vec4<u32>; // (dim, position, num_heads, head_dim)
262
263@compute @workgroup_size(256)
264fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
265 let idx = gid.x;
266 let dim = params.x;
267 let position = params.y;
268 let head_dim = params.w;
269
270 if (idx >= dim) { return; }
271
272 let half_hd = head_dim / 2u;
273 let head_idx = idx / head_dim;
274 let pos_in_head = idx % head_dim;
275
276 if (pos_in_head >= half_hd) { return; }
277
278 let theta = pow(1000000.0, -f32(pos_in_head * 2u) / f32(head_dim));
279 let angle = f32(position) * theta;
280 let cos_a = cos(angle);
281 let sin_a = sin(angle);
282
283 let i0 = head_idx * head_dim + pos_in_head;
284 let i1 = i0 + half_hd;
285
286 let x0 = qk[i0];
287 let x1 = qk[i1];
288 qk[i0] = x0 * cos_a - x1 * sin_a;
289 qk[i1] = x0 * sin_a + x1 * cos_a;
290}
291"#;
292
293impl WgslForwardPass {
294 pub fn rmsnorm_shader() -> &'static str {
296 RMSNORM_SHADER
297 }
298 pub fn silu_mul_shader() -> &'static str {
299 SILU_MUL_SHADER
300 }
301 pub fn residual_shader() -> &'static str {
302 RESIDUAL_SHADER
303 }
304 pub fn rope_shader() -> &'static str {
305 ROPE_SHADER
306 }
307
308 pub fn new(
313 device: wgpu::Device,
314 queue: wgpu::Queue,
315 hidden_dim: usize,
316 num_heads: usize,
317 num_kv_heads: usize,
318 head_dim: usize,
319 intermediate_dim: usize,
320 ) -> Self {
321 let q_dim = num_heads * head_dim;
322 let kv_dim = num_kv_heads * head_dim;
323
324 let matmul_shader = device.create_shader_module(wgpu::ShaderModuleDescriptor {
326 label: Some("matmul"),
327 source: wgpu::ShaderSource::Wgsl(crate::backends::gpu::shaders::MATMUL_SHADER.into()),
328 });
329 let rmsnorm_shader = device.create_shader_module(wgpu::ShaderModuleDescriptor {
330 label: Some("rmsnorm"),
331 source: wgpu::ShaderSource::Wgsl(RMSNORM_SHADER.into()),
332 });
333 let silu_mul_shader = device.create_shader_module(wgpu::ShaderModuleDescriptor {
334 label: Some("silu_mul"),
335 source: wgpu::ShaderSource::Wgsl(SILU_MUL_SHADER.into()),
336 });
337 let rope_shader_mod = device.create_shader_module(wgpu::ShaderModuleDescriptor {
338 label: Some("rope"),
339 source: wgpu::ShaderSource::Wgsl(ROPE_SHADER.into()),
340 });
341 let residual_shader_mod = device.create_shader_module(wgpu::ShaderModuleDescriptor {
342 label: Some("residual"),
343 source: wgpu::ShaderSource::Wgsl(RESIDUAL_SHADER.into()),
344 });
345
346 let matmul_bgl = device.create_bind_group_layout(&wgpu::BindGroupLayoutDescriptor {
348 label: Some("matmul_bgl"),
349 entries: &[
350 bgl_storage(0, true),
351 bgl_storage(1, true),
352 bgl_storage(2, false),
353 bgl_uniform(3),
354 ],
355 });
356 let elementwise_bgl = device.create_bind_group_layout(&wgpu::BindGroupLayoutDescriptor {
357 label: Some("ew_bgl"),
358 entries: &[
359 bgl_storage(0, true),
360 bgl_storage(1, true),
361 bgl_storage(2, false),
362 bgl_uniform(3),
363 ],
364 });
365
366 let make_pipeline =
368 |shader: &wgpu::ShaderModule, bgl: &wgpu::BindGroupLayout, label: &str| {
369 let pl = device.create_pipeline_layout(&wgpu::PipelineLayoutDescriptor {
370 label: Some(label),
371 bind_group_layouts: &[bgl],
372 push_constant_ranges: &[],
373 });
374 device.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
375 label: Some(label),
376 layout: Some(&pl),
377 module: shader,
378 entry_point: Some("main"),
379 compilation_options: Default::default(),
380 cache: None,
381 })
382 };
383
384 let matmul_pipeline = make_pipeline(&matmul_shader, &matmul_bgl, "matmul_pipe");
385
386 let tiled_matmul_shader = device.create_shader_module(wgpu::ShaderModuleDescriptor {
388 label: Some("tiled_matmul"),
389 source: wgpu::ShaderSource::Wgsl(
390 crate::backends::gpu::shaders::TILED_GEMM_SHADER.into(),
391 ),
392 });
393 let tiled_matmul_pipeline =
394 make_pipeline(&tiled_matmul_shader, &matmul_bgl, "tiled_matmul_pipe");
395
396 let attention_shader = device.create_shader_module(wgpu::ShaderModuleDescriptor {
398 label: Some("causal_attention"),
399 source: wgpu::ShaderSource::Wgsl(
400 crate::backends::gpu::shaders::CAUSAL_ATTENTION_SHADER.into(),
401 ),
402 });
403 let attention_bgl = device.create_bind_group_layout(&wgpu::BindGroupLayoutDescriptor {
404 label: Some("attn_bgl"),
405 entries: &[
406 bgl_storage(0, true), bgl_storage(1, true), bgl_storage(2, true), bgl_storage(3, false), bgl_uniform(4), ],
412 });
413 let attention_pl = device.create_pipeline_layout(&wgpu::PipelineLayoutDescriptor {
414 label: Some("attn_pl"),
415 bind_group_layouts: &[&attention_bgl],
416 push_constant_ranges: &[],
417 });
418 let attention_pipeline = device.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
419 label: Some("attn_pipe"),
420 layout: Some(&attention_pl),
421 module: &attention_shader,
422 entry_point: Some("main"),
423 compilation_options: Default::default(),
424 cache: None,
425 });
426
427 let gemv_shader = device.create_shader_module(wgpu::ShaderModuleDescriptor {
429 label: Some("gemv"),
430 source: wgpu::ShaderSource::Wgsl(crate::backends::gpu::shaders::GEMV_SHADER.into()),
431 });
432 let gemv_pipeline = make_pipeline(&gemv_shader, &matmul_bgl, "gemv_pipe");
433
434 let q4k_gemv_shader = device.create_shader_module(wgpu::ShaderModuleDescriptor {
436 label: Some("q4k_gemv"),
437 source: wgpu::ShaderSource::Wgsl(crate::backends::gpu::shaders::Q4K_GEMV_SHADER.into()),
438 });
439 let q4k_gemv_pipeline = make_pipeline(&q4k_gemv_shader, &matmul_bgl, "q4k_gemv_pipe");
440
441 let rmsnorm_pipeline = make_pipeline(&rmsnorm_shader, &elementwise_bgl, "rmsnorm_pipe");
442 let silu_mul_pipeline = make_pipeline(&silu_mul_shader, &elementwise_bgl, "silu_pipe");
443 let residual_pipeline = make_pipeline(&residual_shader_mod, &elementwise_bgl, "res_pipe");
444
445 let rope_bgl = device.create_bind_group_layout(&wgpu::BindGroupLayoutDescriptor {
447 label: Some("rope_bgl"),
448 entries: &[bgl_storage(0, false), bgl_uniform(1)],
449 });
450 let rope_pipeline = {
451 let pl = device.create_pipeline_layout(&wgpu::PipelineLayoutDescriptor {
452 label: Some("rope_pl"),
453 bind_group_layouts: &[&rope_bgl],
454 push_constant_ranges: &[],
455 });
456 device.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
457 label: Some("rope_pipe"),
458 layout: Some(&pl),
459 module: &rope_shader_mod,
460 entry_point: Some("main"),
461 compilation_options: Default::default(),
462 cache: None,
463 })
464 };
465
466 let batch_rope_shader_mod = device.create_shader_module(wgpu::ShaderModuleDescriptor {
468 label: Some("batch_rope"),
469 source: wgpu::ShaderSource::Wgsl(BATCH_ROPE_SHADER.into()),
470 });
471 let batch_rope_bgl = device.create_bind_group_layout(&wgpu::BindGroupLayoutDescriptor {
472 label: Some("batch_rope_bgl"),
473 entries: &[bgl_storage(0, false), bgl_uniform(1)],
474 });
475 let batch_rope_pipeline = {
476 let pl = device.create_pipeline_layout(&wgpu::PipelineLayoutDescriptor {
477 label: Some("batch_rope_pl"),
478 bind_group_layouts: &[&batch_rope_bgl],
479 push_constant_ranges: &[],
480 });
481 device.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
482 label: Some("batch_rope_pipe"),
483 layout: Some(&pl),
484 module: &batch_rope_shader_mod,
485 entry_point: Some("main"),
486 compilation_options: Default::default(),
487 cache: None,
488 })
489 };
490
491 let buf = |size: usize, label: &str| -> wgpu::Buffer {
493 device.create_buffer(&wgpu::BufferDescriptor {
494 label: Some(label),
495 size: (size * 4) as u64,
496 usage: wgpu::BufferUsages::STORAGE
497 | wgpu::BufferUsages::COPY_SRC
498 | wgpu::BufferUsages::COPY_DST,
499 mapped_at_creation: false,
500 })
501 };
502
503 let max_seq = 2048;
507 let hidden_buf = buf(max_seq * hidden_dim, "hidden");
508 let q_buf = buf(max_seq * q_dim, "q");
509 let k_buf = buf(max_seq * kv_dim, "k");
510 let v_buf = buf(max_seq * kv_dim, "v");
511 let attn_out_buf = buf(max_seq * hidden_dim, "attn_out");
512 let ffn_gate_buf = buf(max_seq * intermediate_dim, "ffn_gate");
513 let ffn_up_buf = buf(max_seq * intermediate_dim, "ffn_up");
514 let ffn_silu_buf = buf(max_seq * intermediate_dim, "ffn_silu");
515 let ffn_out_buf = buf(max_seq * hidden_dim, "ffn_out");
516 let norm_buf = buf(max_seq * hidden_dim, "norm");
517
518 let max_out = max_seq * hidden_dim.max(intermediate_dim);
519 let staging_buf = device.create_buffer(&wgpu::BufferDescriptor {
520 label: Some("staging"),
521 size: (max_out * 4) as u64,
522 usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
523 mapped_at_creation: false,
524 });
525
526 Self {
527 device,
528 queue,
529 matmul_pipeline,
530 tiled_matmul_pipeline,
531 attention_pipeline,
532 attention_bgl,
533 gemv_pipeline,
534 q4k_gemv_pipeline,
535 rmsnorm_pipeline,
536 silu_mul_pipeline,
537 rope_pipeline,
538 batch_rope_pipeline,
539 batch_rope_bgl,
540 residual_pipeline,
541 matmul_bgl,
542 elementwise_bgl,
543 weight_buffers: HashMap::new(),
544 q4k_weights: HashMap::new(),
545 kv_cache_k: Vec::new(),
546 kv_cache_v: Vec::new(),
547 cpu_biases: HashMap::new(),
548 hidden_buf,
549 q_buf,
550 k_buf,
551 v_buf,
552 attn_out_buf,
553 ffn_gate_buf,
554 ffn_up_buf,
555 ffn_silu_buf,
556 ffn_out_buf,
557 norm_buf,
558 staging_buf,
559 hidden_dim: hidden_dim as u32,
560 num_heads: num_heads as u32,
561 num_kv_heads: num_kv_heads as u32,
562 head_dim: head_dim as u32,
563 intermediate_dim: intermediate_dim as u32,
564 rms_norm_eps: DEFAULT_RMS_NORM_EPS,
565 }
566 }
567
568 pub fn set_rms_norm_eps(&mut self, eps: f32) {
571 self.rms_norm_eps = eps;
572 }
573
574 #[must_use]
576 pub fn rms_norm_eps(&self) -> f32 {
577 self.rms_norm_eps
578 }
579
580 pub fn upload_weight(&mut self, name: &str, data: &[f32]) {
583 if name.contains("bias") {
584 self.cpu_biases.insert(name.to_string(), data.to_vec());
586 return;
587 }
588 let size_bytes = (data.len() * 4) as u64;
590 let max_binding = self.device.limits().max_storage_buffer_binding_size as u64;
591 if size_bytes > max_binding {
592 eprintln!(
593 "[wgpu] Skipping weight '{}' ({:.1} MB > {:.1} MB limit) — CPU fallback",
594 name,
595 size_bytes as f64 / 1e6,
596 max_binding as f64 / 1e6
597 );
598 return;
599 }
600 use wgpu::util::DeviceExt;
601 let buffer = self.device.create_buffer_init(&wgpu::util::BufferInitDescriptor {
602 label: Some(name),
603 contents: bytemuck::cast_slice(data),
604 usage: wgpu::BufferUsages::STORAGE,
605 });
606 self.weight_buffers.insert(name.to_string(), buffer);
607 }
608
609 pub fn upload_q4k_weight(&mut self, name: &str, data: &[u8]) {
611 use wgpu::util::DeviceExt;
612 let buffer = self.device.create_buffer_init(&wgpu::util::BufferInitDescriptor {
613 label: Some(name),
614 contents: data,
615 usage: wgpu::BufferUsages::STORAGE,
616 });
617 self.q4k_weights.insert(name.to_string(), buffer);
618 }
619
620 pub fn init_kv_cache(&mut self, num_layers: usize) {
622 let kv_dim = (self.num_kv_heads * self.head_dim) as u64;
623 let max_seq = 2048u64;
624 for _ in 0..num_layers {
625 let k = self.device.create_buffer(&wgpu::BufferDescriptor {
626 label: Some("kv_cache_k"),
627 size: max_seq * kv_dim * 4,
628 usage: wgpu::BufferUsages::STORAGE
629 | wgpu::BufferUsages::COPY_DST
630 | wgpu::BufferUsages::COPY_SRC,
631 mapped_at_creation: false,
632 });
633 let v = self.device.create_buffer(&wgpu::BufferDescriptor {
634 label: Some("kv_cache_v"),
635 size: max_seq * kv_dim * 4,
636 usage: wgpu::BufferUsages::STORAGE
637 | wgpu::BufferUsages::COPY_DST
638 | wgpu::BufferUsages::COPY_SRC,
639 mapped_at_creation: false,
640 });
641 self.kv_cache_k.push(k);
642 self.kv_cache_v.push(v);
643 }
644 }
645
646 pub fn weight_count(&self) -> usize {
648 self.weight_buffers.len()
649 }
650
651 pub fn weight_buffer(&self, name: &str) -> Option<&wgpu::Buffer> {
654 self.weight_buffers.get(name)
655 }
656
657 pub fn device_ref(&self) -> &wgpu::Device {
659 &self.device
660 }
661
662 pub fn queue_ref(&self) -> &wgpu::Queue {
664 &self.queue
665 }
666
667 pub fn hidden_buffer(&self) -> &wgpu::Buffer {
669 &self.hidden_buf
670 }
671
672 pub fn q_buffer(&self) -> &wgpu::Buffer {
674 &self.q_buf
675 }
676
677 pub fn k_buffer(&self) -> &wgpu::Buffer {
679 &self.k_buf
680 }
681
682 pub fn v_buffer(&self) -> &wgpu::Buffer {
684 &self.v_buf
685 }
686
687 pub fn gpu_residual_add(
689 &self,
690 a: &wgpu::Buffer,
691 b: &wgpu::Buffer,
692 output: &wgpu::Buffer,
693 len: u32,
694 ) {
695 let mut encoder = self.device.create_command_encoder(&Default::default());
696 self.encode_residual(&mut encoder, a, b, output, len);
697 self.queue.submit(Some(encoder.finish()));
698 }
699
700 pub fn gpu_rmsnorm(&self, weight: &wgpu::Buffer, output: &wgpu::Buffer, _seq_len: u32) {
703 let mut encoder = self.device.create_command_encoder(&Default::default());
704 self.encode_rmsnorm(&mut encoder, &self.hidden_buf, weight, output, self.hidden_dim);
705 self.queue.submit(Some(encoder.finish()));
706 }
707
708 pub fn download_hidden(&self, len: usize) -> Vec<f32> {
710 let size = (len * 4) as u64;
711 let staging = self.device.create_buffer(&wgpu::BufferDescriptor {
712 label: Some("hidden_download"),
713 size,
714 usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
715 mapped_at_creation: false,
716 });
717 let mut encoder = self.device.create_command_encoder(&Default::default());
718 encoder.copy_buffer_to_buffer(&self.hidden_buf, 0, &staging, 0, size);
719 self.queue.submit(Some(encoder.finish()));
720
721 let slice = staging.slice(..size);
722 let (tx, rx) = std::sync::mpsc::channel();
723 slice.map_async(wgpu::MapMode::Read, move |r| {
724 tx.send(r).ok();
725 });
726 self.device.poll(wgpu::PollType::Wait { submission_index: None, timeout: None }).ok();
727 rx.recv()
728 .expect("GPU map_async callback channel disconnected")
729 .expect("GPU buffer mapping failed");
730
731 let data = slice.get_mapped_range();
732 let result: Vec<f32> = bytemuck::cast_slice(&data)[..len].to_vec();
733 drop(data);
734 staging.unmap();
735 result
736 }
737
738 pub fn total_vram_bytes(&self) -> usize {
740 let weight_bytes: usize = self.weight_buffers.values().map(|b| b.size() as usize).sum();
741 let intermediate_bytes = (self.hidden_dim as usize * 4) * 4 + (self.num_heads as usize * self.head_dim as usize * 4) + (self.num_kv_heads as usize * self.head_dim as usize * 4) * 2 + (self.intermediate_dim as usize * 4) * 2; weight_bytes + intermediate_bytes
746 }
747
748 #[provable_contracts_macros::contract("wgpu-forward-pass-v1", equation = "rmsnorm_correctness")]
754 pub fn forward_model(
755 &self,
756 token_id: u32,
757 position: usize,
758 num_layers: usize,
759 token_embedding: &[f32],
760 output_norm_weight: &[f32],
761 lm_head_weight: &[f32],
762 vocab_size: usize,
763 eps: f32,
764 kv_caches: &mut Vec<(Vec<f32>, Vec<f32>)>,
765 ) -> Result<Vec<f32>, String> {
766 let hd = self.hidden_dim as usize;
767
768 let embed_start = token_id as usize * hd;
770 if embed_start + hd > token_embedding.len() {
771 return Err(format!(
772 "Token {} out of range (embedding size {})",
773 token_id,
774 token_embedding.len() / hd
775 ));
776 }
777 let mut hidden: Vec<f32> = token_embedding[embed_start..embed_start + hd].to_vec();
778
779 while kv_caches.len() < num_layers {
782 kv_caches.push((Vec::new(), Vec::new()));
783 }
784 for layer_idx in 0..num_layers {
785 let prefix = format!("layer.{layer_idx}");
786 let (ref mut k_cache, ref mut v_cache) = kv_caches[layer_idx];
787 self.forward_layer(&mut hidden, &prefix, position, k_cache, v_cache)?;
788 }
789
790 let rms = (hidden.iter().map(|x| x * x).sum::<f32>() / hd as f32 + eps).sqrt();
792 for i in 0..hd {
793 hidden[i] = (hidden[i] / rms) * output_norm_weight[i];
794 }
795
796 let mut logits = vec![0.0f32; vocab_size];
802 for v in 0..vocab_size {
803 let mut sum = 0.0f32;
804 let row_start = v * hd;
805 for j in 0..hd {
806 sum += lm_head_weight[row_start + j] * hidden[j];
807 }
808 logits[v] = sum;
809 }
810 Ok(logits)
811 }
812
813 pub fn forward_layer(
820 &self,
821 hidden: &mut [f32],
822 layer_prefix: &str,
823 _position: usize,
824 kv_cache_k: &mut Vec<f32>, kv_cache_v: &mut Vec<f32>, ) -> Result<(), String> {
827 let hd = self.hidden_dim;
828
829 self.queue.write_buffer(&self.hidden_buf, 0, bytemuck::cast_slice(hidden));
831
832 let mut encoder = self.device.create_command_encoder(&Default::default());
833
834 let norm_w = self
836 .weight_buffers
837 .get(&format!("{layer_prefix}.attn_norm"))
838 .ok_or_else(|| format!("Missing {layer_prefix}.attn_norm"))?;
839 self.encode_rmsnorm(&mut encoder, &self.hidden_buf, norm_w, &self.norm_buf, hd);
840
841 let q_dim = self.num_heads * self.head_dim;
843 let kv_dim = self.num_kv_heads * self.head_dim;
844
845 self.encode_matmul(
846 &mut encoder,
847 &self.norm_buf,
848 layer_prefix,
849 "q_proj",
850 &self.q_buf,
851 1,
852 hd,
853 q_dim,
854 );
855 self.encode_matmul(
856 &mut encoder,
857 &self.norm_buf,
858 layer_prefix,
859 "k_proj",
860 &self.k_buf,
861 1,
862 hd,
863 kv_dim,
864 );
865 self.encode_matmul(
866 &mut encoder,
867 &self.norm_buf,
868 layer_prefix,
869 "v_proj",
870 &self.v_buf,
871 1,
872 hd,
873 kv_dim,
874 );
875
876 let q_bytes = (q_dim * 4) as u64;
879 let kv_bytes = (kv_dim * 4) as u64;
880
881 let q_staging = self.device.create_buffer(&wgpu::BufferDescriptor {
883 label: Some("q_stg"),
884 size: q_bytes,
885 usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
886 mapped_at_creation: false,
887 });
888 let k_staging = self.device.create_buffer(&wgpu::BufferDescriptor {
889 label: Some("k_stg"),
890 size: kv_bytes,
891 usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
892 mapped_at_creation: false,
893 });
894 let v_staging = self.device.create_buffer(&wgpu::BufferDescriptor {
895 label: Some("v_stg"),
896 size: kv_bytes,
897 usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
898 mapped_at_creation: false,
899 });
900 encoder.copy_buffer_to_buffer(&self.q_buf, 0, &q_staging, 0, q_bytes);
901 encoder.copy_buffer_to_buffer(&self.k_buf, 0, &k_staging, 0, kv_bytes);
902 encoder.copy_buffer_to_buffer(&self.v_buf, 0, &v_staging, 0, kv_bytes);
903 self.queue.submit(Some(encoder.finish()));
904
905 let mut q_data = vec![0.0f32; q_dim as usize];
907 {
908 let slice = q_staging.slice(..q_bytes);
909 let (tx, rx) = std::sync::mpsc::channel();
910 slice.map_async(wgpu::MapMode::Read, move |r| {
911 tx.send(r).ok();
912 });
913 self.device.poll(wgpu::PollType::Wait { submission_index: None, timeout: None }).ok();
914 rx.recv().map_err(|e| format!("q recv: {e}"))?.map_err(|e| format!("q map: {e:?}"))?;
915 let data = slice.get_mapped_range();
916 q_data.copy_from_slice(&bytemuck::cast_slice::<u8, f32>(&data)[..q_dim as usize]);
917 }
918 q_staging.unmap();
919
920 let mut k_data = vec![0.0f32; kv_dim as usize];
922 {
923 let slice = k_staging.slice(..kv_bytes);
924 let (tx, rx) = std::sync::mpsc::channel();
925 slice.map_async(wgpu::MapMode::Read, move |r| {
926 tx.send(r).ok();
927 });
928 self.device.poll(wgpu::PollType::Wait { submission_index: None, timeout: None }).ok();
929 rx.recv().map_err(|e| format!("k recv: {e}"))?.map_err(|e| format!("k map: {e:?}"))?;
930 let data = slice.get_mapped_range();
931 k_data.copy_from_slice(&bytemuck::cast_slice::<u8, f32>(&data)[..kv_dim as usize]);
932 }
933 k_staging.unmap();
934
935 let mut v_data = vec![0.0f32; kv_dim as usize];
937 {
938 let slice = v_staging.slice(..kv_bytes);
939 let (tx, rx) = std::sync::mpsc::channel();
940 slice.map_async(wgpu::MapMode::Read, move |r| {
941 tx.send(r).ok();
942 });
943 self.device.poll(wgpu::PollType::Wait { submission_index: None, timeout: None }).ok();
944 rx.recv().map_err(|e| format!("v recv: {e}"))?.map_err(|e| format!("v map: {e:?}"))?;
945 let data = slice.get_mapped_range();
946 v_data.copy_from_slice(&bytemuck::cast_slice::<u8, f32>(&data)[..kv_dim as usize]);
947 }
948 v_staging.unmap();
949
950 if let Some(q_bias) = self.cpu_biases.get(&format!("{layer_prefix}.q_bias")) {
952 for (q, b) in q_data.iter_mut().zip(q_bias.iter()) {
953 *q += *b;
954 }
955 }
956 if let Some(k_bias) = self.cpu_biases.get(&format!("{layer_prefix}.k_bias")) {
957 for (k, b) in k_data.iter_mut().zip(k_bias.iter()) {
958 *k += *b;
959 }
960 }
961 if let Some(v_bias) = self.cpu_biases.get(&format!("{layer_prefix}.v_bias")) {
962 for (v, b) in v_data.iter_mut().zip(v_bias.iter()) {
963 *v += *b;
964 }
965 }
966
967 let head_dim = self.head_dim as usize;
969 let position = _position; let rope_theta = 1_000_000.0f64; for h in 0..(self.num_heads as usize) {
974 let offset = h * head_dim;
975 let half = head_dim / 2;
976 for i in 0..half {
977 let theta = rope_theta.powf(-((2 * i) as f64) / head_dim as f64);
978 let angle = position as f64 * theta;
979 let cos_a = angle.cos() as f32;
980 let sin_a = angle.sin() as f32;
981 let x0 = q_data[offset + i];
982 let x1 = q_data[offset + i + half];
983 q_data[offset + i] = x0 * cos_a - x1 * sin_a;
984 q_data[offset + i + half] = x0 * sin_a + x1 * cos_a;
985 }
986 }
987
988 for h in 0..(self.num_kv_heads as usize) {
990 let offset = h * head_dim;
991 let half = head_dim / 2;
992 for i in 0..half {
993 let theta = rope_theta.powf(-((2 * i) as f64) / head_dim as f64);
994 let angle = position as f64 * theta;
995 let cos_a = angle.cos() as f32;
996 let sin_a = angle.sin() as f32;
997 let x0 = k_data[offset + i];
998 let x1 = k_data[offset + i + half];
999 k_data[offset + i] = x0 * cos_a - x1 * sin_a;
1000 k_data[offset + i + half] = x0 * sin_a + x1 * cos_a;
1001 }
1002 }
1003
1004 let head_dim = self.head_dim as usize;
1006 let num_heads = self.num_heads as usize;
1007 let num_kv_heads = self.num_kv_heads as usize;
1008 let kv_dim_usize = kv_dim as usize;
1009
1010 kv_cache_k.extend_from_slice(&k_data);
1011 kv_cache_v.extend_from_slice(&v_data);
1012 let seq_len = kv_cache_k.len() / kv_dim_usize;
1013
1014 let kv_group = num_heads / num_kv_heads;
1016 let scale = 1.0 / (head_dim as f32).sqrt();
1017 let mut attn_out = vec![0.0f32; q_dim as usize];
1018
1019 for h in 0..num_heads {
1020 let kv_h = h / kv_group;
1021 let q_offset = h * head_dim;
1022 let kv_offset = kv_h * head_dim;
1023
1024 let scores = Self::attention_scores(
1026 &q_data[q_offset..q_offset + head_dim],
1027 kv_cache_k,
1028 kv_offset,
1029 kv_dim_usize,
1030 seq_len,
1031 scale,
1032 );
1033
1034 let out_offset = h * head_dim;
1036 Self::attention_weighted_v(
1037 &scores,
1038 kv_cache_v,
1039 &mut attn_out[out_offset..out_offset + head_dim],
1040 kv_offset,
1041 kv_dim_usize,
1042 );
1043 }
1044
1045 self.queue.write_buffer(&self.q_buf, 0, bytemuck::cast_slice(&attn_out));
1047
1048 let mut encoder = self.device.create_command_encoder(&Default::default());
1050
1051 self.encode_matmul(
1053 &mut encoder,
1054 &self.q_buf,
1055 layer_prefix,
1056 "o_proj",
1057 &self.attn_out_buf,
1058 1,
1059 q_dim,
1060 hd,
1061 );
1062
1063 self.encode_residual(
1065 &mut encoder,
1066 &self.hidden_buf,
1067 &self.attn_out_buf,
1068 &self.ffn_out_buf,
1069 hd,
1070 );
1071
1072 let ffn_norm_w = self
1074 .weight_buffers
1075 .get(&format!("{layer_prefix}.ffn_norm"))
1076 .ok_or_else(|| format!("Missing {layer_prefix}.ffn_norm"))?;
1077 self.encode_rmsnorm(&mut encoder, &self.ffn_out_buf, ffn_norm_w, &self.norm_buf, hd);
1078
1079 let inter = self.intermediate_dim;
1081 self.encode_matmul(
1082 &mut encoder,
1083 &self.norm_buf,
1084 layer_prefix,
1085 "gate_proj",
1086 &self.ffn_gate_buf,
1087 1,
1088 hd,
1089 inter,
1090 );
1091 self.encode_matmul(
1092 &mut encoder,
1093 &self.norm_buf,
1094 layer_prefix,
1095 "up_proj",
1096 &self.ffn_up_buf,
1097 1,
1098 hd,
1099 inter,
1100 );
1101
1102 self.encode_silu_mul(
1108 &mut encoder,
1109 &self.ffn_gate_buf,
1110 &self.ffn_up_buf,
1111 &self.ffn_silu_buf,
1112 inter,
1113 );
1114
1115 self.encode_matmul(
1117 &mut encoder,
1118 &self.ffn_silu_buf,
1119 layer_prefix,
1120 "down_proj",
1121 &self.norm_buf,
1122 1,
1123 inter,
1124 hd,
1125 );
1126
1127 self.encode_residual(&mut encoder, &self.ffn_out_buf, &self.norm_buf, &self.hidden_buf, hd);
1129
1130 encoder.copy_buffer_to_buffer(&self.hidden_buf, 0, &self.staging_buf, 0, (hd * 4) as u64);
1132 self.queue.submit(Some(encoder.finish()));
1133
1134 let slice = self.staging_buf.slice(..(hd as u64 * 4));
1136 let (tx, rx) = std::sync::mpsc::channel();
1137 slice.map_async(wgpu::MapMode::Read, move |r| {
1138 tx.send(r).ok();
1139 });
1140 self.device.poll(wgpu::PollType::Wait { submission_index: None, timeout: None }).ok();
1141 rx.recv().map_err(|e| format!("recv: {e}"))?.map_err(|e| format!("map: {e:?}"))?;
1142 {
1143 let data = slice.get_mapped_range();
1144 hidden.copy_from_slice(
1145 &bytemuck::cast_slice::<u8, f32>(&data)[..self.hidden_dim as usize],
1146 );
1147 }
1148 self.staging_buf.unmap();
1149
1150 Ok(())
1151 }
1152
1153 fn attention_scores(
1158 q_head: &[f32],
1159 kv_cache_k: &[f32],
1160 kv_offset: usize,
1161 kv_dim: usize,
1162 seq_len: usize,
1163 scale: f32,
1164 ) -> Vec<f32> {
1165 let mut scores = vec![0.0f32; seq_len];
1166 for (s, score) in scores.iter_mut().enumerate() {
1167 let k_offset = s * kv_dim + kv_offset;
1168 let mut dot = 0.0f32;
1169 for (d, q) in q_head.iter().enumerate() {
1170 dot += q * kv_cache_k[k_offset + d];
1171 }
1172 *score = dot * scale;
1173 }
1174
1175 let max_score = scores.iter().copied().fold(f32::NEG_INFINITY, f32::max);
1177 let mut sum = 0.0f32;
1178 for s in scores.iter_mut() {
1179 *s = (*s - max_score).exp();
1180 sum += *s;
1181 }
1182 if sum > 0.0 {
1183 for s in scores.iter_mut() {
1184 *s /= sum;
1185 }
1186 }
1187 scores
1188 }
1189
1190 fn attention_weighted_v(
1194 scores: &[f32],
1195 kv_cache_v: &[f32],
1196 out_head: &mut [f32],
1197 kv_offset: usize,
1198 kv_dim: usize,
1199 ) {
1200 for (d, out) in out_head.iter_mut().enumerate() {
1201 let mut val = 0.0f32;
1202 for (s, score) in scores.iter().enumerate() {
1203 let v_offset = s * kv_dim + kv_offset;
1204 val += score * kv_cache_v[v_offset + d];
1205 }
1206 *out = val;
1207 }
1208 }
1209
1210 pub fn encode_forward_layer_training(
1226 &self,
1227 encoder: &mut wgpu::CommandEncoder,
1228 seq_len: u32,
1229 layer_prefix: &str,
1230 saved: &LayerActivations,
1231 lora: Option<&QkvLoRA<'_>>,
1232 ) -> Result<(), String> {
1233 let hd = self.hidden_dim;
1234 let q_dim = self.num_heads * self.head_dim;
1235 let kv_dim = self.num_kv_heads * self.head_dim;
1236 let inter = self.intermediate_dim;
1237 let s = seq_len as usize;
1238
1239 let norm_w = self
1241 .weight_buffers
1242 .get(&format!("{layer_prefix}.attn_norm"))
1243 .ok_or_else(|| format!("Missing {layer_prefix}.attn_norm"))?;
1244 self.encode_rmsnorm(encoder, &self.hidden_buf, norm_w, &self.norm_buf, hd);
1245
1246 encoder.copy_buffer_to_buffer(
1248 &self.norm_buf,
1249 0,
1250 &saved.attn_norm_out,
1251 0,
1252 (s * hd as usize * 4) as u64,
1253 );
1254
1255 self.encode_matmul(
1257 encoder,
1258 &self.norm_buf,
1259 layer_prefix,
1260 "q_proj",
1261 &self.q_buf,
1262 seq_len,
1263 hd,
1264 q_dim,
1265 );
1266 self.encode_matmul(
1267 encoder,
1268 &self.norm_buf,
1269 layer_prefix,
1270 "k_proj",
1271 &self.k_buf,
1272 seq_len,
1273 hd,
1274 kv_dim,
1275 );
1276 self.encode_matmul(
1277 encoder,
1278 &self.norm_buf,
1279 layer_prefix,
1280 "v_proj",
1281 &self.v_buf,
1282 seq_len,
1283 hd,
1284 kv_dim,
1285 );
1286
1287 if let Some(lora) = lora {
1290 self.encode_lora_addmm(
1291 encoder,
1292 &saved.attn_norm_out,
1293 lora.q_a,
1294 lora.q_b,
1295 &self.q_buf,
1296 seq_len,
1297 lora.in_dim,
1298 lora.rank,
1299 lora.q_dim,
1300 lora.scale,
1301 lora.lora_pipeline,
1302 lora.lora_bgl,
1303 );
1304 self.encode_lora_addmm(
1305 encoder,
1306 &saved.attn_norm_out,
1307 lora.k_a,
1308 lora.k_b,
1309 &self.k_buf,
1310 seq_len,
1311 lora.in_dim,
1312 lora.rank,
1313 lora.kv_dim,
1314 lora.scale,
1315 lora.lora_pipeline,
1316 lora.lora_bgl,
1317 );
1318 self.encode_lora_addmm(
1319 encoder,
1320 &saved.attn_norm_out,
1321 lora.v_a,
1322 lora.v_b,
1323 &self.v_buf,
1324 seq_len,
1325 lora.in_dim,
1326 lora.rank,
1327 lora.kv_dim,
1328 lora.scale,
1329 lora.lora_pipeline,
1330 lora.lora_bgl,
1331 );
1332 }
1333
1334 if let Some(q_bias) = self.cpu_biases.get(&format!("{layer_prefix}.q_bias")) {
1336 self.encode_broadcast_bias(encoder, &self.q_buf, q_bias, seq_len);
1337 }
1338 if let Some(k_bias) = self.cpu_biases.get(&format!("{layer_prefix}.k_bias")) {
1339 self.encode_broadcast_bias(encoder, &self.k_buf, k_bias, seq_len);
1340 }
1341 if let Some(v_bias) = self.cpu_biases.get(&format!("{layer_prefix}.v_bias")) {
1342 self.encode_broadcast_bias(encoder, &self.v_buf, v_bias, seq_len);
1343 }
1344
1345 self.encode_batch_rope(encoder, &self.q_buf, seq_len, self.num_heads, self.head_dim);
1347 self.encode_batch_rope(encoder, &self.k_buf, seq_len, self.num_kv_heads, self.head_dim);
1348
1349 self.encode_attention(encoder, seq_len);
1351
1352 encoder.copy_buffer_to_buffer(
1354 &self.attn_out_buf,
1355 0,
1356 &saved.attn_output,
1357 0,
1358 (s * q_dim as usize * 4) as u64,
1359 );
1360
1361 self.encode_matmul(
1363 encoder,
1364 &self.attn_out_buf,
1365 layer_prefix,
1366 "o_proj",
1367 &self.q_buf,
1368 seq_len,
1369 q_dim,
1370 hd,
1371 );
1372
1373 self.encode_residual(
1375 encoder,
1376 &self.hidden_buf,
1377 &self.q_buf,
1378 &self.ffn_out_buf,
1379 hd * seq_len,
1380 );
1381
1382 let ffn_norm_w = self
1384 .weight_buffers
1385 .get(&format!("{layer_prefix}.ffn_norm"))
1386 .ok_or_else(|| format!("Missing {layer_prefix}.ffn_norm"))?;
1387 self.encode_rmsnorm(encoder, &self.ffn_out_buf, ffn_norm_w, &self.norm_buf, hd);
1388
1389 encoder.copy_buffer_to_buffer(
1391 &self.norm_buf,
1392 0,
1393 &saved.ffn_norm_out,
1394 0,
1395 (s * hd as usize * 4) as u64,
1396 );
1397
1398 self.encode_matmul(
1400 encoder,
1401 &self.norm_buf,
1402 layer_prefix,
1403 "gate_proj",
1404 &self.ffn_gate_buf,
1405 seq_len,
1406 hd,
1407 inter,
1408 );
1409 self.encode_matmul(
1410 encoder,
1411 &self.norm_buf,
1412 layer_prefix,
1413 "up_proj",
1414 &self.ffn_up_buf,
1415 seq_len,
1416 hd,
1417 inter,
1418 );
1419
1420 self.encode_silu_mul(
1422 encoder,
1423 &self.ffn_gate_buf,
1424 &self.ffn_up_buf,
1425 &self.ffn_silu_buf,
1426 inter * seq_len,
1427 );
1428
1429 encoder.copy_buffer_to_buffer(
1431 &self.ffn_silu_buf,
1432 0,
1433 &saved.silu_gate_output,
1434 0,
1435 (s * inter as usize * 4) as u64,
1436 );
1437
1438 self.encode_matmul(
1440 encoder,
1441 &self.ffn_silu_buf,
1442 layer_prefix,
1443 "down_proj",
1444 &self.norm_buf,
1445 seq_len,
1446 inter,
1447 hd,
1448 );
1449
1450 self.encode_residual(
1452 encoder,
1453 &self.ffn_out_buf,
1454 &self.norm_buf,
1455 &self.hidden_buf,
1456 hd * seq_len,
1457 );
1458
1459 Ok(())
1460 }
1461
1462 pub fn forward_layer_traced(
1465 &self,
1466 seq_len: u32,
1467 layer_prefix: &str,
1468 saved: &LayerActivations,
1469 lora: Option<&QkvLoRA<'_>>,
1470 ) -> Result<(), String> {
1471 let hd = self.hidden_dim;
1472 let q_dim = self.num_heads * self.head_dim;
1473 let kv_dim = self.num_kv_heads * self.head_dim;
1474 let inter = self.intermediate_dim;
1475 let s = seq_len as usize;
1476
1477 let norm_w = self
1478 .weight_buffers
1479 .get(&format!("{layer_prefix}.attn_norm"))
1480 .ok_or_else(|| format!("Missing {layer_prefix}.attn_norm"))?;
1481
1482 let mut trace = Vec::new();
1483 let mut run = |name: &str, f: &dyn Fn(&mut wgpu::CommandEncoder)| {
1484 let mut enc = self.device.create_command_encoder(&Default::default());
1485 f(&mut enc);
1486 self.queue.submit(Some(enc.finish()));
1487 let t = std::time::Instant::now();
1488 self.device.poll(wgpu::PollType::Wait { submission_index: None, timeout: None }).ok();
1489 trace.push((name.to_string(), t.elapsed().as_millis() as u64));
1490 };
1491
1492 run("rmsnorm1", &|e| self.encode_rmsnorm(e, &self.hidden_buf, norm_w, &self.norm_buf, hd));
1493 {
1494 let mut e = self.device.create_command_encoder(&Default::default());
1495 e.copy_buffer_to_buffer(
1496 &self.norm_buf,
1497 0,
1498 &saved.attn_norm_out,
1499 0,
1500 (s * hd as usize * 4) as u64,
1501 );
1502 self.queue.submit(Some(e.finish()));
1503 }
1504 run("q_proj", &|e| {
1505 self.encode_matmul(
1506 e,
1507 &self.norm_buf,
1508 layer_prefix,
1509 "q_proj",
1510 &self.q_buf,
1511 seq_len,
1512 hd,
1513 q_dim,
1514 )
1515 });
1516 run("k_proj", &|e| {
1517 self.encode_matmul(
1518 e,
1519 &self.norm_buf,
1520 layer_prefix,
1521 "k_proj",
1522 &self.k_buf,
1523 seq_len,
1524 hd,
1525 kv_dim,
1526 )
1527 });
1528 run("v_proj", &|e| {
1529 self.encode_matmul(
1530 e,
1531 &self.norm_buf,
1532 layer_prefix,
1533 "v_proj",
1534 &self.v_buf,
1535 seq_len,
1536 hd,
1537 kv_dim,
1538 )
1539 });
1540 if let Some(lr) = lora {
1541 run("lora_qkv", &|e| {
1542 self.encode_lora_addmm(
1543 e,
1544 &saved.attn_norm_out,
1545 lr.q_a,
1546 lr.q_b,
1547 &self.q_buf,
1548 seq_len,
1549 lr.in_dim,
1550 lr.rank,
1551 lr.q_dim,
1552 lr.scale,
1553 lr.lora_pipeline,
1554 lr.lora_bgl,
1555 );
1556 self.encode_lora_addmm(
1557 e,
1558 &saved.attn_norm_out,
1559 lr.k_a,
1560 lr.k_b,
1561 &self.k_buf,
1562 seq_len,
1563 lr.in_dim,
1564 lr.rank,
1565 lr.kv_dim,
1566 lr.scale,
1567 lr.lora_pipeline,
1568 lr.lora_bgl,
1569 );
1570 self.encode_lora_addmm(
1571 e,
1572 &saved.attn_norm_out,
1573 lr.v_a,
1574 lr.v_b,
1575 &self.v_buf,
1576 seq_len,
1577 lr.in_dim,
1578 lr.rank,
1579 lr.kv_dim,
1580 lr.scale,
1581 lr.lora_pipeline,
1582 lr.lora_bgl,
1583 );
1584 });
1585 }
1586 if let Some(q_bias) = self.cpu_biases.get(&format!("{layer_prefix}.q_bias")) {
1588 run("q_bias", &|e| self.encode_broadcast_bias(e, &self.q_buf, q_bias, seq_len));
1589 }
1590 if let Some(k_bias) = self.cpu_biases.get(&format!("{layer_prefix}.k_bias")) {
1591 run("k_bias", &|e| self.encode_broadcast_bias(e, &self.k_buf, k_bias, seq_len));
1592 }
1593 if let Some(v_bias) = self.cpu_biases.get(&format!("{layer_prefix}.v_bias")) {
1594 run("v_bias", &|e| self.encode_broadcast_bias(e, &self.v_buf, v_bias, seq_len));
1595 }
1596 run("rope_q", &|e| {
1597 self.encode_batch_rope(e, &self.q_buf, seq_len, self.num_heads, self.head_dim)
1598 });
1599 run("rope_k", &|e| {
1600 self.encode_batch_rope(e, &self.k_buf, seq_len, self.num_kv_heads, self.head_dim)
1601 });
1602 run("attention", &|e| self.encode_attention(e, seq_len));
1603 {
1604 let mut e = self.device.create_command_encoder(&Default::default());
1605 e.copy_buffer_to_buffer(
1606 &self.attn_out_buf,
1607 0,
1608 &saved.attn_output,
1609 0,
1610 (s * q_dim as usize * 4) as u64,
1611 );
1612 self.queue.submit(Some(e.finish()));
1613 }
1614 run("o_proj", &|e| {
1615 self.encode_matmul(
1616 e,
1617 &self.attn_out_buf,
1618 layer_prefix,
1619 "o_proj",
1620 &self.q_buf,
1621 seq_len,
1622 q_dim,
1623 hd,
1624 )
1625 });
1626 run("residual1", &|e| {
1627 self.encode_residual(e, &self.hidden_buf, &self.q_buf, &self.ffn_out_buf, hd * seq_len)
1628 });
1629 let ffn_norm_w = self
1630 .weight_buffers
1631 .get(&format!("{layer_prefix}.ffn_norm"))
1632 .ok_or_else(|| format!("Missing {layer_prefix}.ffn_norm"))?;
1633 run("rmsnorm2", &|e| {
1634 self.encode_rmsnorm(e, &self.ffn_out_buf, ffn_norm_w, &self.norm_buf, hd)
1635 });
1636 {
1637 let mut e = self.device.create_command_encoder(&Default::default());
1638 e.copy_buffer_to_buffer(
1639 &self.norm_buf,
1640 0,
1641 &saved.ffn_norm_out,
1642 0,
1643 (s * hd as usize * 4) as u64,
1644 );
1645 self.queue.submit(Some(e.finish()));
1646 }
1647 run("gate_proj", &|e| {
1648 self.encode_matmul(
1649 e,
1650 &self.norm_buf,
1651 layer_prefix,
1652 "gate_proj",
1653 &self.ffn_gate_buf,
1654 seq_len,
1655 hd,
1656 inter,
1657 )
1658 });
1659 run("up_proj", &|e| {
1660 self.encode_matmul(
1661 e,
1662 &self.norm_buf,
1663 layer_prefix,
1664 "up_proj",
1665 &self.ffn_up_buf,
1666 seq_len,
1667 hd,
1668 inter,
1669 )
1670 });
1671 run("silu", &|e| {
1672 self.encode_silu_mul(
1673 e,
1674 &self.ffn_gate_buf,
1675 &self.ffn_up_buf,
1676 &self.ffn_silu_buf,
1677 inter * seq_len,
1678 )
1679 });
1680 {
1681 let mut e = self.device.create_command_encoder(&Default::default());
1682 e.copy_buffer_to_buffer(
1683 &self.ffn_silu_buf,
1684 0,
1685 &saved.silu_gate_output,
1686 0,
1687 (s * inter as usize * 4) as u64,
1688 );
1689 self.queue.submit(Some(e.finish()));
1690 }
1691 run("down_proj", &|e| {
1692 self.encode_matmul(
1693 e,
1694 &self.ffn_silu_buf,
1695 layer_prefix,
1696 "down_proj",
1697 &self.norm_buf,
1698 seq_len,
1699 inter,
1700 hd,
1701 )
1702 });
1703 run("residual2", &|e| {
1704 self.encode_residual(
1705 e,
1706 &self.ffn_out_buf,
1707 &self.norm_buf,
1708 &self.hidden_buf,
1709 hd * seq_len,
1710 )
1711 });
1712
1713 let total: u64 = trace.iter().map(|(_, ms)| ms).sum();
1714 let parts: Vec<String> = trace.iter().map(|(n, ms)| format!("{n}={ms}")).collect();
1715 eprintln!("[OP-TRACE] layer {} total={}ms: {}", layer_prefix, total, parts.join(" "));
1716 Ok(())
1717 }
1718
1719 pub fn alloc_layer_activations(&self, seq_len: u32) -> LayerActivations {
1721 let s = seq_len as usize;
1722 let buf = |size: usize, label: &str| -> wgpu::Buffer {
1723 self.device.create_buffer(&wgpu::BufferDescriptor {
1724 label: Some(label),
1725 size: (size * 4) as u64,
1726 usage: wgpu::BufferUsages::STORAGE
1727 | wgpu::BufferUsages::COPY_SRC
1728 | wgpu::BufferUsages::COPY_DST,
1729 mapped_at_creation: false,
1730 })
1731 };
1732 LayerActivations {
1733 attn_norm_out: buf(s * self.hidden_dim as usize, "saved_attn_norm"),
1734 attn_output: buf(s * (self.num_heads * self.head_dim) as usize, "saved_attn_out"),
1735 ffn_norm_out: buf(s * self.hidden_dim as usize, "saved_ffn_norm"),
1736 silu_gate_output: buf(s * self.intermediate_dim as usize, "saved_silu"),
1737 rstd_attn: buf(s, "saved_rstd_attn"),
1738 rstd_ffn: buf(s, "saved_rstd_ffn"),
1739 softmax_logsumexp: buf(self.num_heads as usize * s, "saved_logsumexp"),
1740 }
1741 }
1742
1743 pub fn forward_layer_training(
1745 &self,
1746 seq_len: u32,
1747 layer_prefix: &str,
1748 ) -> Result<LayerActivations, String> {
1749 let saved = self.alloc_layer_activations(seq_len);
1750 let mut encoder = self.device.create_command_encoder(&Default::default());
1751 self.encode_forward_layer_training(&mut encoder, seq_len, layer_prefix, &saved, None)?;
1752 self.queue.submit(Some(encoder.finish()));
1753 Ok(saved)
1754 }
1755
1756 pub fn forward_all_layers_training(
1758 &self,
1759 seq_len: u32,
1760 num_layers: usize,
1761 ) -> Result<Vec<LayerActivations>, String> {
1762 let mut encoder = self.device.create_command_encoder(&Default::default());
1763 let mut all_saved = Vec::with_capacity(num_layers);
1764
1765 for layer_idx in 0..num_layers {
1766 let prefix = format!("layer.{layer_idx}");
1767 let saved = self.alloc_layer_activations(seq_len);
1768 self.encode_forward_layer_training(&mut encoder, seq_len, &prefix, &saved, None)?;
1769 all_saved.push(saved);
1770 }
1771
1772 self.queue.submit(Some(encoder.finish()));
1774 Ok(all_saved)
1775 }
1776 pub fn encode_broadcast_bias(
1784 &self,
1785 encoder: &mut wgpu::CommandEncoder,
1786 buf: &wgpu::Buffer,
1787 bias: &[f32],
1788 seq_len: u32,
1789 ) {
1790 let dim = bias.len();
1791 let mut full_bias = Vec::with_capacity(seq_len as usize * dim);
1793 for _ in 0..seq_len {
1794 full_bias.extend_from_slice(bias);
1795 }
1796 let bias_buf = self.device.create_buffer(&wgpu::BufferDescriptor {
1797 label: Some("broadcast_bias"),
1798 size: (full_bias.len() * 4) as u64,
1799 usage: wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_DST,
1800 mapped_at_creation: false,
1801 });
1802 self.queue.write_buffer(&bias_buf, 0, bytemuck::cast_slice(&full_bias));
1803
1804 let total = seq_len * dim as u32;
1806 let tmp = self.device.create_buffer(&wgpu::BufferDescriptor {
1807 label: Some("bias_tmp"),
1808 size: (total as usize * 4) as u64,
1809 usage: wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_SRC,
1810 mapped_at_creation: false,
1811 });
1812 self.encode_residual(encoder, buf, &bias_buf, &tmp, total);
1813 encoder.copy_buffer_to_buffer(&tmp, 0, buf, 0, (total as u64) * 4);
1814 }
1815
1816 fn encode_batch_rope(
1819 &self,
1820 encoder: &mut wgpu::CommandEncoder,
1821 qk_buf: &wgpu::Buffer,
1822 seq_len: u32,
1823 num_heads: u32,
1824 head_dim: u32,
1825 ) {
1826 #[repr(C)]
1827 #[derive(Copy, Clone, bytemuck::Pod, bytemuck::Zeroable)]
1828 struct RopeParams {
1829 seq_len: u32,
1830 num_heads: u32,
1831 head_dim: u32,
1832 _pad: u32,
1833 }
1834 let params = RopeParams { seq_len, num_heads, head_dim, _pad: 0 };
1835 let params_buf = self.device.create_buffer(&wgpu::BufferDescriptor {
1836 label: Some("batch_rope_params"),
1837 size: 16,
1838 usage: wgpu::BufferUsages::UNIFORM | wgpu::BufferUsages::COPY_DST,
1839 mapped_at_creation: false,
1840 });
1841 self.queue.write_buffer(¶ms_buf, 0, bytemuck::bytes_of(¶ms));
1842
1843 let bg = self.device.create_bind_group(&wgpu::BindGroupDescriptor {
1844 label: Some("batch_rope_bg"),
1845 layout: &self.batch_rope_bgl,
1846 entries: &[
1847 wgpu::BindGroupEntry { binding: 0, resource: qk_buf.as_entire_binding() },
1848 wgpu::BindGroupEntry { binding: 1, resource: params_buf.as_entire_binding() },
1849 ],
1850 });
1851 let total = seq_len * num_heads * head_dim;
1852 let wg = total.div_ceil(256);
1853 let mut pass = encoder.begin_compute_pass(&Default::default());
1854 pass.set_pipeline(&self.batch_rope_pipeline);
1855 pass.set_bind_group(0, &bg, &[]);
1856 pass.dispatch_workgroups(wg, 1, 1);
1857 }
1858
1859 fn encode_attention(&self, encoder: &mut wgpu::CommandEncoder, seq_len: u32) {
1860 let params = [seq_len, self.num_heads, self.num_kv_heads, self.head_dim];
1861 let params_buf = self.make_uniform(¶ms);
1862 let _q_dim = self.num_heads * self.head_dim;
1863
1864 let bg = self.device.create_bind_group(&wgpu::BindGroupDescriptor {
1869 label: None,
1870 layout: &self.attention_bgl,
1871 entries: &[
1872 wgpu::BindGroupEntry { binding: 0, resource: self.q_buf.as_entire_binding() },
1873 wgpu::BindGroupEntry { binding: 1, resource: self.k_buf.as_entire_binding() },
1874 wgpu::BindGroupEntry { binding: 2, resource: self.v_buf.as_entire_binding() },
1875 wgpu::BindGroupEntry {
1876 binding: 3,
1877 resource: self.attn_out_buf.as_entire_binding(),
1878 },
1879 wgpu::BindGroupEntry { binding: 4, resource: params_buf.as_entire_binding() },
1880 ],
1881 });
1882 let mut pass = encoder.begin_compute_pass(&Default::default());
1883 pass.set_pipeline(&self.attention_pipeline);
1884 pass.set_bind_group(0, &bg, &[]);
1885 pass.dispatch_workgroups(self.num_heads, seq_len, 1);
1887 }
1888
1889 #[allow(clippy::too_many_arguments)]
1899 fn encode_lora_addmm(
1900 &self,
1901 encoder: &mut wgpu::CommandEncoder,
1902 input: &wgpu::Buffer,
1903 lora_a: &wgpu::Buffer,
1904 lora_b: &wgpu::Buffer,
1905 output: &wgpu::Buffer,
1906 seq_len: u32,
1907 in_dim: u32,
1908 rank: u32,
1909 out_dim: u32,
1910 scale: f32,
1911 _pipeline: &wgpu::ComputePipeline,
1912 _bgl: &wgpu::BindGroupLayout,
1913 ) {
1914 let temp_size = (seq_len * rank) as u64 * 4;
1916 let temp = self.device.create_buffer(&wgpu::BufferDescriptor {
1917 label: Some("lora_temp"),
1918 size: temp_size,
1919 usage: wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_SRC,
1920 mapped_at_creation: false,
1921 });
1922 self.encode_tiled_gemm(encoder, input, lora_a, &temp, seq_len, in_dim, rank, 1.0);
1923
1924 let delta_size = (seq_len * out_dim) as u64 * 4;
1926 let delta = self.device.create_buffer(&wgpu::BufferDescriptor {
1927 label: Some("lora_delta"),
1928 size: delta_size,
1929 usage: wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_SRC,
1930 mapped_at_creation: false,
1931 });
1932 self.encode_tiled_gemm(encoder, &temp, lora_b, &delta, seq_len, rank, out_dim, scale);
1933
1934 let sum_buf = self.device.create_buffer(&wgpu::BufferDescriptor {
1936 label: Some("lora_sum"),
1937 size: delta_size,
1938 usage: wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_SRC,
1939 mapped_at_creation: false,
1940 });
1941 self.encode_residual(encoder, output, &delta, &sum_buf, seq_len * out_dim);
1942 encoder.copy_buffer_to_buffer(&sum_buf, 0, output, 0, delta_size);
1943 }
1944
1945 fn encode_tiled_gemm(
1947 &self,
1948 encoder: &mut wgpu::CommandEncoder,
1949 a: &wgpu::Buffer,
1950 b: &wgpu::Buffer,
1951 c: &wgpu::Buffer,
1952 m: u32,
1953 k: u32,
1954 n: u32,
1955 alpha: f32,
1956 ) {
1957 let params = [m, k, n, alpha.to_bits()];
1958 let params_buf = self.make_uniform(¶ms);
1959 let bg = self.device.create_bind_group(&wgpu::BindGroupDescriptor {
1960 label: None,
1961 layout: &self.matmul_bgl,
1962 entries: &[
1963 wgpu::BindGroupEntry { binding: 0, resource: a.as_entire_binding() },
1964 wgpu::BindGroupEntry { binding: 1, resource: b.as_entire_binding() },
1965 wgpu::BindGroupEntry { binding: 2, resource: c.as_entire_binding() },
1966 wgpu::BindGroupEntry { binding: 3, resource: params_buf.as_entire_binding() },
1967 ],
1968 });
1969 let mut pass = encoder.begin_compute_pass(&Default::default());
1970 pass.set_pipeline(&self.tiled_matmul_pipeline);
1971 pass.set_bind_group(0, &bg, &[]);
1972 pass.dispatch_workgroups(n.div_ceil(64), m.div_ceil(64), 1);
1973 }
1974
1975 fn encode_rmsnorm(
1976 &self,
1977 encoder: &mut wgpu::CommandEncoder,
1978 input: &wgpu::Buffer,
1979 weight: &wgpu::Buffer,
1980 output: &wgpu::Buffer,
1981 dim: u32,
1982 ) {
1983 let params = rmsnorm_params(dim, self.rms_norm_eps);
1984 let params_buf = self.make_uniform(¶ms);
1985 let bg = self.device.create_bind_group(&wgpu::BindGroupDescriptor {
1986 label: None,
1987 layout: &self.elementwise_bgl,
1988 entries: &[
1989 wgpu::BindGroupEntry { binding: 0, resource: input.as_entire_binding() },
1990 wgpu::BindGroupEntry { binding: 1, resource: weight.as_entire_binding() },
1991 wgpu::BindGroupEntry { binding: 2, resource: output.as_entire_binding() },
1992 wgpu::BindGroupEntry { binding: 3, resource: params_buf.as_entire_binding() },
1993 ],
1994 });
1995 let num_rows = (input.size() / (dim as u64 * 4)).max(1) as u32;
1998 let mut pass = encoder.begin_compute_pass(&Default::default());
1999 pass.set_pipeline(&self.rmsnorm_pipeline);
2000 pass.set_bind_group(0, &bg, &[]);
2001 pass.dispatch_workgroups(1, num_rows, 1);
2002 }
2003
2004 fn encode_matmul(
2005 &self,
2006 encoder: &mut wgpu::CommandEncoder,
2007 input: &wgpu::Buffer,
2008 layer_prefix: &str,
2009 proj_name: &str,
2010 output: &wgpu::Buffer,
2011 m: u32,
2012 k: u32,
2013 n: u32,
2014 ) {
2015 if m == 1 && self.encode_q4k_gemv(encoder, input, output, layer_prefix, proj_name, n, k) {
2017 return;
2018 }
2019 let weight_key = format!("{layer_prefix}.{proj_name}");
2020 let weight = match self.weight_buffers.get(&weight_key) {
2021 Some(w) => w,
2022 None => return, };
2024 let params = if m == 1 { [n, k, 0u32, 0u32] } else { [m, k, n, 1.0_f32.to_bits()] };
2029 let params_buf = self.make_uniform(¶ms);
2030 let bg = self.device.create_bind_group(&wgpu::BindGroupDescriptor {
2031 label: None,
2032 layout: &self.matmul_bgl,
2033 entries: &[
2034 wgpu::BindGroupEntry { binding: 0, resource: input.as_entire_binding() },
2035 wgpu::BindGroupEntry { binding: 1, resource: weight.as_entire_binding() },
2036 wgpu::BindGroupEntry { binding: 2, resource: output.as_entire_binding() },
2037 wgpu::BindGroupEntry { binding: 3, resource: params_buf.as_entire_binding() },
2038 ],
2039 });
2040 let mut pass = encoder.begin_compute_pass(&Default::default());
2041 if m == 1 {
2042 pass.set_pipeline(&self.gemv_pipeline);
2044 pass.set_bind_group(0, &bg, &[]);
2045 pass.dispatch_workgroups(n, 1, 1);
2046 } else if m >= 4 {
2047 pass.set_pipeline(&self.tiled_matmul_pipeline);
2050 pass.set_bind_group(0, &bg, &[]);
2051 pass.dispatch_workgroups(n.div_ceil(64), m.div_ceil(64), 1);
2052 } else {
2053 pass.set_pipeline(&self.matmul_pipeline);
2055 pass.set_bind_group(0, &bg, &[]);
2056 pass.dispatch_workgroups(m.div_ceil(16), n.div_ceil(16), 1);
2057 }
2058 }
2059
2060 fn encode_q4k_gemv(
2064 &self,
2065 encoder: &mut wgpu::CommandEncoder,
2066 input: &wgpu::Buffer,
2067 output: &wgpu::Buffer,
2068 layer_prefix: &str,
2069 proj_name: &str,
2070 n: u32,
2071 k: u32,
2072 ) -> bool {
2073 let weight_key = format!("{layer_prefix}.{proj_name}");
2074 let weight = match self.q4k_weights.get(&weight_key) {
2075 Some(w) => w,
2076 None => return false,
2077 };
2078 let num_superblocks = (k + 255) / 256;
2079 let params = [n, k, num_superblocks, 0u32];
2080 let params_buf = self.make_uniform(¶ms);
2081 let bg = self.device.create_bind_group(&wgpu::BindGroupDescriptor {
2082 label: None,
2083 layout: &self.matmul_bgl,
2084 entries: &[
2085 wgpu::BindGroupEntry { binding: 0, resource: input.as_entire_binding() },
2086 wgpu::BindGroupEntry { binding: 1, resource: weight.as_entire_binding() },
2087 wgpu::BindGroupEntry { binding: 2, resource: output.as_entire_binding() },
2088 wgpu::BindGroupEntry { binding: 3, resource: params_buf.as_entire_binding() },
2089 ],
2090 });
2091 let mut pass = encoder.begin_compute_pass(&Default::default());
2092 pass.set_pipeline(&self.q4k_gemv_pipeline);
2093 pass.set_bind_group(0, &bg, &[]);
2094 pass.dispatch_workgroups(n, 1, 1);
2095 true
2096 }
2097
2098 fn encode_silu_mul(
2099 &self,
2100 encoder: &mut wgpu::CommandEncoder,
2101 gate: &wgpu::Buffer,
2102 up: &wgpu::Buffer,
2103 output: &wgpu::Buffer,
2104 dim: u32,
2105 ) {
2106 let params = [dim, 0u32, 0, 0];
2107 let params_buf = self.make_uniform(¶ms);
2108 let bg = self.device.create_bind_group(&wgpu::BindGroupDescriptor {
2109 label: None,
2110 layout: &self.elementwise_bgl,
2111 entries: &[
2112 wgpu::BindGroupEntry { binding: 0, resource: gate.as_entire_binding() },
2113 wgpu::BindGroupEntry { binding: 1, resource: up.as_entire_binding() },
2114 wgpu::BindGroupEntry { binding: 2, resource: output.as_entire_binding() },
2115 wgpu::BindGroupEntry { binding: 3, resource: params_buf.as_entire_binding() },
2116 ],
2117 });
2118 let mut pass = encoder.begin_compute_pass(&Default::default());
2119 pass.set_pipeline(&self.silu_mul_pipeline);
2120 pass.set_bind_group(0, &bg, &[]);
2121 pass.dispatch_workgroups(dim.div_ceil(256), 1, 1);
2122 }
2123
2124 fn encode_residual(
2125 &self,
2126 encoder: &mut wgpu::CommandEncoder,
2127 a: &wgpu::Buffer,
2128 b: &wgpu::Buffer,
2129 output: &wgpu::Buffer,
2130 dim: u32,
2131 ) {
2132 let params = [dim, 0u32, 0, 0];
2133 let params_buf = self.make_uniform(¶ms);
2134 let bg = self.device.create_bind_group(&wgpu::BindGroupDescriptor {
2135 label: None,
2136 layout: &self.elementwise_bgl,
2137 entries: &[
2138 wgpu::BindGroupEntry { binding: 0, resource: a.as_entire_binding() },
2139 wgpu::BindGroupEntry { binding: 1, resource: b.as_entire_binding() },
2140 wgpu::BindGroupEntry { binding: 2, resource: output.as_entire_binding() },
2141 wgpu::BindGroupEntry { binding: 3, resource: params_buf.as_entire_binding() },
2142 ],
2143 });
2144 let mut pass = encoder.begin_compute_pass(&Default::default());
2145 pass.set_pipeline(&self.residual_pipeline);
2146 pass.set_bind_group(0, &bg, &[]);
2147 pass.dispatch_workgroups(dim.div_ceil(256), 1, 1);
2148 }
2149
2150 fn make_uniform(&self, data: &[u32; 4]) -> wgpu::Buffer {
2151 use wgpu::util::DeviceExt;
2152 self.device.create_buffer_init(&wgpu::util::BufferInitDescriptor {
2153 label: None,
2154 contents: bytemuck::cast_slice(data),
2155 usage: wgpu::BufferUsages::UNIFORM,
2156 })
2157 }
2158}
2159
2160fn bgl_storage(binding: u32, read_only: bool) -> wgpu::BindGroupLayoutEntry {
2161 wgpu::BindGroupLayoutEntry {
2162 binding,
2163 visibility: wgpu::ShaderStages::COMPUTE,
2164 ty: wgpu::BindingType::Buffer {
2165 ty: wgpu::BufferBindingType::Storage { read_only },
2166 has_dynamic_offset: false,
2167 min_binding_size: None,
2168 },
2169 count: None,
2170 }
2171}
2172
2173fn bgl_uniform(binding: u32) -> wgpu::BindGroupLayoutEntry {
2174 wgpu::BindGroupLayoutEntry {
2175 binding,
2176 visibility: wgpu::ShaderStages::COMPUTE,
2177 ty: wgpu::BindingType::Buffer {
2178 ty: wgpu::BufferBindingType::Uniform,
2179 has_dynamic_offset: false,
2180 min_binding_size: None,
2181 },
2182 count: None,
2183 }
2184}
2185
2186#[cfg(test)]
2187mod rmsnorm_eps_tests {
2188 use super::{rmsnorm_params, DEFAULT_RMS_NORM_EPS};
2189
2190 #[test]
2193 fn rmsnorm_params_carry_eps_bits() {
2194 let p = rmsnorm_params(896, 1e-5);
2195 assert_eq!(p[0], 896);
2196 assert_eq!(f32::from_bits(p[1]), 1e-5);
2197 assert_ne!(f32::from_bits(p[1]), DEFAULT_RMS_NORM_EPS);
2198 assert_eq!(f32::from_bits(rmsnorm_params(1, DEFAULT_RMS_NORM_EPS)[1]), 1e-6);
2199 }
2200
2201 #[test]
2204 fn rmsnorm_shader_validates_on_device() {
2205 let Ok(gpu) = crate::backends::gpu::GpuDevice::new() else {
2206 return;
2207 };
2208 gpu.device.push_error_scope(wgpu::ErrorFilter::Validation);
2209 let _module = gpu.device.create_shader_module(wgpu::ShaderModuleDescriptor {
2210 label: Some("rmsnorm_eps_test"),
2211 source: wgpu::ShaderSource::Wgsl(super::RMSNORM_SHADER.into()),
2212 });
2213 let err = pollster::block_on(gpu.device.pop_error_scope());
2214 assert!(err.is_none(), "RMSNORM_SHADER failed validation: {err:?}");
2215 }
2216}