spark_model/model/nllb/
model_impl.rs

1// SPDX-License-Identifier: AGPL-3.0-only
2
3//! `impl Model for NllbGpuModel`. The model owns all KV, returns bf16 logits
4//! (the scheduler's default sampling path), and drives the encoder inside
5//! `prefill` (seeding the decoder with `[decoder_start, forced_bos]` so the
6//! forced target-language token stays out of the sampled stream). Speculative /
7//! SSM / MTP / graph paths are not applicable and are stubbed.
8
9use anyhow::{Context, Result, bail};
10use spark_runtime::gpu::DevicePtr;
11
12use super::NllbGpuModel;
13use super::kv::NllbSeqKv;
14use crate::traits::{Model, SequenceState};
15
16impl NllbGpuModel {
17    /// Encoder pass + decoder seed, shared by `prefill` and `prefill_chunk`.
18    fn prefill_impl(&self, tokens: &[u32], seq: &mut SequenceState) -> Result<DevicePtr> {
19        let slot = seq.slot_idx;
20        // Per-request LoRA: apply the adapter only when the request selected it.
21        self.set_lora_active(seq.adapter_slot);
22        // Per-request source/target language (0 → deployment default).
23        let src_id = if seq.src_lang_id != 0 {
24            seq.src_lang_id
25        } else {
26            self.lang.src_lang_id
27        };
28        let tgt_id = if seq.tgt_lang_id != 0 {
29            seq.tgt_lang_id
30        } else {
31            self.lang.tgt_lang_id
32        };
33        let src = self.lang.encoder_input_with(src_id, tokens);
34        {
35            let mut map = self.kv.lock().unwrap();
36            let kv = map
37                .get_mut(&slot)
38                .context("nllb prefill: no KV state for slot (alloc_sequence not called?)")?;
39            self.run_encoder(&src, kv)?;
40            // Seed the decoder: step 0 consumes `decoder_start` (logits ignored),
41            // step 1 consumes `forced_bos = tgt_lang` and yields logits for the
42            // FIRST real translation token — so the forced token never enters the
43            // sampled stream and the scheduler samples normally from here. Each
44            // concurrent prefill writes its OWN row (by slot) so they don't race.
45            let row = self.prefill_logits_row(slot);
46            self.forward_one(self.lang.decoder_start_id, kv, row)?;
47            self.forward_one(tgt_id, kv, row)?;
48        }
49        self.sync()?;
50        // Decoder-side bookkeeping (the scheduler appends generated tokens and
51        // uses seq_len for its accounting; NLLB's own KV drives the real compute).
52        seq.tokens = vec![self.lang.decoder_start_id, tgt_id];
53        seq.seq_len = seq.tokens.len();
54        seq.prompt_len = seq.tokens.len();
55        Ok(self.prefill_logits_row(slot))
56    }
57}
58
59impl Model for NllbGpuModel {
60    fn prefill(&self, tokens: &[u32], seq: &mut SequenceState, _stream: u64) -> Result<DevicePtr> {
61        self.prefill_impl(tokens, seq)
62    }
63
64    fn prefill_chunk(
65        &self,
66        tokens: &[u32],
67        seq: &mut SequenceState,
68        chunk_start: usize,
69        chunk_len: usize,
70        is_last_chunk: bool,
71        _stream: u64,
72    ) -> Result<DevicePtr> {
73        if chunk_start != 0 || !is_last_chunk || chunk_len != tokens.len() {
74            bail!("nllb: chunked prefill unsupported — the source must arrive in one chunk");
75        }
76        self.prefill_impl(tokens, seq)
77    }
78
79    fn decode(&self, token: u32, seq: &mut SequenceState, _stream: u64) -> Result<DevicePtr> {
80        let slot = seq.slot_idx;
81        self.set_lora_active(seq.adapter_slot);
82        let row = self.decode_logits_row(0);
83        {
84            let mut map = self.kv.lock().unwrap();
85            let kv = map
86                .get_mut(&slot)
87                .context("nllb decode: no KV state for slot")?;
88            self.forward_one(token, kv, row)?;
89        }
90        self.sync()?;
91        seq.tokens.push(token);
92        seq.seq_len += 1;
93        Ok(row)
94    }
95
96    /// Batched decode: one `forward_one` per sequence into CONTIGUOUS logit rows
97    /// `0..n` (batch position `i` ↔ `seqs[i]`), the scheduler's row contract.
98    /// Each sequence's own per-slot KV is looked up by `slot_idx`, so batch
99    /// order is irrelevant. Sequences are processed serially on the default
100    /// stream (shared decode scratch); the returned base pointer is `[n, vocab]`.
101    fn decode_batch(
102        &self,
103        tokens: &[u32],
104        seqs: &mut [&mut SequenceState],
105        _stream: u64,
106    ) -> Result<DevicePtr> {
107        let n = seqs.len();
108        if n == 0 || tokens.len() != n {
109            bail!(
110                "nllb decode_batch: tokens/seqs length mismatch ({}, {n})",
111                tokens.len()
112            );
113        }
114        if n > self.max_batch {
115            bail!(
116                "nllb decode_batch: n={n} exceeds max_batch={}",
117                self.max_batch
118            );
119        }
120        {
121            let mut map = self.kv.lock().unwrap();
122            for (i, seq) in seqs.iter().enumerate() {
123                // Per-request LoRA gate, re-armed for each sequence in the batch.
124                self.set_lora_active(seq.adapter_slot);
125                let kv = map
126                    .get_mut(&seq.slot_idx)
127                    .context("nllb decode_batch: no KV state for slot")?;
128                self.forward_one(tokens[i], kv, self.decode_logits_row(i))?;
129            }
130        }
131        self.sync()?;
132        for (i, seq) in seqs.iter_mut().enumerate() {
133            seq.tokens.push(tokens[i]);
134            seq.seq_len += 1;
135        }
136        Ok(self.decode_logits_row(0))
137    }
138
139    fn vocab_size(&self) -> usize {
140        self.vocab
141    }
142
143    fn supports_beam(&self) -> bool {
144        true
145    }
146
147    fn generate_beam_batch(&self, reqs: &[crate::traits::BeamReq]) -> Result<Vec<Vec<u32>>> {
148        // Phase c: the C requests' Σ beams decode as ONE M=(Σ beams) batch per
149        // step (cross-request co-dispatch). Single-adapter batches fuse; mixed
150        // adapters fall back to serial inside `beam_batched_multi`.
151        self.beam_batched_multi(reqs)
152    }
153
154    fn bind_gpu_to_thread(&self) -> Result<()> {
155        self.gpu.bind_to_thread()
156    }
157
158    fn alloc_sequence(&self) -> Result<SequenceState> {
159        let slot = self.slots.lock().unwrap().claim();
160        let kv = NllbSeqKv::new(self.gpu.as_ref(), self.dec_layers, self.cache_rows, self.d)?;
161        self.kv.lock().unwrap().insert(slot, kv);
162        // All-defaults host-side state; NLLB's KV lives in `self.kv`,
163        // keyed by `slot` (SSOT for the field defaults: `host_only`).
164        Ok(SequenceState::host_only(slot))
165    }
166
167    fn free_sequence(&self, seq: &mut SequenceState) -> Result<()> {
168        let slot = seq.slot_idx;
169        if slot == usize::MAX {
170            return Ok(()); // migrated to a survivor by compact_sequence
171        }
172        if let Some(kv) = self.kv.lock().unwrap().remove(&slot) {
173            kv.free(self.gpu.as_ref())?;
174        }
175        self.slots.lock().unwrap().release(slot);
176        Ok(())
177    }
178
179    fn compact_sequence(&self, seq: &mut SequenceState, new_slot: usize) -> Result<()> {
180        let old = seq.slot_idx;
181        let mut map = self.kv.lock().unwrap();
182        if let Some(kv) = map.remove(&old) {
183            map.insert(new_slot, kv);
184        }
185        seq.slot_idx = new_slot;
186        Ok(())
187    }
188
189    fn detach_slot_for_reuse(&self, seq: &mut SequenceState) {
190        seq.slot_idx = usize::MAX;
191    }
192
193    fn cache_sequence(&self, _seq: &SequenceState) {
194        // No prefix reuse for translation.
195    }
196
197    fn num_free_blocks(&self) -> usize {
198        // The model owns its KV outside the paged block cache; report ample
199        // headroom so the scheduler's swap/admission math never rejects.
200        1 << 20
201    }
202
203    fn copy_logits_to_host(&self, logits_ptr: DevicePtr, dst: &mut [u8]) -> Result<()> {
204        self.gpu.copy_d2h(logits_ptr, dst)
205    }
206
207    fn logits_buffer_ptr(&self) -> DevicePtr {
208        self.decode_logits_row(0)
209    }
210
211    fn argmax_on_device(&self, logits_ptr: DevicePtr, _stream: u64) -> Result<u32> {
212        self.argmax_of(logits_ptr)
213    }
214
215    fn argmax_batch(&self, logits_ptr: DevicePtr, n: usize, _stream: u64) -> Result<Vec<u32>> {
216        // `logits_ptr` is `[n, vocab]` bf16; argmax each contiguous row.
217        (0..n)
218            .map(|i| self.argmax_of(logits_ptr.offset(i * self.vocab * 2)))
219            .collect()
220    }
221
222    // ── inapplicable paths: NLLB is non-speculative / non-SSM / non-MTP ──
223
224    fn hidden_after_norm(&self) -> DevicePtr {
225        DevicePtr(0)
226    }
227
228    fn decode_verify(&self, _t: &[u32], _s: &mut SequenceState, _st: u64) -> Result<Vec<u32>> {
229        bail!("nllb: speculative verify not supported")
230    }
231
232    fn checkpoint_ssm_states(&self, _seq: &mut SequenceState) -> Result<()> {
233        Ok(())
234    }
235
236    fn rollback_ssm_states(&self, _seq: &mut SequenceState, _n: usize) -> Result<()> {
237        Ok(())
238    }
239
240    fn generate_speculative(
241        &self,
242        _tokens: &[u32],
243        _params: &spark_runtime::sampler::SamplingParams,
244        _num_drafts: usize,
245    ) -> Result<crate::engine::GenerateResult> {
246        bail!("nllb: speculative decoding not supported")
247    }
248
249    fn has_proposer(&self) -> bool {
250        false
251    }
252
253    fn has_self_speculative(&self) -> bool {
254        false
255    }
256
257    fn decode_draft(
258        &self,
259        _token: u32,
260        _seq: &mut SequenceState,
261        _stream: u64,
262    ) -> Result<DevicePtr> {
263        bail!("nllb: self-speculative draft not supported")
264    }
265
266    fn decode_verify_graphed(
267        &self,
268        _t: &[u32; 2],
269        _s: &mut SequenceState,
270        _st: u64,
271    ) -> Result<[u32; 2]> {
272        bail!("nllb: graphed verify not supported")
273    }
274
275    fn decode_verify_graphed_k3(
276        &self,
277        _t: &[u32; 3],
278        _s: &mut SequenceState,
279        _st: u64,
280    ) -> Result<[u32; 3]> {
281        bail!("nllb: graphed verify not supported")
282    }
283
284    fn decode_verify_graphed_k4(
285        &self,
286        _t: &[u32; 4],
287        _s: &mut SequenceState,
288        _st: u64,
289    ) -> Result<[u32; 4]> {
290        bail!("nllb: graphed verify not supported")
291    }
292
293    fn save_hidden_for_mtp(&self, _token_idx: usize, _stream: u64) -> Result<()> {
294        Ok(())
295    }
296
297    fn run_mtp_propose(
298        &self,
299        _token: u32,
300        _position: usize,
301        _seq: &mut SequenceState,
302        _stream: u64,
303    ) -> Result<Option<u32>> {
304        Ok(None)
305    }
306
307    fn run_mtp_propose_multi(
308        &self,
309        _token: u32,
310        _position: usize,
311        _num_drafts: usize,
312        _seq: &mut SequenceState,
313        _stream: u64,
314        _grammar_bitmask: Option<&[i32]>,
315    ) -> Result<Vec<u32>> {
316        Ok(Vec::new())
317    }
318
319    fn trim_proposer_state(
320        &self,
321        _seq: &mut SequenceState,
322        _num_accepted: usize,
323        _stream: u64,
324    ) -> Result<()> {
325        Ok(())
326    }
327}