1#![forbid(unsafe_code)]
12
13use std::fmt;
14use std::future::Future;
15use std::pin::Pin;
16use std::sync::atomic::{AtomicUsize, Ordering};
17use std::sync::{Arc, Mutex, PoisonError};
18use std::task::{Context, Poll, Waker};
19use std::time::{Duration, Instant};
20
21use kime_core::answer::{LAYA_MODEL, Response, Temperatures, laya_answer};
22use kime_core::render::{compat_question, compat_state};
23use kime_core::request::{Limits, Problem, Question, Request, parse};
24use kime_model::Model;
25use kime_tensor::{BatchBuf, Buckets, Executor, Outputs};
26use kime_tok::Tokenizer;
27use kime_tok::layout::{CompatBudget, CompatSequence, Cut};
28use serde_json::Value;
29
30pub mod hub;
31mod split;
32
33#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
35pub enum Device {
36 #[default]
38 Auto,
39 Cpu {
41 threads: usize,
43 },
44 Cuda(usize),
46 Metal,
48 Ane,
50}
51
52#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
54pub enum Precision {
55 #[default]
57 F16,
58 F32,
60 Int8,
64}
65
66#[derive(Debug)]
68pub enum Error {
69 Invalid(Vec<Problem>),
71 NotFound(String),
73 Model(kime_model::Error),
75 Tokenizer(String),
77 Backend(kime_tensor::Error),
79 TooLong {
82 question: String,
84 options: usize,
86 fit: usize,
88 },
89 Unsupported(String),
91}
92
93impl fmt::Display for Error {
94 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
95 match self {
96 Error::Invalid(p) => {
97 write!(f, "invalid request:")?;
98 for p in p {
99 write!(f, " {};", p.to_json())?;
100 }
101 Ok(())
102 }
103 Error::NotFound(m) | Error::Tokenizer(m) | Error::Unsupported(m) => f.write_str(m),
104 Error::Model(e) => write!(f, "{e}"),
105 Error::Backend(e) => write!(f, "{e}"),
106 Error::TooLong { question, options, fit } => write!(
107 f,
108 "question {question:?} has {options} options but only {fit} fit the head budget"
109 ),
110 }
111 }
112}
113
114impl std::error::Error for Error {}
115
116impl From<kime_tensor::Error> for Error {
117 fn from(e: kime_tensor::Error) -> Self {
118 Error::Backend(e)
119 }
120}
121
122#[derive(Debug, Clone, Default)]
124pub struct Builder {
125 model: Option<String>,
126 device: Device,
127 precision: Precision,
128 preload: bool,
129}
130
131impl Builder {
132 #[must_use]
135 pub fn model(mut self, name: impl Into<String>) -> Self {
136 self.model = Some(name.into());
137 self
138 }
139
140 #[must_use]
142 pub fn device(mut self, d: Device) -> Self {
143 self.device = d;
144 self
145 }
146
147 #[must_use]
149 pub fn precision(mut self, p: Precision) -> Self {
150 self.precision = p;
151 self
152 }
153
154 #[must_use]
157 pub fn preload(mut self, yes: bool) -> Self {
158 self.preload = yes;
159 self
160 }
161
162 pub fn build(self) -> Result<Kime, Error> {
168 let name = self.model.unwrap_or_else(|| "laya".into());
169 let path = hub::resolve(&name).map_err(Error::NotFound)?;
170 let model = Model::open(&path).map_err(Error::Model)?;
171 let tok_json = model
172 .file("tokenizer/tokenizer.json")
173 .ok_or_else(|| Error::Tokenizer(format!("{}: no tokenizer.json", path.display())))?;
174 let tok = Tokenizer::from_bytes(tok_json, model.file("tokenizer/tokenizer_config.json"))
175 .map_err(|e| Error::Tokenizer(e.to_string()))?;
176 let agent = &model.spec.agent;
177 let temps = Temperatures::new(agent.temperature, &agent.temperature_by_options);
178 let budget = CompatBudget { max_len: agent.max_len, head_max_len: agent.head_max_len };
179 let mut runner = Runner::open(&model, self.device, self.precision)?;
180 if self.preload {
181 for b in Buckets::default().stage("compat").iter().take(4) {
182 runner.prepare(*b)?;
183 }
184 }
185 let buckets = Buckets::default().stage("compat").to_vec();
186 let (weights, plans) = runner.memory();
187 Ok(Kime {
188 inner: Arc::new(Inner {
189 id: model.spec.id.clone(),
190 mask: tok.mask_text().to_string(),
191 tok,
192 budget,
193 temps,
194 buckets,
195 memory: [AtomicUsize::new(weights), AtomicUsize::new(plans)],
196 runner: Mutex::new(Session {
197 runner,
198 buf: BatchBuf::default(),
199 out: Outputs::default(),
200 }),
201 }),
202 })
203 }
204}
205
206enum Runner {
207 Cpu(Box<Executor<kime_cpu::CpuBackend>>),
208 #[cfg(feature = "cuda")]
209 Cuda(Box<Executor<kime_cuda::CudaBackend>>),
210}
211
212impl Runner {
213 fn open(model: &Model, device: Device, precision: Precision) -> Result<Self, Error> {
214 let cpu = |threads: usize| {
215 let t = if threads == 0 { kime_cpu::par::available() } else { threads };
216 let backend = kime_cpu::CpuBackend::new(t).with_int8(precision == Precision::Int8);
217 Ok(Runner::Cpu(Box::new(kime_cpu::executor_with(model, backend)?)))
218 };
219 match device {
220 Device::Cpu { threads } => cpu(threads),
221 #[cfg(feature = "cuda")]
222 Device::Cuda(n) => {
223 Ok(Runner::Cuda(Box::new(kime_cuda::executor(model, n, cuda(precision)?)?)))
224 }
225 #[cfg(feature = "cuda")]
226 Device::Auto => {
227 match cuda(precision).and_then(|p| Ok(kime_cuda::executor(model, 0, p)?)) {
228 Ok(e) => Ok(Runner::Cuda(Box::new(e))),
229 Err(_) => cpu(0),
230 }
231 }
232 #[cfg(not(feature = "cuda"))]
233 Device::Auto => cpu(0),
234 #[cfg(not(feature = "cuda"))]
235 Device::Cuda(_) => Err(Error::Unsupported("this build has no CUDA backend".into())),
236 Device::Metal | Device::Ane => {
237 Err(Error::Unsupported("the Apple backends arrive with M4".into()))
238 }
239 }
240 }
241
242 fn prepare(&mut self, b: kime_tensor::Bucket) -> Result<(), Error> {
243 match self {
244 Runner::Cpu(e) => e.prepare(b)?,
245 #[cfg(feature = "cuda")]
246 Runner::Cuda(e) => e.prepare(b)?,
247 };
248 Ok(())
249 }
250
251 fn run(&mut self, buf: &BatchBuf, out: &mut Outputs) -> Result<(), Error> {
252 match self {
253 Runner::Cpu(e) => e.run(&buf.batch(), out)?,
254 #[cfg(feature = "cuda")]
255 Runner::Cuda(e) => e.run(&buf.batch(), out)?,
256 };
257 Ok(())
258 }
259
260 fn memory(&self) -> (usize, usize) {
261 match self {
262 Runner::Cpu(e) => e.memory(),
263 #[cfg(feature = "cuda")]
264 Runner::Cuda(e) => e.memory(),
265 }
266 }
267
268 fn describe(&self) -> String {
269 match self {
270 Runner::Cpu(e) => {
271 let int8 = if e.backend().int8() { ", int8" } else { "" };
272 format!("cpu, {} threads{int8}", kime_tensor::Backend::caps(e.backend()).threads)
273 }
274 #[cfg(feature = "cuda")]
275 Runner::Cuda(e) => format!("cuda, {}", e.backend().name()),
276 }
277 }
278}
279
280#[cfg(feature = "cuda")]
281fn cuda(p: Precision) -> Result<kime_cuda::Precision, Error> {
282 match p {
283 Precision::F16 => Ok(kime_cuda::Precision::F16),
284 Precision::F32 => Ok(kime_cuda::Precision::F32),
285 Precision::Int8 => Err(Error::Unsupported("INT8 runs on the CPU only for now".into())),
286 }
287}
288
289struct Session {
290 runner: Runner,
291 buf: BatchBuf,
292 out: Outputs,
293}
294
295struct Inner {
296 id: String,
297 tok: Tokenizer,
298 mask: String,
299 budget: CompatBudget,
300 temps: Temperatures,
301 buckets: Vec<kime_tensor::Bucket>,
303 runner: Mutex<Session>,
304 memory: [AtomicUsize; 2],
306}
307
308#[derive(Clone)]
310pub struct Kime {
311 inner: Arc<Inner>,
312}
313
314impl fmt::Debug for Kime {
315 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
316 f.debug_struct("Kime").field("model", &self.inner.id).finish_non_exhaustive()
317 }
318}
319
320#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
322pub struct Timing {
323 pub tokenize: Duration,
325 pub device: Duration,
327 pub batches: usize,
329 pub truncated: usize,
331 pub cut_tokens: usize,
333}
334
335#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
337pub struct Memory {
338 pub weights: usize,
340 pub plans: usize,
342}
343
344struct Item<'a> {
346 req: usize,
347 q: &'a Question,
348 seq: CompatSequence,
349 logits: Vec<f32>,
350 act: [f32; 2],
351}
352
353impl Kime {
354 #[must_use]
356 pub fn builder() -> Builder {
357 Builder::default()
358 }
359
360 #[must_use]
362 pub fn model_id(&self) -> &str {
363 &self.inner.id
364 }
365
366 #[must_use]
368 pub fn device(&self) -> String {
369 self.lock().runner.describe()
370 }
371
372 #[must_use]
375 pub fn max_row_tokens(&self) -> usize {
376 self.inner.budget.max_len
377 }
378
379 #[must_use]
381 pub fn memory(&self) -> Memory {
382 let m = &self.inner.memory;
383 Memory { weights: m[0].load(Ordering::Relaxed), plans: m[1].load(Ordering::Relaxed) }
384 }
385
386 fn lock(&self) -> std::sync::MutexGuard<'_, Session> {
387 self.inner.runner.lock().unwrap_or_else(PoisonError::into_inner)
388 }
389
390 #[must_use]
394 pub fn count_tokens(&self, req: &Request) -> usize {
395 let inner = &*self.inner;
396 let mut n = inner.tok.encode(&compat_state(&req.state, &inner.mask)).len();
397 for q in &req.questions {
398 let text = compat_question(q, &inner.mask);
399 n += inner.tok.encode(&text.head).len();
400 n += text.options.iter().map(|o| inner.tok.encode(o).len()).sum::<usize>();
401 }
402 n
403 }
404
405 pub fn decide(&self, req: &Request) -> Result<Response, Error> {
412 Ok(self.decide_batch(std::slice::from_ref(req))?.remove(0))
413 }
414
415 pub fn decide_batch(&self, reqs: &[Request]) -> Result<Vec<Response>, Error> {
422 Ok(self.decide_batch_timed(reqs)?.0)
423 }
424
425 pub fn decide_batch_timed(&self, reqs: &[Request]) -> Result<(Vec<Response>, Timing), Error> {
431 let inner = &*self.inner;
432 let t0 = Instant::now();
433 let mut parsed = Vec::with_capacity(reqs.len());
434 for r in reqs {
435 parsed.push(parse(&r.to_json(), &Limits::LAYA).map_err(Error::Invalid)?);
436 }
437 let mut items = Vec::new();
438 for (i, r) in parsed.iter().enumerate() {
439 if r.questions.is_empty() {
440 continue;
441 }
442 let cut = if matches!(r.state, Value::Array(_)) { Cut::Head } else { Cut::Tail };
444 let state = inner.tok.encode_state(&compat_state(&r.state, &inner.mask));
445 for q in &r.questions {
446 let text = compat_question(q, &inner.mask);
447 let seq =
448 inner.tok.compat_sequence(&text.head, &text.options, &state, inner.budget, cut);
449 if seq.markers.len() != q.criteria.len() {
450 return Err(Error::TooLong {
451 question: q.id.clone(),
452 options: q.criteria.len(),
453 fit: seq.markers.len(),
454 });
455 }
456 items.push(Item { req: i, q, seq, logits: Vec::new(), act: [0.0; 2] });
457 }
458 }
459 let tokenize = t0.elapsed();
460 let t1 = Instant::now();
461 let batches = self.run(&mut items)?;
462 let cut = items.iter().map(|it| it.seq.state_tokens - it.seq.state_tokens_used);
463 let timing = Timing {
464 tokenize,
465 device: t1.elapsed(),
466 batches,
467 truncated: cut.clone().filter(|&n| n > 0).count(),
468 cut_tokens: cut.sum(),
469 };
470 let mut out: Vec<Response> = parsed
471 .iter()
472 .map(|_| Response { model: LAYA_MODEL.into(), answers: Vec::new(), input_tokens: 0 })
473 .collect();
474 for it in &items {
475 let res = &mut out[it.req];
476 res.input_tokens += it.seq.ids.len();
477 res.answers
478 .push((it.q.id.clone(), laya_answer(it.q, &it.logits, it.act, &inner.temps)));
479 }
480 Ok((out, timing))
481 }
482
483 fn run(&self, items: &mut [Item<'_>]) -> Result<usize, Error> {
485 let sizes: Vec<(usize, usize)> =
486 items.iter().map(|it| (it.seq.ids.len(), it.seq.markers.len())).collect();
487 let batches = split::split(&self.inner.buckets, &sizes);
488 let mut s = self.lock();
489 let Session { runner, buf, out } = &mut *s;
490 for batch in &batches {
491 buf.clear();
492 for &i in batch {
493 let it = &items[i];
494 buf.push(&it.seq.ids, &it.seq.markers, it.q.qtype.index() as u8);
495 }
496 runner.run(buf, out)?;
497 let mut at = 0;
498 for (&i, a) in batch.iter().zip(&out.act) {
499 let it = &mut items[i];
500 let k = it.seq.markers.len();
501 it.logits = out.logits[at..at + k].to_vec();
502 it.act = *a;
503 at += k;
504 }
505 }
506 let (weights, plans) = runner.memory();
507 self.inner.memory[0].store(weights, Ordering::Relaxed);
508 self.inner.memory[1].store(plans, Ordering::Relaxed);
509 Ok(batches.len())
510 }
511
512 #[must_use]
515 pub fn decide_async(&self, req: &Request) -> Decision {
516 let shared = Arc::new(Mutex::new((None, None::<Waker>)));
517 let (kime, req, done) = (self.clone(), req.clone(), shared.clone());
518 std::thread::spawn(move || {
519 let r = kime.decide(&req);
520 let mut g = done.lock().unwrap_or_else(PoisonError::into_inner);
521 g.0 = Some(r);
522 if let Some(w) = g.1.take() {
523 w.wake();
524 }
525 });
526 Decision { shared }
527 }
528}
529
530#[derive(Debug)]
532pub struct Decision {
533 #[allow(clippy::type_complexity)]
534 shared: Arc<Mutex<(Option<Result<Response, Error>>, Option<Waker>)>>,
535}
536
537impl Future for Decision {
538 type Output = Result<Response, Error>;
539
540 fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
541 let mut g = self.shared.lock().unwrap_or_else(PoisonError::into_inner);
542 match g.0.take() {
543 Some(r) => Poll::Ready(r),
544 None => {
545 g.1 = Some(cx.waker().clone());
546 Poll::Pending
547 }
548 }
549 }
550}