#![forbid(unsafe_code)]
use std::fmt;
use std::future::Future;
use std::pin::Pin;
use std::sync::{Arc, Mutex, PoisonError};
use std::task::{Context, Poll, Waker};
use kime_core::answer::{LAYA_MODEL, Response, Temperatures, laya_answer};
use kime_core::render::{compat_question, compat_state};
use kime_core::request::{Limits, Problem, Question, Request, parse};
use kime_model::Model;
use kime_tensor::{BatchBuf, Buckets, Executor, Outputs};
use kime_tok::Tokenizer;
use kime_tok::layout::{CompatBudget, CompatSequence, Cut};
use serde_json::Value;
pub mod hub;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum Device {
#[default]
Auto,
Cpu {
threads: usize,
},
Cuda(usize),
Metal,
Ane,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum Precision {
#[default]
F16,
F32,
Int8,
}
#[derive(Debug)]
pub enum Error {
Invalid(Vec<Problem>),
NotFound(String),
Model(kime_model::Error),
Tokenizer(String),
Backend(kime_tensor::Error),
TooLong {
question: String,
options: usize,
fit: usize,
},
Unsupported(String),
}
impl fmt::Display for Error {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Error::Invalid(p) => {
write!(f, "invalid request:")?;
for p in p {
write!(f, " {};", p.to_json())?;
}
Ok(())
}
Error::NotFound(m) | Error::Tokenizer(m) | Error::Unsupported(m) => f.write_str(m),
Error::Model(e) => write!(f, "{e}"),
Error::Backend(e) => write!(f, "{e}"),
Error::TooLong { question, options, fit } => write!(
f,
"question {question:?} has {options} options but only {fit} fit the head budget"
),
}
}
}
impl std::error::Error for Error {}
impl From<kime_tensor::Error> for Error {
fn from(e: kime_tensor::Error) -> Self {
Error::Backend(e)
}
}
#[derive(Debug, Clone, Default)]
pub struct Builder {
model: Option<String>,
device: Device,
precision: Precision,
preload: bool,
}
impl Builder {
#[must_use]
pub fn model(mut self, name: impl Into<String>) -> Self {
self.model = Some(name.into());
self
}
#[must_use]
pub fn device(mut self, d: Device) -> Self {
self.device = d;
self
}
#[must_use]
pub fn precision(mut self, p: Precision) -> Self {
self.precision = p;
self
}
#[must_use]
pub fn preload(mut self, yes: bool) -> Self {
self.preload = yes;
self
}
pub fn build(self) -> Result<Kime, Error> {
let name = self.model.unwrap_or_else(|| "laya".into());
let path = hub::resolve(&name).map_err(Error::NotFound)?;
let model = Model::open(&path).map_err(Error::Model)?;
let tok_json = model
.file("tokenizer/tokenizer.json")
.ok_or_else(|| Error::Tokenizer(format!("{}: no tokenizer.json", path.display())))?;
let tok = Tokenizer::from_bytes(tok_json, model.file("tokenizer/tokenizer_config.json"))
.map_err(|e| Error::Tokenizer(e.to_string()))?;
let agent = &model.spec.agent;
let temps = Temperatures::new(agent.temperature, &agent.temperature_by_options);
let budget = CompatBudget { max_len: agent.max_len, head_max_len: agent.head_max_len };
let mut runner = Runner::open(&model, self.device, self.precision)?;
if self.preload {
for b in Buckets::default().stage("compat").iter().take(4) {
runner.prepare(*b)?;
}
}
let buckets = Buckets::default();
let largest = buckets.stage("compat").last().copied().unwrap_or(kime_tensor::Bucket {
tokens: 16384,
seqs: 1024,
markers: 8192,
});
Ok(Kime {
inner: Arc::new(Inner {
id: model.spec.id.clone(),
mask: tok.mask_text().to_string(),
tok,
budget,
temps,
limits: (largest.tokens, largest.seqs, largest.markers),
runner: Mutex::new(Session {
runner,
buf: BatchBuf::default(),
out: Outputs::default(),
}),
}),
})
}
}
enum Runner {
Cpu(Box<Executor<kime_cpu::CpuBackend>>),
#[cfg(feature = "cuda")]
Cuda(Box<Executor<kime_cuda::CudaBackend>>),
}
impl Runner {
fn open(model: &Model, device: Device, precision: Precision) -> Result<Self, Error> {
let cpu = |threads: usize| {
let t = if threads == 0 { kime_cpu::par::available() } else { threads };
let backend = kime_cpu::CpuBackend::new(t).with_int8(precision == Precision::Int8);
Ok(Runner::Cpu(Box::new(kime_cpu::executor_with(model, backend)?)))
};
match device {
Device::Cpu { threads } => cpu(threads),
#[cfg(feature = "cuda")]
Device::Cuda(n) => {
Ok(Runner::Cuda(Box::new(kime_cuda::executor(model, n, cuda(precision)?)?)))
}
#[cfg(feature = "cuda")]
Device::Auto => {
match cuda(precision).and_then(|p| Ok(kime_cuda::executor(model, 0, p)?)) {
Ok(e) => Ok(Runner::Cuda(Box::new(e))),
Err(_) => cpu(0),
}
}
#[cfg(not(feature = "cuda"))]
Device::Auto => cpu(0),
#[cfg(not(feature = "cuda"))]
Device::Cuda(_) => Err(Error::Unsupported("this build has no CUDA backend".into())),
Device::Metal | Device::Ane => {
Err(Error::Unsupported("the Apple backends arrive with M4".into()))
}
}
}
fn prepare(&mut self, b: kime_tensor::Bucket) -> Result<(), Error> {
match self {
Runner::Cpu(e) => e.prepare(b)?,
#[cfg(feature = "cuda")]
Runner::Cuda(e) => e.prepare(b)?,
};
Ok(())
}
fn run(&mut self, buf: &BatchBuf, out: &mut Outputs) -> Result<(), Error> {
match self {
Runner::Cpu(e) => e.run(&buf.batch(), out)?,
#[cfg(feature = "cuda")]
Runner::Cuda(e) => e.run(&buf.batch(), out)?,
};
Ok(())
}
fn describe(&self) -> String {
match self {
Runner::Cpu(e) => {
let int8 = if e.backend().int8() { ", int8" } else { "" };
format!("cpu, {} threads{int8}", kime_tensor::Backend::caps(e.backend()).threads)
}
#[cfg(feature = "cuda")]
Runner::Cuda(e) => format!("cuda, {}", e.backend().name()),
}
}
}
#[cfg(feature = "cuda")]
fn cuda(p: Precision) -> Result<kime_cuda::Precision, Error> {
match p {
Precision::F16 => Ok(kime_cuda::Precision::F16),
Precision::F32 => Ok(kime_cuda::Precision::F32),
Precision::Int8 => Err(Error::Unsupported("INT8 runs on the CPU only for now".into())),
}
}
struct Session {
runner: Runner,
buf: BatchBuf,
out: Outputs,
}
struct Inner {
id: String,
tok: Tokenizer,
mask: String,
budget: CompatBudget,
temps: Temperatures,
limits: (usize, usize, usize),
runner: Mutex<Session>,
}
#[derive(Clone)]
pub struct Kime {
inner: Arc<Inner>,
}
impl fmt::Debug for Kime {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Kime").field("model", &self.inner.id).finish_non_exhaustive()
}
}
struct Item<'a> {
req: usize,
q: &'a Question,
seq: CompatSequence,
logits: Vec<f32>,
act: [f32; 2],
}
impl Kime {
#[must_use]
pub fn builder() -> Builder {
Builder::default()
}
#[must_use]
pub fn model_id(&self) -> &str {
&self.inner.id
}
#[must_use]
pub fn device(&self) -> String {
self.lock().runner.describe()
}
fn lock(&self) -> std::sync::MutexGuard<'_, Session> {
self.inner.runner.lock().unwrap_or_else(PoisonError::into_inner)
}
pub fn decide(&self, req: &Request) -> Result<Response, Error> {
Ok(self.decide_batch(std::slice::from_ref(req))?.remove(0))
}
pub fn decide_batch(&self, reqs: &[Request]) -> Result<Vec<Response>, Error> {
let inner = &*self.inner;
let mut parsed = Vec::with_capacity(reqs.len());
for r in reqs {
parsed.push(parse(&r.to_json(), &Limits::LAYA).map_err(Error::Invalid)?);
}
let mut items = Vec::new();
for (i, r) in parsed.iter().enumerate() {
if r.questions.is_empty() {
continue;
}
let cut = if matches!(r.state, Value::Array(_)) { Cut::Head } else { Cut::Tail };
let state = inner.tok.encode_state(&compat_state(&r.state, &inner.mask));
for q in &r.questions {
let text = compat_question(q, &inner.mask);
let seq =
inner.tok.compat_sequence(&text.head, &text.options, &state, inner.budget, cut);
if seq.markers.len() != q.criteria.len() {
return Err(Error::TooLong {
question: q.id.clone(),
options: q.criteria.len(),
fit: seq.markers.len(),
});
}
items.push(Item { req: i, q, seq, logits: Vec::new(), act: [0.0; 2] });
}
}
self.run(&mut items)?;
let mut out: Vec<Response> = parsed
.iter()
.map(|_| Response { model: LAYA_MODEL.into(), answers: Vec::new(), input_tokens: 0 })
.collect();
for it in &items {
let res = &mut out[it.req];
res.input_tokens += it.seq.ids.len();
res.answers
.push((it.q.id.clone(), laya_answer(it.q, &it.logits, it.act, &inner.temps)));
}
Ok(out)
}
fn run(&self, items: &mut [Item<'_>]) -> Result<(), Error> {
let (max_t, max_s, max_m) = self.inner.limits;
let mut s = self.lock();
let Session { runner, buf, out } = &mut *s;
let mut start = 0;
while start < items.len() {
let (mut t, mut m, mut end) = (0, 0, start);
while end < items.len() && end - start < max_s {
let it = &items[end];
if end > start && (t + it.seq.ids.len() > max_t || m + it.seq.markers.len() > max_m)
{
break;
}
t += it.seq.ids.len();
m += it.seq.markers.len();
end += 1;
}
buf.clear();
for it in &items[start..end] {
buf.push(&it.seq.ids, &it.seq.markers, it.q.qtype.index() as u8);
}
runner.run(buf, out)?;
let mut at = 0;
for (it, a) in items[start..end].iter_mut().zip(&out.act) {
let k = it.seq.markers.len();
it.logits = out.logits[at..at + k].to_vec();
it.act = *a;
at += k;
}
start = end;
}
Ok(())
}
#[must_use]
pub fn decide_async(&self, req: &Request) -> Decision {
let shared = Arc::new(Mutex::new((None, None::<Waker>)));
let (kime, req, done) = (self.clone(), req.clone(), shared.clone());
std::thread::spawn(move || {
let r = kime.decide(&req);
let mut g = done.lock().unwrap_or_else(PoisonError::into_inner);
g.0 = Some(r);
if let Some(w) = g.1.take() {
w.wake();
}
});
Decision { shared }
}
}
#[derive(Debug)]
pub struct Decision {
#[allow(clippy::type_complexity)]
shared: Arc<Mutex<(Option<Result<Response, Error>>, Option<Waker>)>>,
}
impl Future for Decision {
type Output = Result<Response, Error>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let mut g = self.shared.lock().unwrap_or_else(PoisonError::into_inner);
match g.0.take() {
Some(r) => Poll::Ready(r),
None => {
g.1 = Some(cx.waker().clone());
Poll::Pending
}
}
}
}