1pub 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#[derive(Debug, Clone, Default)]
42pub struct Options {
43 pub device: Option<Device>,
45 pub dtype: Option<DType>,
47}
48
49#[derive(Debug, Clone, Serialize)]
51#[serde(tag = "type", rename_all = "lowercase")]
52pub enum Answer {
53 Choice {
54 choice: String,
56 probabilities: Map<String, Value>,
57 confidence: Value,
59 rl_agent: Meta,
60 },
61 Score {
62 score: Value,
64 legend: Map<String, Value>,
65 probabilities: Map<String, Value>,
66 confidence: Value,
67 rl_agent: Meta,
68 },
69 Noul {
70 noul: Value,
72 rl_agent: Meta,
73 },
74}
75
76#[derive(Debug, Clone, Serialize)]
77pub struct Meta {
78 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#[derive(Debug, Clone, Copy, Serialize)]
97#[serde(rename_all = "lowercase")]
98pub enum Segment {
99 Special,
101 Head,
103 Option,
105 State,
107}
108
109#[derive(Debug, Clone, Serialize)]
111pub struct TokenView {
112 pub text: String,
114 pub id: u32,
115 pub segment: Segment,
116 pub marker: bool,
118 #[serde(skip_serializing_if = "Option::is_none")]
120 pub option: Option<usize>,
121}
122
123#[derive(Debug, Clone, Serialize)]
125pub struct PromptView {
126 pub tokens: Vec<TokenView>,
127 pub options: Vec<String>,
129 pub markers: Vec<usize>,
131 pub total_tokens: usize,
132 pub max_len: usize,
133 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 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 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 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 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 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 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 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 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 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
434fn 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
445fn 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}