pub struct NgramEmbedding {
pub dims: NgramDims,
pub word: DenseWeight,
pub tables: Vec<NgramTable>,
pub projs: Vec<DenseWeight>,
/* private fields */
}Expand description
GPU n-gram embedding: base word embedding fused with the K*(N-1)
hashed-table lookups, composed entirely from existing kernels
(batched_embed gathers + dense_gemm_bf16_pipelined projections +
bf16_scaled_add accumulation). A fused single-kernel version is a
later optimization; per-token cost here is 1+T gathers, T tiny GEMMs
and T+1 scaled adds (T = 12 for LongCat-Lite) — negligible next to a
Fields§
§dims: NgramDims§word: DenseWeightBase word embedding [vocab, hidden] BF16.
tables: Vec<NgramTable>The K*(N-1) lookup tables, index order (ngram-2)*K + split,
each [table_rows(i), table_dim] — BF16 or FP8-quantized.
projs: Vec<DenseWeight>Per-table projections [hidden, table_dim] BF16 (nn.Linear layout).
Implementations§
Source§impl NgramEmbedding
impl NgramEmbedding
pub fn new( dims: NgramDims, word: DenseWeight, tables: Vec<NgramTable>, projs: Vec<DenseWeight>, max_tokens: usize, gpu: &dyn GpuBackend, ) -> Result<Self>
Sourcepub fn embed(
&mut self,
ctx_tokens: &[u32],
seq_len: usize,
out: DevicePtr,
gpu: &dyn GpuBackend,
stream: u64,
) -> Result<()>
pub fn embed( &mut self, ctx_tokens: &[u32], seq_len: usize, out: DevicePtr, gpu: &dyn GpuBackend, stream: u64, ) -> Result<()>
Fused embedding for the LAST seq_len tokens of ctx (ctx =
up to n-1 cached context tokens followed by the new tokens —
exactly the reference NgramCache contract). Writes
[seq_len, hidden] BF16 to out.