use crate::config::Florence2Config;
use crate::weight_source::CloningWeightSource;
use crate::weights::lang as lk;
use anyhow::{Result, anyhow};
use rayon::prelude::*;
use rlx_core::flow_util::compile_built;
use rlx_core::weight_map::WeightMap;
use rlx_runtime::{CompiledGraph, Device};
use std::collections::HashMap;
pub struct Florence2Model {
cfg: Florence2Config,
device: Device,
weights: WeightMap,
shared_table: Vec<f32>,
final_logits_bias: Vec<f32>,
vision: Option<CompiledGraph>,
encoders: HashMap<usize, CompiledGraph>,
decoders: HashMap<usize, CompiledGraph>,
embed_scratch: Vec<f32>,
img_size: usize,
}
impl Florence2Model {
pub fn config(&self) -> &Florence2Config {
&self.cfg
}
pub fn device(&self) -> Device {
self.device
}
pub fn load(
weights_path: &std::path::Path,
cfg: Florence2Config,
device: Device,
) -> Result<Self> {
rlx_core::validate_standard_device("florence2", device)?;
let mut weights = if weights_path.is_dir() {
WeightMap::from_safetensors_dir(weights_path)?
} else {
WeightMap::from_file(
weights_path
.to_str()
.ok_or_else(|| anyhow!("non-UTF8 weights path"))?,
)?
};
let (shared_table, _) = weights.take(lk::SHARED)?;
let final_logits_bias = weights
.take(lk::FINAL_LOGITS_BIAS)
.map(|(d, _)| d)
.unwrap_or_else(|_| vec![0.0; cfg.vocab_size]);
Ok(Self {
cfg,
device,
weights,
shared_table,
final_logits_bias,
vision: None,
encoders: HashMap::new(),
decoders: HashMap::new(),
embed_scratch: Vec::new(),
img_size: 768,
})
}
pub fn encode_image(&mut self, pixel_values: &[f32], img_size: usize) -> Result<Vec<f32>> {
if self.vision.is_none() || self.img_size != img_size {
let mut src = CloningWeightSource(&self.weights);
let built =
crate::flow::build_vision_built(&self.cfg, &mut src, 1, img_size, img_size)?;
self.vision = Some(compile_built(built, self.device)?);
self.img_size = img_size;
}
let g = self.vision.as_mut().unwrap();
let out = g.run(&[("pixel", pixel_values)]);
out.into_iter()
.next()
.ok_or_else(|| anyhow!("vision graph produced no output"))
}
pub fn debug_vision(&self, pixel: &[f32], img_size: usize, device: Device) -> Result<Vec<f32>> {
let mut src = CloningWeightSource(&self.weights);
let built = crate::flow::build_vision_built(&self.cfg, &mut src, 1, img_size, img_size)?;
let mut g = compile_built(built, device)?;
Ok(g.run(&[("pixel", pixel)]).into_iter().next().unwrap())
}
pub fn embed_text(&self, token_ids: &[u32]) -> Vec<f32> {
let d = self.cfg.text.d_model;
let scale = self.cfg.embed_scale();
let mut out = vec![0f32; token_ids.len() * d];
for (i, &tok) in token_ids.iter().enumerate() {
let src = (tok as usize) * d;
let dst = i * d;
for j in 0..d {
out[dst + j] = self.shared_table[src + j] * scale;
}
}
out
}
pub fn merge_embeds(&self, image_features: &[f32], text_embeds: &[f32]) -> Vec<f32> {
let mut merged = Vec::with_capacity(image_features.len() + text_embeds.len());
merged.extend_from_slice(image_features);
merged.extend_from_slice(text_embeds);
merged
}
pub fn encode(&mut self, inputs_embeds: &[f32], seq: usize) -> Result<Vec<f32>> {
if !self.encoders.contains_key(&seq) {
let mut src = CloningWeightSource(&self.weights);
let built = crate::flow::build_encoder_built(&self.cfg, &mut src, 1, seq)?;
self.encoders
.insert(seq, compile_built(built, self.device)?);
}
let g = self.encoders.get_mut(&seq).unwrap();
let out = g.run(&[("inputs_embeds", inputs_embeds)]);
out.into_iter()
.next()
.ok_or_else(|| anyhow!("encoder graph produced no output"))
}
pub fn decode_logits(
&mut self,
token_ids: &[u32],
encoder_hidden: &[f32],
enc_seq: usize,
cap: usize,
) -> Result<Vec<f32>> {
let d = self.cfg.text.d_model;
let cur = token_ids.len();
debug_assert!(cur >= 1 && cur <= cap);
let cap = bucket_len(cur).min(cap).max(cur);
let key = cap * 1_000_000 + enc_seq;
if !self.decoders.contains_key(&key) {
let mut src = CloningWeightSource(&self.weights);
let built =
crate::flow::build_decoder_hidden_built(&self.cfg, &mut src, 1, cap, enc_seq)?;
self.decoders
.insert(key, compile_built(built, self.device)?);
}
let pad = self.cfg.text.pad_token_id as usize;
self.embed_scratch.resize(cap * d, 0.0);
for i in 0..cap {
let tok = if i < cur { token_ids[i] as usize } else { pad };
let src = tok * d;
self.embed_scratch[i * d..(i + 1) * d]
.copy_from_slice(&self.shared_table[src..src + d]);
}
let hidden = {
let g = self.decoders.get_mut(&key).unwrap();
g.run(&[
("decoder_inputs_embeds", self.embed_scratch.as_slice()),
("encoder_hidden", encoder_hidden),
])
.into_iter()
.next()
.ok_or_else(|| anyhow!("decoder graph produced no output"))?
};
let row = &hidden[(cur - 1) * d..cur * d];
Ok(self.lm_head(row))
}
pub fn decoder_logits(
&mut self,
token_ids: &[u32],
encoder_hidden: &[f32],
enc_seq: usize,
) -> Result<Vec<f32>> {
self.decode_logits(token_ids, encoder_hidden, enc_seq, token_ids.len())
}
fn lm_head(&self, hidden_row: &[f32]) -> Vec<f32> {
let d = self.cfg.text.d_model;
let vocab = self.cfg.vocab_size;
let table = &self.shared_table;
let bias = &self.final_logits_bias;
let mut logits = vec![0f32; vocab];
logits.par_iter_mut().enumerate().for_each(|(v, out)| {
let row = &table[v * d..v * d + d];
let mut acc = 0f32;
for j in 0..d {
acc += hidden_row[j] * row[j];
}
*out = acc + bias[v];
});
logits
}
}
fn bucket_len(n: usize) -> usize {
const BUCKETS: [usize; 11] = [4, 8, 16, 24, 32, 48, 64, 96, 128, 192, 256];
for &b in &BUCKETS {
if b >= n {
return b;
}
}
n.next_power_of_two()
}