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}
112
113// WGSL shader source for RMSNorm (multi-row via workgroup_id.y)
114// Dispatch: (1, seq_len, 1) — one workgroup per row.
115const 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
163// WGSL shader for SiLU(gate) * up
164const 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
180// WGSL shader for residual add: output = a + b
181const 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
195// Batch RoPE shader — applies RoPE to all positions in a sequence at once.
196// PMAT-509: Training forward path was missing RoPE entirely, causing loss > random.
197// Input: qk[seq_len * num_heads * head_dim], applies position-dependent rotation.
198const 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
245// RoPE shader (NeoX-style interleaved) — single position (inference)
246const 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    /// Get the shader sources for external inspection/testing
282    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    /// PMAT-325: Create a new WGSL forward pass context.
296    ///
297    /// Compiles all shader pipelines and allocates persistent intermediate buffers.
298    /// Call once at model init. All GPU resources persist until dropped.
299    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        // Compile shaders
312        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        // Bind group layouts
334        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        // Pipelines
354        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        // CUTLASS-style tiled GEMM for M>=4 (training batch, prefill)
374        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        // Causal attention pipeline for training
384        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),  // Q
394                bgl_storage(1, true),  // K
395                bgl_storage(2, true),  // V
396                bgl_storage(3, false), // output
397                bgl_uniform(4),        // params
398            ],
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        // PMAT-327: GEMV pipeline — same bind group layout as matmul but cooperative reduction
415        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        // C-WGPU-Q4K-001: Q4K GEMV — dequantize on-the-fly, no F32 weight buffer
422        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        // RoPE has a 2-binding layout (in-place + uniform)
433        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        // PMAT-509: Batch RoPE for training (all positions at once)
454        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        // Allocate persistent intermediate buffers
479        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        // Buffer sizes: max_seq × dim for training, or 1 × dim for inference.
491        // Training calls forward_layer_training with seq_len > 1.
492        // Allocate for max_seq=2048 to support both.
493        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    /// Upload a weight matrix (call once per layer at init).
555    /// PMAT-342: Bias weights (name contains "bias") are stored CPU-side.
556    pub fn upload_weight(&mut self, name: &str, data: &[f32]) {
557        if name.contains("bias") {
558            // Biases are small, keep on CPU for easy access in attention
559            self.cpu_biases.insert(name.to_string(), data.to_vec());
560            return;
561        }
562        // Skip weights that exceed the device's max buffer binding size (e.g., lm_head > 2 GB)
563        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    /// GH-560: Upload raw Q4K weight bytes for fused dequant+GEMV on GPU.
584    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    /// GH-560: Initialize per-layer KV cache buffers on GPU.
595    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    /// Number of uploaded weight buffers.
621    pub fn weight_count(&self) -> usize {
622        self.weight_buffers.len()
623    }
624
625    /// Access a dequantized weight buffer by name (e.g. "layer.0.down_proj").
626    /// Used by backward pass for gradient propagation through frozen base weights.
627    pub fn weight_buffer(&self, name: &str) -> Option<&wgpu::Buffer> {
628        self.weight_buffers.get(name)
629    }
630
631    /// Reference to the wgpu device.
632    pub fn device_ref(&self) -> &wgpu::Device {
633        &self.device
634    }
635
636    /// Reference to the wgpu queue.
637    pub fn queue_ref(&self) -> &wgpu::Queue {
638        &self.queue
639    }
640
641    /// Reference to the hidden state buffer (for writing input).
642    pub fn hidden_buffer(&self) -> &wgpu::Buffer {
643        &self.hidden_buf
644    }
645
646    /// Reference to Q buffer (for LoRA addmm after Q projection).
647    pub fn q_buffer(&self) -> &wgpu::Buffer {
648        &self.q_buf
649    }
650
651    /// Reference to K buffer.
652    pub fn k_buffer(&self) -> &wgpu::Buffer {
653        &self.k_buf
654    }
655
656    /// Reference to V buffer.
657    pub fn v_buffer(&self) -> &wgpu::Buffer {
658        &self.v_buf
659    }
660
661    /// Elementwise add: output = a + b. Dispatches residual add shader.
662    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    /// Apply RMSNorm on GPU: normed = rmsnorm(hidden_buf, weight) → output_buf.
675    /// Contract: gpu-output-norm-v1 / gpu_resident — hidden state never leaves GPU.
676    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    /// Download hidden state from GPU.
683    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    /// Total VRAM used by all buffers (bytes).
713    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  // hidden, attn_out, ffn_out, norm
716            + (self.num_heads as usize * self.head_dim as usize * 4) // q
717            + (self.num_kv_heads as usize * self.head_dim as usize * 4) * 2 // k, v
718            + (self.intermediate_dim as usize * 4) * 2; // gate, up
719        weight_bytes + intermediate_bytes
720    }
721
722    /// PMAT-336: Full model forward — embedding + all layers + output norm + LM head.
723    ///
724    /// Returns logits [vocab_size] for the given token at the given position.
725    /// Embedding lookup and final LM head are CPU-side (not yet GPU-accelerated).
726    /// PMAT-344: Added kv_caches for multi-token context
727    #[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        // 1. Embedding lookup (CPU)
743        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        // 2. Transformer layers (GPU via forward_layer with KV cache)
754        // Initialize KV caches if empty
755        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        // 3. Output RMSNorm (CPU — small, not worth GPU dispatch)
765        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        // 4. LM head — CPU matmul
771        // PMAT-346: GPU tiled GEMM expects weight in [K,N] layout but lm_head is [N,K].
772        // CPU path reads weight[v * hd + j] which matches the [vocab, hidden] layout.
773        // GPU LM head via GEMV is blocked by vocab > 65535 dispatch limit.
774        // Deferred (PMAT-751): add upload_weight_transposed() for GPU-accelerated LM head.
775        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    /// PMAT-325: Execute one transformer layer — 14 passes, 1 submit, 1 readback.
788    ///
789    /// Input: hidden state [hidden_dim] on CPU.
790    /// Output: updated hidden state [hidden_dim] on CPU.
791    /// All intermediate computation stays GPU-resident.
792    /// PMAT-344: KV cache parameters for multi-token context
793    pub fn forward_layer(
794        &self,
795        hidden: &mut [f32],
796        layer_prefix: &str,
797        _position: usize,
798        kv_cache_k: &mut Vec<f32>, // accumulated K: [seq_len * kv_dim]
799        kv_cache_v: &mut Vec<f32>, // accumulated V: [seq_len * kv_dim]
800    ) -> Result<(), String> {
801        let hd = self.hidden_dim;
802
803        // Upload hidden state
804        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        // Pass 1: RMSNorm(hidden → norm_buf)
809        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        // Passes 2-4: Q/K/V projections (norm_buf × W → q/k/v_buf)
816        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        // PMAT-342: Submit Q/K/V projections, readback, do attention on CPU
851        // GPU handles the heavy matmuls; CPU handles attention (small at M=1)
852        let q_bytes = (q_dim * 4) as u64;
853        let kv_bytes = (kv_dim * 4) as u64;
854
855        // Readback Q/K/V from GPU
856        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        // Readback Q
880        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        // Readback K
895        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        // Readback V
910        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        // PMAT-342: Add QKV biases (required for Qwen2)
925        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        // PMAT-343: Apply RoPE (NeoX-style interleaved) to Q and K
942        let head_dim = self.head_dim as usize;
943        let position = _position; // Use the position parameter
944        let rope_theta = 1_000_000.0f64; // Qwen2 rope_theta
945
946        // RoPE on Q (num_heads × head_dim)
947        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        // RoPE on K (num_kv_heads × head_dim)
963        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        // PMAT-344: Append K,V to cache and compute full attention
979        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        // Scaled dot-product attention with GQA
989        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            // Compute attention scores: Q[h] · K[kv_h, :seq_len]^T / sqrt(d)
999            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            // Weighted sum of V
1009            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        // Upload attention output back to GPU for O projection
1020        self.queue.write_buffer(&self.q_buf, 0, bytemuck::cast_slice(&attn_out));
1021
1022        // New encoder for remaining passes
1023        let mut encoder = self.device.create_command_encoder(&Default::default());
1024
1025        // Pass 7: O projection (attn_out × W_o → attn_out_buf)
1026        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        // Pass 8: Residual(hidden + attn_out → hidden)
1038        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        // Pass 9: FFN RMSNorm(ffn_out → norm_buf)
1047        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        // Passes 10-11: Gate + Up projections
1054        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        // Pass 12: SiLU(gate) × up → ffn_silu_buf [intermediate_dim]
1077        // BUG FIX: was writing to attn_out_buf (hidden_dim=3584) but needs intermediate_dim=18944.
1078        // attn_out_buf is only hidden_dim — wgpu robustness silently drops OOB writes,
1079        // then down_proj reads zeros past hidden_dim → 81% of FFN truncated → garbage output.
1080        // Cannot alias gate/up buffers (WGSL read/write aliasing UB), so use dedicated buffer.
1081        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        // Pass 13: Down projection (reads ffn_silu_buf [intermediate_dim] → norm_buf [hidden_dim])
1090        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        // Pass 14: Residual(ffn_out + down → hidden)
1102        self.encode_residual(&mut encoder, &self.ffn_out_buf, &self.norm_buf, &self.hidden_buf, hd);
1103
1104        // Single readback
1105        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        // Readback
1109        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    /// Softmaxed attention scores for one head over `seq_len` cached steps.
1128    ///
1129    /// `kv_offset` is this head's offset inside each cached step, and each
1130    /// cached step is `kv_dim` floats wide. Head dimension is `q_head.len()`.
1131    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        // Softmax
1150        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    /// Weighted sum of V for one head, written into `out_head`.
1165    ///
1166    /// `kv_offset` and `kv_dim` are as in `attention_scores`.
1167    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    /// Training forward pass for a single transformer layer.
1185    ///
1186    /// Unlike `forward_layer` (M=1 decode), this processes the full sequence
1187    /// at once (M=seq_len) and keeps everything on GPU. No CPU readback.
1188    ///
1189    /// Saves `norm_output` (pre-projection activations) for backward pass.
1190    ///
1191    /// # Arguments
1192    /// - `seq_len`: number of tokens in the sequence
1193    /// - `layer_prefix`: e.g. "model.layers.0"
1194    /// - `saved_norm_attn`: OUTPUT — saved pre-attention norm for backward (wgpu::Buffer, [seq×hidden])
1195    /// - `saved_norm_ffn`: OUTPUT — saved pre-FFN norm for backward (wgpu::Buffer, [seq×hidden])
1196    ///
1197    /// Forward one layer into an EXISTING encoder (no submit).
1198    /// Caller batches multiple layers into one encoder, submits once.
1199    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        // Pass 1: RMSNorm
1214        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        // SAVE attn_norm_out
1221        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        // Q/K/V projections
1230        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        // LoRA addmm on Q/K/V: output += (saved_input @ A) @ B * scale
1262        // Must happen BEFORE attention consumes Q/K/V buffers.
1263        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        // PMAT-509: Apply QKV biases (required for Qwen2)
1309        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        // PMAT-509: Apply RoPE to Q and K before attention.
1320        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        // Attention — wgpu handles execution ordering within the encoder.
1324        self.encode_attention(encoder, seq_len);
1325
1326        // SAVE attn_output
1327        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        // O projection
1336        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        // Residual
1348        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        // FFN RMSNorm
1357        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        // SAVE ffn_norm_out
1364        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        // Gate + Up
1373        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        // SiLU
1395        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        // SAVE silu_gate_output
1404        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        // Down projection
1413        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        // Residual
1425        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    /// Run one layer with per-operation GPU timing (submit+poll between each op group).
1437    /// Contract: forward-pass-perf-v1 / bottleneck_identified
1438    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        // PMAT-509: QKV biases + RoPE before attention
1561        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    /// Allocate saved activations for one layer.
1694    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    /// Forward one layer with its own encoder + submit (original API, kept for compat).
1718    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    /// Forward ALL layers in one encoder submit. 28 layers → 1 GPU sync.
1731    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        // ONE submit for all 28 layers — eliminates 27 GPU sync barriers
1747        self.queue.submit(Some(encoder.finish()));
1748        Ok(all_saved)
1749    }
1750    // --- Encode helpers (add compute passes to an existing encoder) ---
1751
1752    /// Encode causal multi-head attention on GPU.
1753    /// Q: [seq_len, num_heads * head_dim], K/V: [seq_len, num_kv_heads * head_dim]
1754    /// Output written to q_buf (reused as attn output).
1755    /// PMAT-509: Add broadcast bias to a [seq_len, dim] buffer.
1756    /// bias has shape [dim], applied to each of seq_len rows.
1757    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        // Create a full-size bias buffer by repeating the bias per position
1766        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        // Use existing residual: out = buf + bias_buf (into a temp, then copy back)
1779        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    /// PMAT-509: Encode batch RoPE for all positions in a sequence.
1791    /// Applies position-dependent rotation to Q or K buffer in-place.
1792    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(&params_buf, 0, bytemuck::bytes_of(&params));
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(&params);
1836        let _q_dim = self.num_heads * self.head_dim;
1837
1838        // Attention reads Q and writes to attn_out_buf.
1839        // Then O projection reads attn_out_buf → writes to another buffer.
1840        // We can safely write to norm_buf here since it's not read during attention.
1841        // After attention, we'll copy norm_buf → q_buf for the O projection to read.
1842        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        // One workgroup per (head, position)
1860        pass.dispatch_workgroups(self.num_heads, seq_len, 1);
1861    }
1862
1863    /// Encode LoRA addmm: output += (input @ A) @ B * scale
1864    ///
1865    /// KAIZEN: replaced fused shader (0.11 GFLOPS) with two tiled GEMM dispatches (1000+ GFLOPS).
1866    /// Step 1: temp = input @ A  [seq, rank] via tiled GEMM
1867    /// Step 2: output += scale * (temp @ B) [seq, out_dim] via tiled GEMM with alpha=scale
1868    ///
1869    /// The second GEMM uses alpha=scale in the tiled GEMM shader (C = alpha * A @ B).
1870    /// But we need ADD (+=), not overwrite (=). We use a temp buffer for the delta,
1871    /// then add to output via an elementwise shader.
1872    #[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        // Step 1: temp[seq, rank] = input[seq, in_dim] @ A[in_dim, rank]
1889        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        // Step 2: delta[seq, out_dim] = scale * temp[seq, rank] @ B[rank, out_dim]
1899        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        // Step 3: output += delta (elementwise add, via temp to avoid aliasing)
1909        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    /// Encode tiled GEMM: C = alpha * A[M,K] @ B[K,N]. Uses CUTLASS-style 64×64 tiles.
1920    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(&params);
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(&params);
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        // Dispatch: (1, num_rows, 1). Each workgroup processes one row via wg_id.y.
1970        // For inference (M=1): dispatch (1,1,1). For training (M=seq_len): dispatch (1,seq_len,1).
1971        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        // C-WGPU-Q4K-001: Try Q4K GEMV first for M=1 decode (7x less VRAM)
1990        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, // Skip missing weights silently
1997        };
1998        // PMAT-346: GEMV and matmul have different uniform struct layouts.
1999        // GEMV: Params { n (output dim), k (input dim), _, _ }
2000        // Matmul: Dimensions { M, K, N, _ }
2001        // Tiled GEMM: Dimensions { M, K, N, alpha_bits }
2002        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(&params);
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            // PMAT-327: GEMV for M=1 — cooperative K-reduction, N workgroups
2017            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            // CUTLASS-style tiled GEMM for M>=4 (training batch, prefill)
2022            // 64×64 tiles, 4×4 thread micro-tiles, 10-30x faster than naive
2023            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            // Naive 16×16 GEMM for small M (2-3)
2028            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    /// C-WGPU-Q4K-001: Encode Q4K GEMV — reads raw Q4K weight bytes, dequantizes on-the-fly.
2035    /// Falls back to F32 GEMV if no Q4K weight found for this layer.
2036    /// Returns true if Q4K path was used.
2037    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(&params);
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(&params);
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(&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: 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}