Skip to main content

laya/
lib.rs

1//! Rust inference for [Laya](https://huggingface.co/convaiinnovations/laya), a non-autoregressive
2//! typed-decision model: a ModernBERT-large encoder plus an RL-trained decision head.
3//!
4//! Give it a state (text or JSON) and a set of typed questions; it returns typed answers with
5//! calibrated probabilities in a single forward pass. It never generates text.
6//!
7//! ```no_run
8//! use laya::{Agent, Question};
9//! use serde_json::json;
10//!
11//! let agent = Agent::from_dir("models/laya-base", Default::default())?;
12//! let answers = agent.system_one(
13//!     &json!("My card was charged twice for the same order."),
14//!     &[("department".into(), Question::choice("Which team owns this?", ["billing", "technical", "sales"]))]
15//!         .into_iter()
16//!         .collect(),
17//! )?;
18//! println!("{}", serde_json::to_string_pretty(&answers)?);
19//! # Ok::<(), anyhow::Error>(())
20//! ```
21
22pub mod config;
23pub mod model;
24pub mod question;
25
26use anyhow::{Context, Result};
27use candle_core::{DType, Device, Tensor};
28use candle_nn::VarBuilder;
29use serde::Serialize;
30use serde_json::{Map, Value};
31use std::collections::HashMap;
32use std::path::Path;
33use tokenizers::Tokenizer;
34
35pub use config::{AgentConfig, EncoderConfig};
36pub use question::{QType, Question};
37
38use question::SpecialIds;
39
40/// How to place and run the model.
41#[derive(Debug, Clone, Default)]
42pub struct Options {
43    /// `None` picks the best available: CUDA, then Metal, then CPU.
44    pub device: Option<Device>,
45    /// Currently f32 only; see the note on [`Agent::from_dir`].
46    pub dtype: Option<DType>,
47}
48
49/// One typed answer. The variant matches the question's type.
50#[derive(Debug, Clone, Serialize)]
51#[serde(tag = "type", rename_all = "lowercase")]
52pub enum Answer {
53    Choice {
54        /// The argmax option.
55        choice: String,
56        probabilities: Map<String, Value>,
57        /// 1 - normalized entropy of the answer distribution.
58        confidence: Value,
59        rl_agent: Meta,
60    },
61    Score {
62        /// Expectation over level indices.
63        score: Value,
64        legend: Map<String, Value>,
65        probabilities: Map<String, Value>,
66        confidence: Value,
67        rl_agent: Meta,
68    },
69    Noul {
70        /// P(the statement holds).
71        noul: Value,
72        rl_agent: Meta,
73    },
74}
75
76#[derive(Debug, Clone, Serialize)]
77pub struct Meta {
78    /// The act head's probability of answering rather than escalating.
79    pub act_probability: Value,
80}
81
82#[derive(Debug, Clone, Serialize)]
83pub struct Response {
84    pub model: String,
85    pub answers: Map<String, Value>,
86    pub usage: Usage,
87}
88
89#[derive(Debug, Clone, Serialize)]
90pub struct Usage {
91    pub input_tokens: usize,
92    pub output_tokens: usize,
93}
94
95/// Which part of the rendered sequence a token belongs to.
96#[derive(Debug, Clone, Copy, Serialize)]
97#[serde(rename_all = "lowercase")]
98pub enum Segment {
99    /// `[CLS]` / `[SEP]` structure.
100    Special,
101    /// `"<type> question: <instructions>"`.
102    Head,
103    /// An option's text. `option` carries its index.
104    Option,
105    /// The serialized state.
106    State,
107}
108
109/// One token of the rendered sequence.
110#[derive(Debug, Clone, Serialize)]
111pub struct TokenView {
112    /// The token text, with the byte-level BPE space marker turned back into a space.
113    pub text: String,
114    pub id: u32,
115    pub segment: Segment,
116    /// True for the `[MASK]` marker the head reads this option's logit from.
117    pub marker: bool,
118    /// The option this token belongs to, for `Option` and marker tokens.
119    #[serde(skip_serializing_if = "Option::is_none")]
120    pub option: Option<usize>,
121}
122
123/// The exact sequence the encoder sees for one question, for inspection and debugging.
124#[derive(Debug, Clone, Serialize)]
125pub struct PromptView {
126    pub tokens: Vec<TokenView>,
127    /// Rendered option texts, in label-index order.
128    pub options: Vec<String>,
129    /// Marker positions in the sequence.
130    pub markers: Vec<usize>,
131    pub total_tokens: usize,
132    pub max_len: usize,
133    /// True when the state was cut to fit `max_len`.
134    pub state_truncated: bool,
135}
136
137pub struct Agent {
138    model: model::DecisionModel,
139    tokenizer: Tokenizer,
140    specials: SpecialIds,
141    pub cfg: AgentConfig,
142}
143
144impl Agent {
145    /// Load a checkpoint directory containing `model.safetensors`, `encoder/`, `tokenizer/`
146    /// and `rl_agent_config.json`.
147    pub fn from_dir(dir: impl AsRef<Path>, opts: Options) -> Result<Self> {
148        let dir = dir.as_ref();
149        let device = match opts.device {
150            Some(d) => d,
151            None => default_device(),
152        };
153        // candle's ModernBert builds its attention mask as f32 unconditionally, so a f16 backbone
154        // fails with a dtype mismatch inside the first attention block. Weights are f16 on disk and
155        // are upcast on load; peak memory is therefore about 2.4 GB.
156        let dtype = opts.dtype.unwrap_or(DType::F32);
157        if dtype != DType::F32 {
158            anyhow::bail!(
159                "only f32 is supported: candle's ModernBert forces an f32 attention mask, \
160                 so a {dtype:?} backbone fails inside the first attention block"
161            );
162        }
163
164        let cfg = AgentConfig::load(dir.join("rl_agent_config.json"))?;
165        let enc_cfg = EncoderConfig::load(dir.join("encoder/config.json"))?;
166
167        let tok_path = dir.join("tokenizer/tokenizer.json");
168        let tokenizer = Tokenizer::from_file(&tok_path)
169            .map_err(|e| anyhow::anyhow!("loading tokenizer {}: {e}", tok_path.display()))?;
170        let specials = resolve_specials(&tokenizer)?;
171
172        let weights = dir.join("model.safetensors");
173        let tensors = candle_core::safetensors::load(&weights, &device)
174            .with_context(|| format!("loading weights {}", weights.display()))?;
175        // The checkpoint stores the backbone under `encoder.*`; candle's ModernBert expects `model.*`.
176        let tensors: HashMap<String, Tensor> = tensors
177            .into_iter()
178            .map(|(k, v)| match k.strip_prefix("encoder.") {
179                Some(rest) => (format!("model.{rest}"), v),
180                None => (k, v),
181            })
182            .collect();
183        let vb = VarBuilder::from_tensors(tensors, dtype, &device);
184
185        let head_layers = if cfg.head_layers == 0 {
186            2
187        } else {
188            cfg.head_layers
189        };
190        let model = model::DecisionModel::load(vb, &enc_cfg, head_layers, 2, device, dtype)?;
191
192        Ok(Self {
193            model,
194            tokenizer,
195            specials,
196            cfg,
197        })
198    }
199
200    /// Answer every question about `state` in a single batched forward pass.
201    pub fn system_one(
202        &self,
203        state: &Value,
204        questions: &Vec<(String, Question)>,
205    ) -> Result<Response> {
206        if questions.is_empty() {
207            return Ok(Response {
208                model: self.model_name(),
209                answers: Map::new(),
210                usage: Usage {
211                    input_tokens: 0,
212                    output_tokens: 0,
213                },
214            });
215        }
216
217        let mut encoded = Vec::with_capacity(questions.len());
218        for (qid, q) in questions {
219            let e = question::build_sequence(
220                &self.tokenizer,
221                &self.specials,
222                state,
223                q,
224                self.cfg.max_len,
225                self.cfg.head_max_len,
226            )
227            .with_context(|| format!("question {qid:?}"))?;
228            encoded.push(e);
229        }
230
231        let n = encoded.len();
232        let l = encoded.iter().map(|e| e.ids.len()).max().unwrap();
233        let kmax = encoded.iter().map(|e| e.markers.len()).max().unwrap();
234
235        let mut ids = vec![self.specials.pad; n * l];
236        let mut att = vec![0u32; n * l];
237        let mut mpos = vec![0u32; n * kmax];
238        let mut mmask = vec![vec![false; kmax]; n];
239        let mut qtypes = Vec::with_capacity(n);
240        let mut input_tokens = 0usize;
241
242        for (i, (e, (_, q))) in encoded.iter().zip(questions).enumerate() {
243            ids[i * l..i * l + e.ids.len()].copy_from_slice(&e.ids);
244            for j in 0..e.ids.len() {
245                att[i * l + j] = 1;
246            }
247            input_tokens += e.ids.len();
248            for (j, m) in e.markers.iter().enumerate() {
249                mpos[i * kmax + j] = *m as u32;
250                mmask[i][j] = true;
251            }
252            qtypes.push(q.qtype.index() as u32);
253        }
254
255        let dev = &self.model.device;
256        let input_ids = Tensor::from_vec(ids, (n, l), dev)?;
257        let attention_mask = Tensor::from_vec(att, (n, l), dev)?;
258        let marker_pos = Tensor::from_vec(mpos, (n, kmax), dev)?;
259        let qtype = Tensor::from_vec(qtypes, n, dev)?;
260
261        let out = self
262            .model
263            .forward(&input_ids, &attention_mask, &marker_pos, &mmask, &qtype)?;
264
265        let mut answers = Map::new();
266        for (r, ((qid, q), e)) in questions.iter().zip(&encoded).enumerate() {
267            let k = e.markers.len();
268            let t = self.cfg.temperature_for(q.qtype.index(), k);
269            let z: Vec<f32> = out.logits[r][..k].iter().map(|v| v / t).collect();
270            let p = model::stable_softmax(&z);
271            let meta = Meta {
272                act_probability: json_f32(round4(&out.act[r][0])),
273            };
274
275            let answer = match q.qtype {
276                QType::Choice => {
277                    let keys = q.choice_keys()?;
278                    let best = argmax(&p);
279                    Answer::Choice {
280                        choice: keys[best].clone(),
281                        probabilities: keys
282                            .iter()
283                            .cloned()
284                            .zip(p.iter().map(|v| json_f32(round4(v))))
285                            .collect(),
286                        confidence: json_f32(round4(&confidence_from_probs(&p, k))),
287                        rl_agent: meta,
288                    }
289                }
290                QType::Score => {
291                    let score: f32 = p.iter().enumerate().map(|(i, v)| i as f32 * v).sum();
292                    Answer::Score {
293                        score: json_f32(round4(&score)),
294                        legend: q
295                            .score_levels()
296                            .into_iter()
297                            .enumerate()
298                            .map(|(i, c)| (i.to_string(), Value::String(c)))
299                            .collect(),
300                        probabilities: p
301                            .iter()
302                            .enumerate()
303                            .map(|(i, v)| (i.to_string(), json_f32(round4(v))))
304                            .collect(),
305                        confidence: json_f32(round4(&confidence_from_probs(&p, k))),
306                        rl_agent: meta,
307                    }
308                }
309                QType::Noul => Answer::Noul {
310                    noul: json_f32(round4(&p[1])),
311                    rl_agent: meta,
312                },
313            };
314            answers.insert(qid.clone(), serde_json::to_value(answer)?);
315        }
316
317        Ok(Response {
318            model: self.model_name(),
319            answers,
320            usage: Usage {
321                input_tokens,
322                output_tokens: 0,
323            },
324        })
325    }
326
327    /// The exact token sequence for one question, annotated by segment. Useful for showing
328    /// where the option markers land and whether the state was truncated.
329    pub fn render_prompt(&self, state: &Value, q: &Question) -> Result<PromptView> {
330        let e = question::build_sequence(
331            &self.tokenizer,
332            &self.specials,
333            state,
334            q,
335            self.cfg.max_len,
336            self.cfg.head_max_len,
337        )?;
338        let options = q.render_options()?;
339
340        let mut tokens = Vec::with_capacity(e.ids.len());
341        for (i, id) in e.ids.iter().enumerate() {
342            let text = self
343                .tokenizer
344                .id_to_token(*id)
345                .unwrap_or_else(|| format!("<{id}>"))
346                .replace('\u{0120}', " ")
347                .replace('\u{010a}', "\\n");
348            let marker = e.markers.binary_search(&i).is_ok();
349            // [CLS] head [SEP] (marker opt)* [SEP] state [SEP]
350            let segment =
351                if marker || i == 0 || i == e.head_sep || i == e.opts_sep || i + 1 == e.ids.len() {
352                    Segment::Special
353                } else if i < e.head_sep {
354                    Segment::Head
355                } else if i < e.opts_sep {
356                    Segment::Option
357                } else {
358                    Segment::State
359                };
360            // Markers and option text both belong to the option whose marker most recently opened.
361            let option = if i >= *e.markers.first().unwrap_or(&usize::MAX) && i < e.opts_sep {
362                e.markers.iter().rposition(|m| *m <= i)
363            } else {
364                None
365            };
366            tokens.push(TokenView {
367                text,
368                id: *id,
369                segment,
370                marker,
371                option,
372            });
373        }
374
375        // The state fits when the sequence came in under budget.
376        let state_truncated = e.ids.len() >= self.cfg.max_len;
377
378        Ok(PromptView {
379            total_tokens: e.ids.len(),
380            tokens,
381            options,
382            markers: e.markers,
383            max_len: self.cfg.max_len,
384            state_truncated,
385        })
386    }
387
388    /// The device the model is resident on.
389    pub fn device(&self) -> &Device {
390        &self.model.device
391    }
392
393    fn model_name(&self) -> String {
394        if self.cfg.model_name.is_empty() {
395            "rl-agent".to_string()
396        } else {
397            self.cfg.model_name.clone()
398        }
399    }
400}
401
402fn default_device() -> Device {
403    if let Ok(d) = Device::new_cuda(0) {
404        return d;
405    }
406    if let Ok(d) = Device::new_metal(0) {
407        return d;
408    }
409    Device::Cpu
410}
411
412fn resolve_specials(tok: &Tokenizer) -> Result<SpecialIds> {
413    let id = |t: &str| -> Result<u32> {
414        tok.token_to_id(t)
415            .ok_or_else(|| anyhow::anyhow!("tokenizer has no {t} token"))
416    };
417    Ok(SpecialIds {
418        cls: id("[CLS]")?,
419        sep: id("[SEP]")?,
420        mask: id("[MASK]")?,
421        pad: id("[PAD]")?,
422        mask_text: "[MASK]".to_string(),
423    })
424}
425
426fn argmax(p: &[f32]) -> usize {
427    p.iter()
428        .enumerate()
429        .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
430        .map(|(i, _)| i)
431        .unwrap_or(0)
432}
433
434/// Serialize an f32 without the noise of its f64 widening (0.9641 rather than 0.9641000032424927).
435fn json_f32(v: f32) -> Value {
436    serde_json::Number::from_f64(format!("{v}").parse::<f64>().unwrap_or(v as f64))
437        .map(Value::Number)
438        .unwrap_or(Value::Null)
439}
440
441fn round4(v: &f32) -> f32 {
442    (v * 10_000.0).round() / 10_000.0
443}
444
445/// 1 - normalized entropy of the answer distribution.
446fn confidence_from_probs(p: &[f32], k: usize) -> f32 {
447    if k < 2 {
448        return 1.0;
449    }
450    let ent: f32 = -p[..k]
451        .iter()
452        .map(|v| v * v.clamp(1e-12, 1.0).ln())
453        .sum::<f32>();
454    1.0 - ent / (k as f32).ln()
455}