Skip to main content

trueno/backends/gpu/device/linalg/
wgsl_forward.rs

1#![allow(dead_code, clippy::many_single_char_names)]
2//! PMAT-324: WGSL transformer forward pass — multi-pass single submission.
3//!
4//! Instead of one matmul per CPU call (2ms roundtrip each), this encodes
5//! ALL operations for one transformer layer into a single command encoder.
6//! Only one submit + one readback per layer (or per full forward pass).
7//!
8//! Architecture: separate WGSL kernels per operation type, dispatched
9//! sequentially within one command encoder. All intermediate data stays
10//! GPU-resident in persistent buffers.
11
12use std::collections::HashMap;
13
14/// Saved activations for one transformer layer's backward pass.
15///
16/// Contains the 7 tensors needed for LoRA gradient computation
17/// without replaying the forward pass (§26.11.5, falsification-verified).
18pub struct LayerActivations {
19    /// Input to Q/K/V projections (RMSNorm output). [seq, hidden]
20    pub attn_norm_out: wgpu::Buffer,
21    /// Input to O projection (attention output). [seq, q_dim]
22    pub attn_output: wgpu::Buffer,
23    /// Input to gate/up/down projections (FFN RMSNorm output). [seq, hidden]
24    pub ffn_norm_out: wgpu::Buffer,
25    /// Input to down projection (SiLU(gate)×up). [seq, intermediate]
26    pub silu_gate_output: wgpu::Buffer,
27    /// RMSNorm reciprocal std for attention norm. [seq]
28    pub rstd_attn: wgpu::Buffer,
29    /// RMSNorm reciprocal std for FFN norm. [seq]
30    pub rstd_ffn: wgpu::Buffer,
31    /// Softmax logsumexp for attention backward. [num_heads, seq]
32    pub softmax_logsumexp: wgpu::Buffer,
33}
34
35/// Optional LoRA buffers for Q/K/V projections in a layer's forward pass.
36pub 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
52/// GPU-resident transformer layer state.
53/// All buffers persist across tokens — only input/output change per step.
54pub struct WgslForwardPass {
55    device: wgpu::Device,
56    queue: wgpu::Queue,
57
58    // Kernels (compiled once)
59    matmul_pipeline: wgpu::ComputePipeline,
60    /// CUTLASS-style tiled GEMM for M>=4 (training batch, prefill)
61    tiled_matmul_pipeline: wgpu::ComputePipeline,
62    /// PMAT-327: GEMV pipeline for M=1 decode (cooperative K-reduction)
63    gemv_pipeline: wgpu::ComputePipeline,
64    /// C-WGPU-Q4K-001: Q4K GEMV pipeline — dequantize-on-the-fly, no F32 weights
65    q4k_gemv_pipeline: wgpu::ComputePipeline,
66    /// Causal attention pipeline for training (full sequence, no KV cache)
67    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    // Bind group layouts
77    matmul_bgl: wgpu::BindGroupLayout,
78    elementwise_bgl: wgpu::BindGroupLayout,
79
80    // Weight buffers (persistent, uploaded once)
81    weight_buffers: HashMap<String, wgpu::Buffer>,
82    /// GH-560: Raw Q4K weight buffers for fused dequant+GEMV.
83    q4k_weights: HashMap<String, wgpu::Buffer>,
84    /// PMAT-342: CPU-side bias data (small, not worth GPU dispatch)
85    cpu_biases: HashMap<String, Vec<f32>>,
86    /// GH-560: Per-layer GPU KV cache buffers.
87    kv_cache_k: Vec<wgpu::Buffer>,
88    /// GH-560: Per-layer GPU KV cache buffers (values).
89    kv_cache_v: Vec<wgpu::Buffer>,
90
91    // Intermediate buffers (persistent, reused across calls)
92    // For 1.5B: hidden=1536, kv=256, intermediate=8960
93    hidden_buf: wgpu::Buffer,   // [hidden_dim] working state
94    q_buf: wgpu::Buffer,        // [q_dim]
95    k_buf: wgpu::Buffer,        // [kv_dim]
96    v_buf: wgpu::Buffer,        // [kv_dim]
97    attn_out_buf: wgpu::Buffer, // [hidden_dim]
98    ffn_gate_buf: wgpu::Buffer, // [intermediate_dim]
99    ffn_up_buf: wgpu::Buffer,   // [intermediate_dim]
100    ffn_silu_buf: wgpu::Buffer, // [intermediate_dim] — SiLU(gate)×up output (can't alias inputs)
101    ffn_out_buf: wgpu::Buffer,  // [hidden_dim]
102    norm_buf: wgpu::Buffer,     // [hidden_dim] for RMSNorm output
103    staging_buf: wgpu::Buffer,  // readback
104
105    // Config
106    hidden_dim: u32,
107    num_heads: u32,
108    num_kv_heads: u32,
109    head_dim: u32,
110    intermediate_dim: u32,
111    /// RMSNorm epsilon, passed to the shader as `params.y` (f32 bits). Defaults to
112    /// [`DEFAULT_RMS_NORM_EPS`]; callers set the model's value with
113    /// [`WgslForwardPass::set_rms_norm_eps`] (#4056: Llama-family models use 1e-5).
114    rms_norm_eps: f32,
115}
116
117/// Default RMSNorm epsilon: the constant the shader hardcoded before #4056.
118pub const DEFAULT_RMS_NORM_EPS: f32 = 1e-6;
119
120/// Uniform params for the RMSNorm shader: `(dim, eps as f32 bits, 0, 0)`.
121fn rmsnorm_params(dim: u32, eps: f32) -> [u32; 4] {
122    [dim, eps.to_bits(), 0, 0]
123}
124
125// WGSL shader source for RMSNorm (multi-row via workgroup_id.y)
126// Dispatch: (1, seq_len, 1) — one workgroup per row.
127const 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
176// WGSL shader for SiLU(gate) * up
177const 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
193// WGSL shader for residual add: output = a + b
194const 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
208// Batch RoPE shader — applies RoPE to all positions in a sequence at once.
209// PMAT-509: Training forward path was missing RoPE entirely, causing loss > random.
210// Input: qk[seq_len * num_heads * head_dim], applies position-dependent rotation.
211const 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
258// RoPE shader (NeoX-style interleaved) — single position (inference)
259const 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    /// Get the shader sources for external inspection/testing
295    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    /// PMAT-325: Create a new WGSL forward pass context.
309    ///
310    /// Compiles all shader pipelines and allocates persistent intermediate buffers.
311    /// Call once at model init. All GPU resources persist until dropped.
312    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        // Compile shaders
325        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        // Bind group layouts
347        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        // Pipelines
367        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        // CUTLASS-style tiled GEMM for M>=4 (training batch, prefill)
387        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        // Causal attention pipeline for training
397        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),  // Q
407                bgl_storage(1, true),  // K
408                bgl_storage(2, true),  // V
409                bgl_storage(3, false), // output
410                bgl_uniform(4),        // params
411            ],
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        // PMAT-327: GEMV pipeline — same bind group layout as matmul but cooperative reduction
428        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        // C-WGPU-Q4K-001: Q4K GEMV — dequantize on-the-fly, no F32 weight buffer
435        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        // RoPE has a 2-binding layout (in-place + uniform)
446        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        // PMAT-509: Batch RoPE for training (all positions at once)
467        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        // Allocate persistent intermediate buffers
492        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        // Buffer sizes: max_seq × dim for training, or 1 × dim for inference.
504        // Training calls forward_layer_training with seq_len > 1.
505        // Allocate for max_seq=2048 to support both.
506        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    /// Set the RMSNorm epsilon every norm in this forward pass uses — pass the
569    /// model's configured `rms_norm_eps`. Default: [`DEFAULT_RMS_NORM_EPS`].
570    pub fn set_rms_norm_eps(&mut self, eps: f32) {
571        self.rms_norm_eps = eps;
572    }
573
574    /// The RMSNorm epsilon this forward pass uses.
575    #[must_use]
576    pub fn rms_norm_eps(&self) -> f32 {
577        self.rms_norm_eps
578    }
579
580    /// Upload a weight matrix (call once per layer at init).
581    /// PMAT-342: Bias weights (name contains "bias") are stored CPU-side.
582    pub fn upload_weight(&mut self, name: &str, data: &[f32]) {
583        if name.contains("bias") {
584            // Biases are small, keep on CPU for easy access in attention
585            self.cpu_biases.insert(name.to_string(), data.to_vec());
586            return;
587        }
588        // Skip weights that exceed the device's max buffer binding size (e.g., lm_head > 2 GB)
589        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    /// GH-560: Upload raw Q4K weight bytes for fused dequant+GEMV on GPU.
610    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    /// GH-560: Initialize per-layer KV cache buffers on GPU.
621    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    /// Number of uploaded weight buffers.
647    pub fn weight_count(&self) -> usize {
648        self.weight_buffers.len()
649    }
650
651    /// Access a dequantized weight buffer by name (e.g. "layer.0.down_proj").
652    /// Used by backward pass for gradient propagation through frozen base weights.
653    pub fn weight_buffer(&self, name: &str) -> Option<&wgpu::Buffer> {
654        self.weight_buffers.get(name)
655    }
656
657    /// Reference to the wgpu device.
658    pub fn device_ref(&self) -> &wgpu::Device {
659        &self.device
660    }
661
662    /// Reference to the wgpu queue.
663    pub fn queue_ref(&self) -> &wgpu::Queue {
664        &self.queue
665    }
666
667    /// Reference to the hidden state buffer (for writing input).
668    pub fn hidden_buffer(&self) -> &wgpu::Buffer {
669        &self.hidden_buf
670    }
671
672    /// Reference to Q buffer (for LoRA addmm after Q projection).
673    pub fn q_buffer(&self) -> &wgpu::Buffer {
674        &self.q_buf
675    }
676
677    /// Reference to K buffer.
678    pub fn k_buffer(&self) -> &wgpu::Buffer {
679        &self.k_buf
680    }
681
682    /// Reference to V buffer.
683    pub fn v_buffer(&self) -> &wgpu::Buffer {
684        &self.v_buf
685    }
686
687    /// Elementwise add: output = a + b. Dispatches residual add shader.
688    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    /// Apply RMSNorm on GPU: normed = rmsnorm(hidden_buf, weight) → output_buf.
701    /// Contract: gpu-output-norm-v1 / gpu_resident — hidden state never leaves GPU.
702    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    /// Download hidden state from GPU.
709    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    /// Total VRAM used by all buffers (bytes).
739    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  // hidden, attn_out, ffn_out, norm
742            + (self.num_heads as usize * self.head_dim as usize * 4) // q
743            + (self.num_kv_heads as usize * self.head_dim as usize * 4) * 2 // k, v
744            + (self.intermediate_dim as usize * 4) * 2; // gate, up
745        weight_bytes + intermediate_bytes
746    }
747
748    /// PMAT-336: Full model forward — embedding + all layers + output norm + LM head.
749    ///
750    /// Returns logits [vocab_size] for the given token at the given position.
751    /// Embedding lookup and final LM head are CPU-side (not yet GPU-accelerated).
752    /// PMAT-344: Added kv_caches for multi-token context
753    #[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        // 1. Embedding lookup (CPU)
769        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        // 2. Transformer layers (GPU via forward_layer with KV cache)
780        // Initialize KV caches if empty
781        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        // 3. Output RMSNorm (CPU — small, not worth GPU dispatch)
791        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        // 4. LM head — CPU matmul
797        // PMAT-346: GPU tiled GEMM expects weight in [K,N] layout but lm_head is [N,K].
798        // CPU path reads weight[v * hd + j] which matches the [vocab, hidden] layout.
799        // GPU LM head via GEMV is blocked by vocab > 65535 dispatch limit.
800        // Deferred (PMAT-751): add upload_weight_transposed() for GPU-accelerated LM head.
801        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    /// PMAT-325: Execute one transformer layer — 14 passes, 1 submit, 1 readback.
814    ///
815    /// Input: hidden state [hidden_dim] on CPU.
816    /// Output: updated hidden state [hidden_dim] on CPU.
817    /// All intermediate computation stays GPU-resident.
818    /// PMAT-344: KV cache parameters for multi-token context
819    pub fn forward_layer(
820        &self,
821        hidden: &mut [f32],
822        layer_prefix: &str,
823        _position: usize,
824        kv_cache_k: &mut Vec<f32>, // accumulated K: [seq_len * kv_dim]
825        kv_cache_v: &mut Vec<f32>, // accumulated V: [seq_len * kv_dim]
826    ) -> Result<(), String> {
827        let hd = self.hidden_dim;
828
829        // Upload hidden state
830        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        // Pass 1: RMSNorm(hidden → norm_buf)
835        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        // Passes 2-4: Q/K/V projections (norm_buf × W → q/k/v_buf)
842        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        // PMAT-342: Submit Q/K/V projections, readback, do attention on CPU
877        // GPU handles the heavy matmuls; CPU handles attention (small at M=1)
878        let q_bytes = (q_dim * 4) as u64;
879        let kv_bytes = (kv_dim * 4) as u64;
880
881        // Readback Q/K/V from GPU
882        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        // Readback Q
906        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        // Readback K
921        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        // Readback V
936        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        // PMAT-342: Add QKV biases (required for Qwen2)
951        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        // PMAT-343: Apply RoPE (NeoX-style interleaved) to Q and K
968        let head_dim = self.head_dim as usize;
969        let position = _position; // Use the position parameter
970        let rope_theta = 1_000_000.0f64; // Qwen2 rope_theta
971
972        // RoPE on Q (num_heads × head_dim)
973        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        // RoPE on K (num_kv_heads × head_dim)
989        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        // PMAT-344: Append K,V to cache and compute full attention
1005        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        // Scaled dot-product attention with GQA
1015        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            // Compute attention scores: Q[h] · K[kv_h, :seq_len]^T / sqrt(d)
1025            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            // Weighted sum of V
1035            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        // Upload attention output back to GPU for O projection
1046        self.queue.write_buffer(&self.q_buf, 0, bytemuck::cast_slice(&attn_out));
1047
1048        // New encoder for remaining passes
1049        let mut encoder = self.device.create_command_encoder(&Default::default());
1050
1051        // Pass 7: O projection (attn_out × W_o → attn_out_buf)
1052        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        // Pass 8: Residual(hidden + attn_out → hidden)
1064        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        // Pass 9: FFN RMSNorm(ffn_out → norm_buf)
1073        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        // Passes 10-11: Gate + Up projections
1080        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        // Pass 12: SiLU(gate) × up → ffn_silu_buf [intermediate_dim]
1103        // BUG FIX: was writing to attn_out_buf (hidden_dim=3584) but needs intermediate_dim=18944.
1104        // attn_out_buf is only hidden_dim — wgpu robustness silently drops OOB writes,
1105        // then down_proj reads zeros past hidden_dim → 81% of FFN truncated → garbage output.
1106        // Cannot alias gate/up buffers (WGSL read/write aliasing UB), so use dedicated buffer.
1107        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        // Pass 13: Down projection (reads ffn_silu_buf [intermediate_dim] → norm_buf [hidden_dim])
1116        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        // Pass 14: Residual(ffn_out + down → hidden)
1128        self.encode_residual(&mut encoder, &self.ffn_out_buf, &self.norm_buf, &self.hidden_buf, hd);
1129
1130        // Single readback
1131        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        // Readback
1135        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    /// Softmaxed attention scores for one head over `seq_len` cached steps.
1154    ///
1155    /// `kv_offset` is this head's offset inside each cached step, and each
1156    /// cached step is `kv_dim` floats wide. Head dimension is `q_head.len()`.
1157    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        // Softmax
1176        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    /// Weighted sum of V for one head, written into `out_head`.
1191    ///
1192    /// `kv_offset` and `kv_dim` are as in `attention_scores`.
1193    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    /// Training forward pass for a single transformer layer.
1211    ///
1212    /// Unlike `forward_layer` (M=1 decode), this processes the full sequence
1213    /// at once (M=seq_len) and keeps everything on GPU. No CPU readback.
1214    ///
1215    /// Saves `norm_output` (pre-projection activations) for backward pass.
1216    ///
1217    /// # Arguments
1218    /// - `seq_len`: number of tokens in the sequence
1219    /// - `layer_prefix`: e.g. "model.layers.0"
1220    /// - `saved_norm_attn`: OUTPUT — saved pre-attention norm for backward (wgpu::Buffer, [seq×hidden])
1221    /// - `saved_norm_ffn`: OUTPUT — saved pre-FFN norm for backward (wgpu::Buffer, [seq×hidden])
1222    ///
1223    /// Forward one layer into an EXISTING encoder (no submit).
1224    /// Caller batches multiple layers into one encoder, submits once.
1225    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        // Pass 1: RMSNorm
1240        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        // SAVE attn_norm_out
1247        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        // Q/K/V projections
1256        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        // LoRA addmm on Q/K/V: output += (saved_input @ A) @ B * scale
1288        // Must happen BEFORE attention consumes Q/K/V buffers.
1289        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        // PMAT-509: Apply QKV biases (required for Qwen2)
1335        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        // PMAT-509: Apply RoPE to Q and K before attention.
1346        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        // Attention — wgpu handles execution ordering within the encoder.
1350        self.encode_attention(encoder, seq_len);
1351
1352        // SAVE attn_output
1353        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        // O projection
1362        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        // Residual
1374        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        // FFN RMSNorm
1383        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        // SAVE ffn_norm_out
1390        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        // Gate + Up
1399        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        // SiLU
1421        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        // SAVE silu_gate_output
1430        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        // Down projection
1439        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        // Residual
1451        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    /// Run one layer with per-operation GPU timing (submit+poll between each op group).
1463    /// Contract: forward-pass-perf-v1 / bottleneck_identified
1464    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        // PMAT-509: QKV biases + RoPE before attention
1587        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    /// Allocate saved activations for one layer.
1720    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    /// Forward one layer with its own encoder + submit (original API, kept for compat).
1744    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    /// Forward ALL layers in one encoder submit. 28 layers → 1 GPU sync.
1757    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        // ONE submit for all 28 layers — eliminates 27 GPU sync barriers
1773        self.queue.submit(Some(encoder.finish()));
1774        Ok(all_saved)
1775    }
1776    // --- Encode helpers (add compute passes to an existing encoder) ---
1777
1778    /// Encode causal multi-head attention on GPU.
1779    /// Q: [seq_len, num_heads * head_dim], K/V: [seq_len, num_kv_heads * head_dim]
1780    /// Output written to q_buf (reused as attn output).
1781    /// PMAT-509: Add broadcast bias to a [seq_len, dim] buffer.
1782    /// bias has shape [dim], applied to each of seq_len rows.
1783    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        // Create a full-size bias buffer by repeating the bias per position
1792        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        // Use existing residual: out = buf + bias_buf (into a temp, then copy back)
1805        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    /// PMAT-509: Encode batch RoPE for all positions in a sequence.
1817    /// Applies position-dependent rotation to Q or K buffer in-place.
1818    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(&params_buf, 0, bytemuck::bytes_of(&params));
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(&params);
1862        let _q_dim = self.num_heads * self.head_dim;
1863
1864        // Attention reads Q and writes to attn_out_buf.
1865        // Then O projection reads attn_out_buf → writes to another buffer.
1866        // We can safely write to norm_buf here since it's not read during attention.
1867        // After attention, we'll copy norm_buf → q_buf for the O projection to read.
1868        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        // One workgroup per (head, position)
1886        pass.dispatch_workgroups(self.num_heads, seq_len, 1);
1887    }
1888
1889    /// Encode LoRA addmm: output += (input @ A) @ B * scale
1890    ///
1891    /// KAIZEN: replaced fused shader (0.11 GFLOPS) with two tiled GEMM dispatches (1000+ GFLOPS).
1892    /// Step 1: temp = input @ A  [seq, rank] via tiled GEMM
1893    /// Step 2: output += scale * (temp @ B) [seq, out_dim] via tiled GEMM with alpha=scale
1894    ///
1895    /// The second GEMM uses alpha=scale in the tiled GEMM shader (C = alpha * A @ B).
1896    /// But we need ADD (+=), not overwrite (=). We use a temp buffer for the delta,
1897    /// then add to output via an elementwise shader.
1898    #[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        // Step 1: temp[seq, rank] = input[seq, in_dim] @ A[in_dim, rank]
1915        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        // Step 2: delta[seq, out_dim] = scale * temp[seq, rank] @ B[rank, out_dim]
1925        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        // Step 3: output += delta (elementwise add, via temp to avoid aliasing)
1935        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    /// Encode tiled GEMM: C = alpha * A[M,K] @ B[K,N]. Uses CUTLASS-style 64×64 tiles.
1946    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(&params);
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(&params);
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        // Dispatch: (1, num_rows, 1). Each workgroup processes one row via wg_id.y.
1996        // For inference (M=1): dispatch (1,1,1). For training (M=seq_len): dispatch (1,seq_len,1).
1997        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        // C-WGPU-Q4K-001: Try Q4K GEMV first for M=1 decode (7x less VRAM)
2016        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, // Skip missing weights silently
2023        };
2024        // PMAT-346: GEMV and matmul have different uniform struct layouts.
2025        // GEMV: Params { n (output dim), k (input dim), _, _ }
2026        // Matmul: Dimensions { M, K, N, _ }
2027        // Tiled GEMM: Dimensions { M, K, N, alpha_bits }
2028        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(&params);
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            // PMAT-327: GEMV for M=1 — cooperative K-reduction, N workgroups
2043            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            // CUTLASS-style tiled GEMM for M>=4 (training batch, prefill)
2048            // 64×64 tiles, 4×4 thread micro-tiles, 10-30x faster than naive
2049            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            // Naive 16×16 GEMM for small M (2-3)
2054            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    /// C-WGPU-Q4K-001: Encode Q4K GEMV — reads raw Q4K weight bytes, dequantizes on-the-fly.
2061    /// Falls back to F32 GEMV if no Q4K weight found for this layer.
2062    /// Returns true if Q4K path was used.
2063    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(&params);
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(&params);
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(&params);
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    /// #4056: the shader reads eps from `params.y`; it must carry the configured
2191    /// value bit-exactly, not the old hardcoded 1e-6.
2192    #[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    /// #4056: the shader now reads `bitcast<f32>(params.y)`; naga must accept it.
2202    /// Skips (returns) when no GPU adapter is present.
2203    #[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}