spark_model/layers/ngram_embed/
embed.rs1use anyhow::{Context, Result};
7use spark_runtime::gpu::{DevicePtr, GpuBackend, KernelHandle};
8
9use super::ids::ngram_ids;
10use super::{NgramDims, NgramTable};
11use crate::weight_map::DenseWeight;
12
13pub struct NgramEmbedding {
20 pub dims: NgramDims,
21 pub word: DenseWeight,
23 pub tables: Vec<NgramTable>,
26 pub projs: Vec<DenseWeight>,
28
29 pub(super) batched_embed_k: KernelHandle,
30 pub(super) batched_embed_fp8_k: KernelHandle,
31 pub(super) gemm_k: KernelHandle,
32 pub(super) scaled_add_k: KernelHandle,
33
34 pub(super) ids_dev: DevicePtr,
38 pub(super) gather_buf: DevicePtr,
39 pub(super) proj_buf: DevicePtr,
40 pub(super) max_tokens: usize,
41}
42
43impl NgramEmbedding {
44 pub fn new(
45 dims: NgramDims,
46 word: DenseWeight,
47 tables: Vec<NgramTable>,
48 projs: Vec<DenseWeight>,
49 max_tokens: usize,
50 gpu: &dyn GpuBackend,
51 ) -> Result<Self> {
52 anyhow::ensure!(tables.len() == dims.num_tables(), "table count");
53 anyhow::ensure!(projs.len() == dims.num_tables(), "proj count");
54 let td = dims.table_dim();
55 Ok(Self {
56 dims,
57 word,
58 tables,
59 projs,
60 batched_embed_k: gpu.kernel("embed_from_argmax", "batched_embed")?,
61 batched_embed_fp8_k: gpu.kernel("embed_from_argmax", "batched_embed_fp8")?,
62 gemm_k: gpu.kernel("gemm", "dense_gemm_bf16_pipelined")?,
63 scaled_add_k: gpu.kernel("residual_add", "bf16_scaled_add")?,
64 ids_dev: gpu.alloc(max_tokens * 4)?,
65 gather_buf: gpu.alloc(max_tokens * td * 2)?,
66 proj_buf: gpu.alloc(max_tokens * dims.hidden_size * 2)?,
67 max_tokens,
68 })
69 }
70
71 pub fn embed(
76 &mut self,
77 ctx_tokens: &[u32],
78 seq_len: usize,
79 out: DevicePtr,
80 gpu: &dyn GpuBackend,
81 stream: u64,
82 ) -> Result<()> {
83 use crate::layers::ops;
84 anyhow::ensure!(
85 seq_len <= self.max_tokens,
86 "ngram embed: seq_len over staging"
87 );
88 anyhow::ensure!(seq_len <= ctx_tokens.len(), "ngram embed: seq_len over ctx");
89 let h = self.dims.hidden_size;
90 let td = self.dims.table_dim();
91 let inv_scale = 1.0f32 / (1 + self.dims.num_tables()) as f32;
92
93 gpu.memset(out, 0, seq_len * h * 2)?;
96
97 let new_tokens = &ctx_tokens[ctx_tokens.len() - seq_len..];
99 let ids_bytes: Vec<u8> = new_tokens.iter().flat_map(|t| t.to_le_bytes()).collect();
100 gpu.copy_h2d_async(&ids_bytes, self.ids_dev, stream)?;
101 ops::batched_embed(
102 gpu,
103 self.batched_embed_k,
104 self.ids_dev,
105 self.word.weight,
106 self.proj_buf,
107 seq_len as u32,
108 h as u32,
109 stream,
110 )?;
111 ops::scaled_add(
112 gpu,
113 self.scaled_add_k,
114 out,
115 self.proj_buf,
116 inv_scale,
117 (seq_len * h) as u32,
118 stream,
119 )?;
120
121 let all_ids = ngram_ids(&self.dims, ctx_tokens);
123 for (index, ids) in all_ids.iter().enumerate() {
124 let tail = &ids[ids.len() - seq_len..];
125 let id_bytes: Vec<u8> = tail
126 .iter()
127 .map(|&v| u32::try_from(v).context("ngram id exceeds u32"))
128 .collect::<Result<Vec<u32>>>()?
129 .iter()
130 .flat_map(|v| v.to_le_bytes())
131 .collect();
132 gpu.copy_h2d_async(&id_bytes, self.ids_dev, stream)?;
133 #[cfg(feature = "cuda")]
137 if let NgramTable::Cached(cache) = &mut self.tables[index] {
138 let mut slots: Vec<u32> = Vec::with_capacity(seq_len);
139 cache.resolve(tail, &mut slots)?;
140 let slot_bytes: Vec<u8> =
141 slots.iter().flat_map(|v: &u32| v.to_le_bytes()).collect();
142 gpu.copy_h2d_async(&slot_bytes, self.ids_dev, stream)?;
143 let table = DevicePtr(cache.table_dev_va()?);
144 match cache.scale_dev_va()? {
145 Some(sc) => ops::batched_embed_fp8(
146 gpu,
147 self.batched_embed_fp8_k,
148 self.ids_dev,
149 table,
150 DevicePtr(sc),
151 self.gather_buf,
152 seq_len as u32,
153 td as u32,
154 stream,
155 )?,
156 None => ops::batched_embed(
157 gpu,
158 self.batched_embed_k,
159 self.ids_dev,
160 table,
161 self.gather_buf,
162 seq_len as u32,
163 td as u32,
164 stream,
165 )?,
166 }
167 cache.end_batch();
168 ops::dense_gemm_bf16_pipelined(
169 gpu,
170 self.gemm_k,
171 self.gather_buf,
172 &self.projs[index],
173 self.proj_buf,
174 seq_len as u32,
175 h as u32,
176 td as u32,
177 stream,
178 )?;
179 ops::scaled_add(
180 gpu,
181 self.scaled_add_k,
182 out,
183 self.proj_buf,
184 inv_scale,
185 (seq_len * h) as u32,
186 stream,
187 )?;
188 continue;
189 }
190 match &self.tables[index] {
191 NgramTable::Bf16(w) => ops::batched_embed(
192 gpu,
193 self.batched_embed_k,
194 self.ids_dev,
195 w.weight,
196 self.gather_buf,
197 seq_len as u32,
198 td as u32,
199 stream,
200 )?,
201 NgramTable::Fp8(w) => ops::batched_embed_fp8(
202 gpu,
203 self.batched_embed_fp8_k,
204 self.ids_dev,
205 w.weight,
206 w.row_scale,
207 self.gather_buf,
208 seq_len as u32,
209 td as u32,
210 stream,
211 )?,
212 #[cfg(feature = "cuda")]
213 NgramTable::Cached(_) => unreachable!("resolved above"),
214 }
215 ops::dense_gemm_bf16_pipelined(
216 gpu,
217 self.gemm_k,
218 self.gather_buf,
219 &self.projs[index],
220 self.proj_buf,
221 seq_len as u32,
222 h as u32,
223 td as u32,
224 stream,
225 )?;
226 ops::scaled_add(
227 gpu,
228 self.scaled_add_k,
229 out,
230 self.proj_buf,
231 inv_scale,
232 (seq_len * h) as u32,
233 stream,
234 )?;
235 }
236 Ok(())
237 }
238}