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