spark_model/layers/ngram_embed/
embed.rs

1// SPDX-License-Identifier: AGPL-3.0-only
2
3//! The fused GPU embedding: base word row + the K*(N-1) hashed table
4//! lookups, composed from existing kernels.
5
6use 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
13/// GPU n-gram embedding: base word embedding fused with the K*(N-1)
14/// hashed-table lookups, composed entirely from existing kernels
15/// (`batched_embed` gathers + `dense_gemm_bf16_pipelined` projections +
16/// `bf16_scaled_add` accumulation). A fused single-kernel version is a
17/// later optimization; per-token cost here is 1+T gathers, T tiny GEMMs
18/// and T+1 scaled adds (T = 12 for LongCat-Lite) — negligible next to a
19pub struct NgramEmbedding {
20    pub dims: NgramDims,
21    /// Base word embedding `[vocab, hidden]` BF16.
22    pub word: DenseWeight,
23    /// The K*(N-1) lookup tables, index order `(ngram-2)*K + split`,
24    /// each `[table_rows(i), table_dim]` — BF16 or FP8-quantized.
25    pub tables: Vec<NgramTable>,
26    /// Per-table projections `[hidden, table_dim]` BF16 (nn.Linear layout).
27    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    /// Device staging: ids `[max_tokens]` u32 (reused per table),
35    /// gathered rows `[max_tokens, table_dim]` BF16, projected rows
36    /// `[max_tokens, hidden]` BF16.
37    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    /// Fused embedding for the LAST `seq_len` tokens of `ctx` (`ctx` =
72    /// up to n-1 cached context tokens followed by the new tokens —
73    /// exactly the reference NgramCache contract). Writes
74    /// `[seq_len, hidden]` BF16 to `out`.
75    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        // out = 0; each contribution lands via scaled_add(out += src/13),
94        // matching the reference's final 1/(1+T) over base + all tables.
95        gpu.memset(out, 0, seq_len * h * 2)?;
96
97        // Base word rows for the NEW tokens only.
98        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        // Host-side hash ids over the full ctx, one table at a time.
122        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            // NVMe-backed table: fault the rows in and REPLACE the ids with
134            // their slot indices, then gather from the arena exactly as if it
135            // were a small resident table.
136            #[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}