1use super::*;
4
5impl QTensorEngine {
6 #[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 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(¶ms_buf, 0, bytemuck::bytes_of(¶ms));
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 #[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 #[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 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 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 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(¶ms_buf, 0, bytemuck::bytes_of(&gemm_params));
291
292 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 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 #[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 #[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}