Skip to main content

qualia_core_db/gguf_bridge/
embedding.rs

1//! Quantized token embedding + fused transformer block dispatch
2//! Split from gguf_bridge/mod.rs (structural refactor; no behaviour change).
3use super::*;
4
5impl QTensorEngine {
6    /// Upload raw quantized embedding bytes to the GPU and matmul without CPU dequant.
7    /// Returns `None` when the GGML type has no WGSL kernel (caller uses CPU fallback).
8    #[cfg(not(target_arch = "wasm32"))]
9    pub fn dispatch_quantized_token_embedding(
10        &self,
11        raw_embd: &[u8],
12        ggml_type: u32,
13        n_embd: u32,
14        weight_tensor: &QTensor,
15    ) -> Option<Vec<f32>> {
16        if ggml_type != crate::ggml_quants::GGML_TYPE_Q6_K || raw_embd.is_empty() || n_embd == 0 {
17            return None;
18        }
19
20        let n_output = weight_tensor
21            .shape
22            .first()
23            .copied()
24            .unwrap_or(n_embd as usize) as u32;
25        let n_embd_u = n_embd;
26        let weights_elems = (n_output as usize).saturating_mul(n_embd as usize);
27
28        let params = EmbeddingGpuParams {
29            n_embd: n_embd_u,
30            ggml_type,
31            n_output,
32            raw_byte_len: raw_embd.len() as u32,
33        };
34
35        // WGSL storage uses u32 words; pad mmap slice to 4-byte alignment.
36        let word_bytes = raw_embd.len().div_ceil(4) * 4;
37        let embd_buf = self.gpu_device().create_buffer(&wgpu::BufferDescriptor {
38            label: Some("QuantizedEmbeddingBytes"),
39            size: word_bytes.max(4) as wgpu::BufferAddress,
40            usage: wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_DST,
41            mapped_at_creation: false,
42        });
43        if raw_embd.len() == word_bytes {
44            self.gpu_queue().write_buffer(&embd_buf, 0, raw_embd);
45        } else {
46            const MAX_EMB_ROW_PAD: usize = 8192;
47            if word_bytes > MAX_EMB_ROW_PAD {
48                return None;
49            }
50            let mut padded = [0u8; MAX_EMB_ROW_PAD];
51            padded[..raw_embd.len()].copy_from_slice(raw_embd);
52            self.gpu_queue()
53                .write_buffer(&embd_buf, 0, &padded[..word_bytes]);
54        }
55
56        let params_buf = self.gpu_device().create_buffer(&wgpu::BufferDescriptor {
57            label: Some("EmbeddingParams"),
58            size: std::mem::size_of::<EmbeddingGpuParams>() as wgpu::BufferAddress,
59            usage: wgpu::BufferUsages::UNIFORM | wgpu::BufferUsages::COPY_DST,
60            mapped_at_creation: false,
61        });
62        self.gpu_queue()
63            .write_buffer(&params_buf, 0, bytemuck::bytes_of(&params));
64
65        let weights_size = (weights_elems * 4).max(4) as wgpu::BufferAddress;
66        let weights_buf = self.gpu_device().create_buffer(&wgpu::BufferDescriptor {
67            label: Some("EmbeddingWeights"),
68            size: weights_size,
69            usage: wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_DST,
70            mapped_at_creation: false,
71        });
72        if let Some(mmap) = &self.gguf_mmap {
73            let offset = (self.tensor_data_offset + weight_tensor.byte_offset) as usize;
74            let end = (offset + weights_elems * 4).min(mmap.len());
75            if end > offset {
76                self.gpu_queue()
77                    .write_buffer(&weights_buf, 0, &mmap[offset..end]);
78            }
79        }
80
81        let output_size = (n_output as usize * 4).max(4) as wgpu::BufferAddress;
82        let output_buf = self.gpu_device().create_buffer(&wgpu::BufferDescriptor {
83            label: Some("EmbeddingOutput"),
84            size: output_size,
85            usage: wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_SRC,
86            mapped_at_creation: false,
87        });
88
89        #[cfg(not(target_arch = "wasm32"))]
90        let bind_layout = self.embedding_bind_layout.clone();
91        #[cfg(target_arch = "wasm32")]
92        let bind_layout = self.embedding_pipeline.get_bind_group_layout(0);
93        let bind_group = self
94            .gpu_device()
95            .create_bind_group(&wgpu::BindGroupDescriptor {
96                label: Some("QuantizedEmbeddingBindGroup"),
97                layout: &bind_layout,
98                entries: &[
99                    wgpu::BindGroupEntry {
100                        binding: 0,
101                        resource: embd_buf.as_entire_binding(),
102                    },
103                    wgpu::BindGroupEntry {
104                        binding: 1,
105                        resource: params_buf.as_entire_binding(),
106                    },
107                    wgpu::BindGroupEntry {
108                        binding: 2,
109                        resource: weights_buf.as_entire_binding(),
110                    },
111                    wgpu::BindGroupEntry {
112                        binding: 3,
113                        resource: output_buf.as_entire_binding(),
114                    },
115                ],
116            });
117
118        let mut encoder = self
119            .device()
120            .create_command_encoder(&wgpu::CommandEncoderDescriptor {
121                label: Some("QuantizedEmbeddingEncoder"),
122            });
123        {
124            let mut cpass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
125                label: Some("QuantizedEmbeddingPass"),
126                timestamp_writes: crate::llm_gpu_profiler::pass_writes_both(),
127            });
128            cpass.set_pipeline(&self.embedding_pipeline);
129            cpass.set_bind_group(0, &bind_group, &[]);
130            cpass.dispatch_workgroups((n_output + 63) / 64, 1, 1);
131        }
132
133        let staging_buf = self.gpu_device().create_buffer(&wgpu::BufferDescriptor {
134            label: Some("EmbeddingStaging"),
135            size: output_size,
136            usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
137            mapped_at_creation: false,
138        });
139        encoder.copy_buffer_to_buffer(&output_buf, 0, &staging_buf, 0, output_size);
140        crate::llm_gpu_profiler::resolve(&mut encoder);
141        self.gpu_queue().submit(Some(encoder.finish()));
142        crate::llm_gpu_profiler::accumulate(crate::llm_gpu_profiler::Phase::Embedding);
143
144        let buffer_slice = staging_buf.slice(..);
145        let (sender, receiver) = futures_channel::oneshot::channel();
146        buffer_slice.map_async(wgpu::MapMode::Read, move |v| {
147            let _ = sender.send(v);
148        });
149        self.poll_wait();
150
151        #[cfg(not(target_arch = "wasm32"))]
152        {
153            let handle = tokio::runtime::Handle::try_current().unwrap_or_else(|_| {
154                let rt = Box::leak(Box::new(tokio::runtime::Runtime::new().unwrap()));
155                rt.handle().clone()
156            });
157            if handle.block_on(receiver).ok()?.is_err() {
158                return None;
159            }
160        }
161        #[cfg(target_arch = "wasm32")]
162        {
163            return None;
164        }
165
166        let data = buffer_slice
167            .get_mapped_range()
168            .expect("wgpu buffer map_range failed");
169        let result: Vec<f32> = bytemuck::cast_slice(&data).to_vec();
170        drop(data);
171        staging_buf.unmap();
172
173        crate::telemetry::SIEVE_OPS_COUNT
174            .fetch_add(weights_elems, std::sync::atomic::Ordering::Relaxed);
175        Some(result)
176    }
177
178    #[cfg(not(target_arch = "wasm32"))]
179    pub fn dispatch_fused_transformer_block(
180        &self,
181        tensor: &QTensor,
182        input_activations: &[f32],
183    ) -> Vec<f32> {
184        let rows = tensor.shape.get(0).copied().unwrap_or(4096);
185        let cols = tensor.shape.get(1).copied().unwrap_or(4096);
186
187        // ── DirectML path (Windows) ───────────────────────────────────────────
188        #[cfg(target_os = "windows")]
189        if let Some(dml) = &self.dml {
190            if let Some(mmap) = &self.gguf_mmap {
191                let offset = self.tensor_data_offset + tensor.byte_offset;
192                let q4_bytes_needed = (rows * cols / crate::directml_bridge::Q4_K_BLOCK_SIZE)
193                    * crate::directml_bridge::Q4_K_BLOCK_BYTES;
194                if (offset as usize + q4_bytes_needed) <= mmap.len() {
195                    let q4_slice = &mmap[offset as usize..offset as usize + q4_bytes_needed];
196                    let weights_f32 =
197                        crate::directml_bridge::dequantize_q4_k_tensor(q4_slice, rows * cols);
198                    let op = crate::directml_bridge::DmlGemmOp {
199                        m: input_activations.len() as u32 / cols as u32,
200                        k: cols as u32,
201                        n: rows as u32,
202                    };
203                    if let Ok(result) = op.execute(dml, input_activations, &weights_f32) {
204                        crate::telemetry::SIEVE_OPS_COUNT
205                            .fetch_add(rows * cols, std::sync::atomic::Ordering::Relaxed);
206                        return result;
207                    }
208                }
209            }
210        }
211
212        // ── Accelerate BLAS path (macOS / Apple Silicon AMX) ─────────────────────
213        #[cfg(any(target_os = "macos", target_os = "ios"))]
214        if let Some(mmap) = &self.gguf_mmap {
215            let offset = (self.tensor_data_offset + tensor.byte_offset) as usize;
216            let q4_bytes_needed = (rows * cols / crate::metal_bridge::Q4_K_BLOCK_SIZE)
217                * crate::metal_bridge::Q4_K_BLOCK_BYTES;
218            if offset + q4_bytes_needed <= mmap.len() {
219                let q4_slice = &mmap[offset..offset + q4_bytes_needed];
220                let weights_f32 =
221                    crate::metal_bridge::dequantize_q4_k_tensor(q4_slice, rows * cols);
222                let input_rows = (input_activations.len() / cols).max(1);
223                let result = crate::metal_bridge::accelerate_sgemm(
224                    input_rows,
225                    cols,
226                    rows,
227                    input_activations,
228                    &weights_f32,
229                );
230                crate::telemetry::SIEVE_OPS_COUNT
231                    .fetch_add(rows * cols, std::sync::atomic::Ordering::Relaxed);
232                return result;
233            }
234        }
235
236        // ── wgpu / WGSL fallback (all platforms — Vulkan on Linux/NVIDIA,
237        //    Metal on macOS when mmap not loaded, D3D12 on Windows fallback) ──
238        let input_bytes = bytemuck::cast_slice(input_activations);
239        let input_buf = self.gpu_device().create_buffer(&wgpu::BufferDescriptor {
240            label: Some("Input"),
241            size: input_bytes.len().max(4) as wgpu::BufferAddress,
242            usage: wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_DST,
243            mapped_at_creation: false,
244        });
245        self.gpu_queue().write_buffer(&input_buf, 0, input_bytes);
246
247        // Upload real weights from mmap when available, else use a zero buffer.
248        let weights_size = (rows * cols * 4) as wgpu::BufferAddress;
249        let weights_buf = self.gpu_device().create_buffer(&wgpu::BufferDescriptor {
250            label: Some("Weights"),
251            size: weights_size.max(4),
252            usage: wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_DST,
253            mapped_at_creation: false,
254        });
255        if let Some(mmap) = &self.gguf_mmap {
256            let offset = (self.tensor_data_offset + tensor.byte_offset) as usize;
257            let end = (offset + rows * cols * 4).min(mmap.len());
258            if end > offset {
259                let f32_bytes = &mmap[offset..end];
260                self.gpu_queue().write_buffer(&weights_buf, 0, f32_bytes);
261            }
262        }
263
264        let output_size = (rows * 4).max(4) as wgpu::BufferAddress;
265        let output_buf = self.gpu_device().create_buffer(&wgpu::BufferDescriptor {
266            label: Some("Output"),
267            size: output_size,
268            usage: wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_SRC,
269            mapped_at_creation: false,
270        });
271
272        // Upload GemmGpuParams for fused_transformer.wgsl (binding 2).
273        let gemm_params = GemmGpuParams {
274            n_in: cols as u32,
275            n_out: rows as u32,
276            weight_ggml_type: if tensor.is_quantized_q4_k { 12 } else { 14 },
277            weight_row_elems: cols as u32,
278            weight_byte_len: (rows * cols * 4) as u32,
279            n_batch: 1,
280            in_row_stride: 0,
281            out_row_stride: 0,
282        };
283        let params_buf = self.gpu_device().create_buffer(&wgpu::BufferDescriptor {
284            label: Some("TransformerParams"),
285            size: std::mem::size_of::<GemmGpuParams>() as wgpu::BufferAddress,
286            usage: wgpu::BufferUsages::UNIFORM | wgpu::BufferUsages::COPY_DST,
287            mapped_at_creation: false,
288        });
289        self.gpu_queue()
290            .write_buffer(&params_buf, 0, bytemuck::bytes_of(&gemm_params));
291
292        // The bind group MUST be built from the SAME pipeline it is dispatched
293        // under (below). Both `pipeline` (Fused Transformer) and `mock_pipeline`
294        // (Mock Fused Contraction) are created with `layout: None`, i.e. wgpu
295        // *exclusive* auto-derived layouts — a bind group from one is rejected by
296        // a dispatch that set the other ("exclusive pipelines don't match"). The
297        // real path (`gguf_mmap.is_some()`) dispatches `self.pipeline`; the mock
298        // path (no model, i.e. tests) dispatches `self.mock_pipeline`, so it must
299        // take the mock pipeline's own group-0 layout.
300        let use_mock = self.gguf_mmap.is_none();
301        #[cfg(not(target_arch = "wasm32"))]
302        let bind_group_layout = if use_mock {
303            self.mock_pipeline.get_bind_group_layout(0)
304        } else {
305            self.pipeline_bind_layout.clone()
306        };
307        #[cfg(target_arch = "wasm32")]
308        let bind_group_layout = if use_mock {
309            self.mock_pipeline.get_bind_group_layout(0)
310        } else {
311            self.pipeline.get_bind_group_layout(0)
312        };
313        let bind_group = self
314            .gpu_device()
315            .create_bind_group(&wgpu::BindGroupDescriptor {
316                label: None,
317                layout: &bind_group_layout,
318                entries: &[
319                    wgpu::BindGroupEntry {
320                        binding: 0,
321                        resource: input_buf.as_entire_binding(),
322                    },
323                    wgpu::BindGroupEntry {
324                        binding: 1,
325                        resource: weights_buf.as_entire_binding(),
326                    },
327                    wgpu::BindGroupEntry {
328                        binding: 2,
329                        resource: params_buf.as_entire_binding(),
330                    },
331                    wgpu::BindGroupEntry {
332                        binding: 3,
333                        resource: output_buf.as_entire_binding(),
334                    },
335                ],
336            });
337
338        let mut encoder = self
339            .device()
340            .create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
341        {
342            let mut cpass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
343                label: None,
344                timestamp_writes: crate::llm_gpu_profiler::pass_writes_both(),
345            });
346            // Same selector as the bind-group layout above — they must agree.
347            let pipeline = if use_mock {
348                &self.mock_pipeline
349            } else {
350                &self.pipeline
351            };
352            cpass.set_pipeline(pipeline);
353            cpass.set_bind_group(0, &bind_group, &[]);
354            cpass.dispatch_workgroups((rows as u32 + 63) / 64, 1, 1);
355        }
356
357        let staging_buf = self.gpu_device().create_buffer(&wgpu::BufferDescriptor {
358            label: Some("Staging"),
359            size: output_size,
360            usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
361            mapped_at_creation: false,
362        });
363        encoder.copy_buffer_to_buffer(&output_buf, 0, &staging_buf, 0, output_size);
364        crate::llm_gpu_profiler::resolve(&mut encoder);
365        self.gpu_queue().submit(Some(encoder.finish()));
366        crate::llm_gpu_profiler::accumulate(crate::llm_gpu_profiler::Phase::FusedBlock);
367
368        let buffer_slice = staging_buf.slice(..);
369        let (sender, receiver) = futures_channel::oneshot::channel();
370        buffer_slice.map_async(wgpu::MapMode::Read, move |v| sender.send(v).unwrap());
371        self.poll_wait();
372
373        #[cfg(not(target_arch = "wasm32"))]
374        {
375            let handle = tokio::runtime::Handle::try_current().unwrap_or_else(|_| {
376                let rt = Box::leak(Box::new(tokio::runtime::Runtime::new().unwrap()));
377                rt.handle().clone()
378            });
379            handle.block_on(receiver).unwrap().unwrap();
380        }
381
382        let data = buffer_slice
383            .get_mapped_range()
384            .expect("wgpu buffer map_range failed");
385        let result: Vec<f32> = bytemuck::cast_slice(&data).to_vec();
386        drop(data);
387        staging_buf.unmap();
388
389        crate::telemetry::SIEVE_OPS_COUNT
390            .fetch_add(rows * cols, std::sync::atomic::Ordering::Relaxed);
391        result
392    }
393
394    /// Browser inference must await WebGPU buffer mapping. The synchronous entry
395    /// point therefore reports the documented CPU-fallback sentinel.
396    #[cfg(target_arch = "wasm32")]
397    pub fn dispatch_quantized_token_embedding(
398        &self,
399        raw_embd: &[u8],
400        ggml_type: u32,
401        n_embd: u32,
402        weight_tensor: &QTensor,
403    ) -> Option<Vec<f32>> {
404        wlog(&format!(
405            "[embedding] synchronous browser dispatch unavailable (bytes={}, type={}, dim={}, tensor_offset={})",
406            raw_embd.len(),
407            ggml_type,
408            n_embd,
409            weight_tensor.byte_offset
410        ));
411        None
412    }
413
414    /// A synchronous GPU readback cannot make progress on the browser event
415    /// loop. Callers must use the async inference surface instead.
416    #[cfg(target_arch = "wasm32")]
417    pub fn dispatch_fused_transformer_block(
418        &self,
419        tensor: &QTensor,
420        input_activations: &[f32],
421    ) -> Vec<f32> {
422        panic!(
423            "synchronous browser transformer dispatch is unsupported for tensor at byte offset {} ({} activations); use inferWasmAsync/inferWasmStreaming",
424            tensor.byte_offset,
425            input_activations.len()
426        );
427    }
428}