1use 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 fn prefill_impl(&self, tokens: &[u32], seq: &mut SequenceState) -> Result<DevicePtr> {
19 let slot = seq.slot_idx;
20 self.set_lora_active(seq.adapter_slot);
22 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 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 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 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 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 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 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(()); }
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 }
196
197 fn num_free_blocks(&self) -> usize {
198 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 (0..n)
218 .map(|i| self.argmax_of(logits_ptr.offset(i * self.vocab * 2)))
219 .collect()
220 }
221
222 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}