use candle_core::{DType, Device, Result as CandleResult, Tensor};
use candle_transformers::generation::{LogitsProcessor, Sampling};
use ffai_core::engine::Decoding;
use candle_nn::VarBuilder;
use candle_transformers::models::llama;
#[derive(Debug, Clone, Default, PartialEq)]
pub struct DecodeTrace {
pub prefill_ms: f64,
pub prompt_tokens: usize,
pub steps_ms: Vec<f64>,
}
impl DecodeTrace {
#[must_use]
pub fn decode_ms(&self) -> f64 {
self.steps_ms.iter().sum()
}
#[must_use]
pub fn tokens_per_sec(&self) -> f64 {
let ms = self.decode_ms();
if ms <= 0.0 {
return 0.0;
}
self.steps_ms.len() as f64 / (ms / 1e3)
}
}
pub struct TextDecoder {
model: Option<llama::Llama>,
cache: llama::Cache,
pristine: llama::Cache,
config: llama::Config,
device: Device,
ours: Option<crate::text::TextTower>,
}
fn use_candle_text() -> bool {
use std::sync::atomic::{AtomicU8, Ordering};
static C: AtomicU8 = AtomicU8::new(u8::MAX);
match C.load(Ordering::Relaxed) {
u8::MAX => {
let on = std::env::var("FFAI_ARGUS_CANDLE_TEXT").is_ok_and(|v| v == "1");
C.store(u8::from(on), Ordering::Relaxed);
on
}
v => v == 1,
}
}
impl TextDecoder {
pub fn load(weights: &std::path::Path, config_json: &str, device: &Device) -> Result<Self, String> {
let vb = unsafe {
VarBuilder::from_mmaped_safetensors(std::slice::from_ref(&weights), DType::F32, device)
}
.map_err(|e| format!("load {}: {e}", weights.display()))?;
Self::load_vb(vb, config_json, device)
}
pub fn load_vb(
vb: VarBuilder<'static>,
config_json: &str,
device: &Device,
) -> Result<Self, String> {
let config = text_config_from_json(config_json)?;
let raw = vb.clone();
let vb = vb.rename_f(|name: &str| {
if let Some(rest) = name.strip_prefix("model.") {
format!("model.text_model.{rest}")
} else {
name.to_string()
}
});
let cache = llama::Cache::new(true, DType::F32, &config, device)
.map_err(|e| format!("kv cache: {e}"))?;
let load_candle = |vb: VarBuilder<'static>| {
llama::Llama::load(vb, &config).map_err(|e| format!("text tower: {e}"))
};
let v: serde_json::Value =
serde_json::from_str(config_json).map_err(|e| format!("config.json: {e}"))?;
let t = v.get("text_config").unwrap_or(&v);
let gu = |k: &str, d: u64| t.get(k).and_then(serde_json::Value::as_u64).unwrap_or(d);
let gfl = |k: &str, d: f64| t.get(k).and_then(serde_json::Value::as_f64).unwrap_or(d);
let heads = gu("num_attention_heads", 9) as usize;
let hidden = gu("hidden_size", 576) as usize;
let cfg = crate::text::Cfg {
layers: gu("num_hidden_layers", 30) as usize,
hidden,
heads,
kv_heads: gu("num_key_value_heads", 3) as usize,
head_dim: hidden / heads.max(1),
inter: gu("intermediate_size", 1536) as usize,
eps: gfl("rms_norm_eps", 1e-5),
rope_theta: gfl("rope_theta", 100_000.0) as f32,
max_pos: gu("max_position_embeddings", 8192) as usize,
};
let ours = if use_candle_text() {
None
} else {
crate::text::TextTower::load(&raw, cfg, device).ok()
};
let model = if ours.is_some() { None } else { Some(load_candle(vb)?) };
Ok(Self {
model,
pristine: cache.clone(),
cache,
config,
device: device.clone(),
ours,
})
}
pub fn load_reference(
weights: &std::path::Path,
config_json: &str,
device: &Device,
) -> Result<Self, String> {
let raw = unsafe {
VarBuilder::from_mmaped_safetensors(
std::slice::from_ref(&weights),
DType::F32,
device,
)
}
.map_err(|e| format!("load {}: {e}", weights.display()))?;
let config = text_config_from_json(config_json)?;
let vb = raw.rename_f(|name: &str| {
if let Some(rest) = name.strip_prefix("model.") {
format!("model.text_model.{rest}")
} else {
name.to_string()
}
});
let cache = llama::Cache::new(true, DType::F32, &config, device)
.map_err(|e| format!("kv cache: {e}"))?;
let model = llama::Llama::load(vb, &config).map_err(|e| format!("text tower: {e}"))?;
Ok(Self {
model: Some(model),
pristine: cache.clone(),
cache,
config,
device: device.clone(),
ours: None,
})
}
pub fn reset(&mut self) {
self.cache = self.pristine.clone();
if let Some(t) = self.ours.as_mut() {
t.reset();
}
}
pub fn forward_embeds(&mut self, embeds: &Tensor, index_pos: usize) -> CandleResult<Tensor> {
if let Some(t) = self.ours.as_mut() {
return t.forward(embeds, index_pos);
}
let Some(m) = self.model.as_ref() else {
return Err(candle_core::Error::Msg("no text tower loaded".into()));
};
m.forward_input_embed(embeds, index_pos, &mut self.cache)
}
pub fn embed(&self, ids: &Tensor) -> CandleResult<Tensor> {
if let Some(t) = self.ours.as_ref() {
return t.embed(ids);
}
let Some(m) = self.model.as_ref() else {
return Err(candle_core::Error::Msg("no text tower loaded".into()));
};
m.embed(ids)
}
pub fn generate_greedy(
&mut self,
inputs_embeds: &Tensor,
max_new_tokens: usize,
stop_ids: &[u32],
) -> CandleResult<Vec<u32>> {
self.generate(inputs_embeds, max_new_tokens, stop_ids, &Decoding::Greedy, None)
}
pub fn generate(
&mut self,
inputs_embeds: &Tensor,
max_new_tokens: usize,
stop_ids: &[u32],
decoding: &Decoding,
repetition_penalty: Option<f32>,
) -> CandleResult<Vec<u32>> {
self.generate_traced(
inputs_embeds,
max_new_tokens,
stop_ids,
decoding,
repetition_penalty,
None,
)
}
pub fn generate_traced(
&mut self,
inputs_embeds: &Tensor,
max_new_tokens: usize,
stop_ids: &[u32],
decoding: &Decoding,
repetition_penalty: Option<f32>,
mut trace: Option<&mut DecodeTrace>,
) -> CandleResult<Vec<u32>> {
self.reset();
let mut sampler = match decoding {
Decoding::Greedy => None,
Decoding::Sampled {
temperature,
top_p,
top_k,
seed,
} => {
let t = f64::from(*temperature);
Some(LogitsProcessor::from_sampling(
*seed,
match (top_k, top_p) {
(Some(k), Some(p)) => Sampling::TopKThenTopP {
k: *k,
p: f64::from(*p),
temperature: t,
},
(Some(k), None) => Sampling::TopK {
k: *k,
temperature: t,
},
(None, Some(p)) => Sampling::TopP {
p: f64::from(*p),
temperature: t,
},
(None, None) => Sampling::All { temperature: t },
},
))
}
};
let (_b, prefill_len, _d) = inputs_embeds.dims3()?;
let t_prefill = crate::clock::Instant::now();
let mut logits = self.forward_embeds(inputs_embeds, 0)?;
if let Some(t) = trace.as_deref_mut() {
t.prefill_ms = t_prefill.elapsed().as_secs_f64() * 1e3;
t.prompt_tokens = prefill_len;
}
let mut out: Vec<u32> = Vec::with_capacity(max_new_tokens);
for pos in prefill_len..prefill_len + max_new_tokens {
let t_step = crate::clock::Instant::now();
let mut step = logits.flatten_all()?;
if let Some(p) = repetition_penalty {
#[allow(clippy::float_cmp)]
if p != 1.0 && !out.is_empty() {
step = candle_transformers::utils::apply_repeat_penalty(&step, p, &out)?;
}
}
let next = match sampler.as_mut() {
Some(s) => s.sample(&step)?,
None => argmax(&step)?,
};
if stop_ids.contains(&next) {
break;
}
out.push(next);
let ids = Tensor::new(&[next], &self.device)?.unsqueeze(0)?;
let emb = self.embed(&ids)?;
logits = self.forward_embeds(&emb, pos)?;
if let Some(t) = trace.as_deref_mut() {
t.steps_ms.push(t_step.elapsed().as_secs_f64() * 1e3);
}
}
Ok(out)
}
#[must_use]
pub const fn config(&self) -> &llama::Config {
&self.config
}
}
fn argmax(logits: &Tensor) -> CandleResult<u32> {
let v = logits.flatten_all()?.to_vec1::<f32>()?;
let mut best = 0usize;
let mut best_v = f32::NEG_INFINITY;
for (i, &x) in v.iter().enumerate() {
if x > best_v {
best_v = x;
best = i;
}
}
u32::try_from(best).map_err(|e| candle_core::Error::Msg(format!("token id overflow: {e}")))
}
pub fn text_config_from_json(config_json: &str) -> Result<llama::Config, String> {
let v: serde_json::Value =
serde_json::from_str(config_json).map_err(|e| format!("config.json: {e}"))?;
let tc = v
.get("text_config")
.ok_or("config.json has no text_config")?;
let get_usize = |k: &str| -> Result<usize, String> {
tc.get(k)
.and_then(serde_json::Value::as_u64)
.map(|x| x as usize)
.ok_or_else(|| format!("text_config has no {k}"))
};
Ok(llama::Config {
hidden_size: get_usize("hidden_size")?,
intermediate_size: get_usize("intermediate_size")?,
vocab_size: get_usize("vocab_size")?,
num_hidden_layers: get_usize("num_hidden_layers")?,
num_attention_heads: get_usize("num_attention_heads")?,
num_key_value_heads: get_usize("num_key_value_heads")
.or_else(|_| get_usize("num_attention_heads"))?,
rms_norm_eps: tc
.get("rms_norm_eps")
.and_then(serde_json::Value::as_f64)
.unwrap_or(1e-5),
rope_theta: tc
.get("rope_theta")
.and_then(serde_json::Value::as_f64)
.unwrap_or(10000.0) as f32,
bos_token_id: tc.get("bos_token_id").and_then(serde_json::Value::as_u64).map(|x| x as u32),
eos_token_id: tc
.get("eos_token_id")
.and_then(serde_json::Value::as_u64)
.map(|x| llama::LlamaEosToks::Single(x as u32)),
rope_scaling: None,
max_position_embeddings: get_usize("max_position_embeddings").unwrap_or(8192),
tie_word_embeddings: tc
.get("tie_word_embeddings")
.and_then(serde_json::Value::as_bool)
.unwrap_or(false),
use_flash_attn: false,
})
}