use std::collections::HashMap;
use std::path::Path;
use anyhow::Context;
use super::graph::{build_encoder_graph, build_forward_graph, ForwardSpec};
#[cfg(feature = "validation")]
use super::graph::{build_encoder_debug2_graph, build_encoder_debug_graph};
use super::prepare::{channel_wise_normalize, gather_channel_emb, prepare_tokens};
use super::weights::{
apply_params, build_forward_params, build_prepare_params, load_safetensors, ParamMap,
};
use crate::config::ModelConfig;
#[derive(Clone, Debug)]
pub struct EpochEmbedding {
pub output: Vec<f32>,
pub shape: Vec<usize>,
pub chan_pos: Vec<f32>,
pub n_channels: usize,
}
#[derive(Clone, Copy, Debug)]
pub struct RunEpochOpts {
pub normalize: bool,
}
impl Default for RunEpochOpts {
fn default() -> Self {
Self { normalize: true }
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ClassifierKind {
None,
Luna,
Linear,
Mamba,
}
impl ClassifierKind {
fn detect(raw: &ParamMap) -> Self {
if raw.contains_key("classifier.decoder_attn.in_proj_weight") {
Self::Luna
} else if raw.contains_key("classifier.fc1.weight") {
Self::Linear
} else if raw.keys().any(|k| k.starts_with("classifier.")) {
Self::Mamba
} else {
Self::None
}
}
}
pub struct LuMambaEncoder {
pub model_cfg: ModelConfig,
pub device: rlx::Device,
pub classifier_kind: ClassifierKind,
forward_params: ParamMap,
prepare_params: ParamMap,
session: rlx::Session,
forward_cache: HashMap<u64, rlx::CompiledGraph>,
encoder_cache: HashMap<u64, rlx::CompiledGraph>,
}
fn gelu_host(x: f32) -> f32 {
let s = (x as f64) * std::f64::consts::FRAC_1_SQRT_2;
0.5 * x * (1.0 + libm::erf(s) as f32)
}
impl LuMambaEncoder {
pub fn load(
config_path: &Path,
weights_path: &Path,
device: rlx::Device,
) -> anyhow::Result<(Self, f64)> {
let cfg_str = std::fs::read_to_string(config_path)
.with_context(|| format!("reading config: {}", config_path.display()))?;
let hf_val: serde_json::Value = serde_json::from_str(&cfg_str)?;
let model_cfg: ModelConfig =
serde_json::from_value(hf_val.get("model").cloned().unwrap_or(hf_val))
.context("parsing model config")?;
let t = std::time::Instant::now();
let mut raw =
load_safetensors(weights_path.to_str().context("weights path not valid UTF-8")?)?;
let classifier_kind = ClassifierKind::detect(&raw);
let mut raw_prepare = raw.clone();
let forward_params = build_forward_params(&mut raw, &model_cfg)?;
let prepare_params = build_prepare_params(&mut raw_prepare)?;
let session = rlx::Session::new(device);
let ms = t.elapsed().as_secs_f64() * 1000.0;
Ok((
Self {
model_cfg,
device,
classifier_kind,
forward_params,
prepare_params,
session,
forward_cache: HashMap::new(),
encoder_cache: HashMap::new(),
},
ms,
))
}
pub fn describe(&self) -> String {
let c = &self.model_cfg;
if c.num_classes > 0 {
format!(
"LuMamba classifier (RLX, dev={:?}) embed_dim={} classes={}",
self.device, c.embed_dim, c.num_classes,
)
} else {
format!(
"LuMamba encoder (RLX, dev={:?}) E={} Q={} blocks={} d_inner={} d_state={} patch={}",
self.device, c.embed_dim, c.num_queries, c.num_blocks,
c.d_inner(), c.d_state, c.patch_size,
)
}
}
fn spec(&self, b: usize, c: usize, t: usize) -> ForwardSpec {
let cfg = &self.model_cfg;
let s = t / cfg.patch_size;
ForwardSpec {
b,
c,
s,
bt: b * s,
d: cfg.embed_dim,
q: cfg.num_queries,
hidden: cfg.hidden_dim(),
nh_ca: cfg.num_heads,
dh_ca: cfg.cross_head_dim(),
ff_ca: cfg.ffn_cross_dim(),
patch_size: cfg.patch_size,
norm_eps: cfg.norm_eps as f32,
num_blocks: cfg.num_blocks,
d_inner: cfg.d_inner(),
d_state: cfg.d_state,
d_conv: cfg.d_conv,
dt_rank: cfg.dt_rank(),
bidir_multiply: cfg.bidir_multiply(),
num_classes: cfg.num_classes,
nh_cls: cfg.num_heads,
}
}
fn expand_queries(&self, bt: usize) -> Vec<f32> {
let q = self.model_cfg.num_queries;
let d = self.model_cfg.embed_dim;
let embed = &self.forward_params["cross_attn.query_embed"];
let flat = if embed.shape == vec![1, q, d] {
embed.data.clone()
} else {
embed.data[..q * d].to_vec()
};
let mut out = vec![0f32; bt * q * d];
for i in 0..bt {
out[i * q * d..(i + 1) * q * d].copy_from_slice(&flat);
}
out
}
fn expand_agg_query(&self, b: usize) -> Vec<f32> {
let hidden = self.model_cfg.hidden_dim();
let embed = &self.forward_params["classifier.learned_agg"];
let flat = if embed.shape == vec![1, 1, hidden] {
embed.data.clone()
} else {
embed.data[..hidden].to_vec()
};
let mut out = vec![0f32; b * hidden];
for i in 0..b {
out[i * hidden..(i + 1) * hidden].copy_from_slice(&flat);
}
out
}
fn channel_emb_slice(&self, indices: Option<&[i32]>, b: usize, c: usize) -> Option<Vec<f32>> {
let table = self.prepare_params.get("channel_emb.weight")?;
let d = self.model_cfg.embed_dim;
let vocab = table.shape[0] as i32;
let idx = indices?;
let clamped: Vec<i32> = idx.iter().map(|&i| i.clamp(0, vocab - 1)).collect();
Some(gather_channel_emb(table, &clamped, b, c, d))
}
fn cache_key(&self, b: usize, c: usize, t: usize) -> u64 {
(b as u64) << 40
| (c as u64) << 20
| (t as u64)
| ((self.model_cfg.num_classes as u64) << 60)
}
fn compiled_for(&mut self, b: usize, c: usize, t: usize) -> &mut rlx::CompiledGraph {
let key = self.cache_key(b, c, t);
if !self.forward_cache.contains_key(&key) {
let spec = self.spec(b, c, t);
let graph = build_forward_graph(&spec);
let mut compiled = self.session.compile(graph);
apply_params(&mut compiled, &self.forward_params);
self.forward_cache.insert(key, compiled);
}
self.forward_cache.get_mut(&key).expect("just inserted")
}
pub fn run_epoch(
&mut self,
signal: &[f32],
chan_pos: &[f32],
channel_indices: Option<&[i32]>,
n_channels: usize,
n_samples: usize,
) -> anyhow::Result<EpochEmbedding> {
self.run_epoch_opts(
signal,
chan_pos,
channel_indices,
n_channels,
n_samples,
RunEpochOpts::default(),
)
}
pub fn run_epoch_opts(
&mut self,
signal: &[f32],
chan_pos: &[f32],
channel_indices: Option<&[i32]>,
n_channels: usize,
n_samples: usize,
opts: RunEpochOpts,
) -> anyhow::Result<EpochEmbedding> {
let (outs, c, t) =
self.run_all_outputs(signal, chan_pos, channel_indices, n_channels, n_samples, opts)?;
let num_classes = self.model_cfg.num_classes;
let output = outs
.into_iter()
.next()
.ok_or_else(|| anyhow::anyhow!("forward graph produced no output"))?;
let shape = if num_classes > 0 { vec![num_classes] } else { vec![c, t] };
Ok(EpochEmbedding {
output,
shape,
chan_pos: chan_pos.to_vec(),
n_channels: c,
})
}
pub fn encode(
&mut self,
signal: &[f32],
chan_pos: &[f32],
channel_indices: Option<&[i32]>,
n_channels: usize,
n_samples: usize,
) -> anyhow::Result<(Vec<f32>, Vec<usize>)> {
let (mut outs, _c, t) = self.run_all_outputs(
signal,
chan_pos,
channel_indices,
n_channels,
n_samples,
RunEpochOpts::default(),
)?;
let s = t / self.model_cfg.patch_size;
let hidden = self.model_cfg.hidden_dim();
anyhow::ensure!(outs.len() >= 2, "encoder latent output missing");
let latent = outs.swap_remove(1);
Ok((latent, vec![s, hidden]))
}
fn run_all_outputs(
&mut self,
signal: &[f32],
chan_pos: &[f32],
channel_indices: Option<&[i32]>,
n_channels: usize,
n_samples: usize,
opts: RunEpochOpts,
) -> anyhow::Result<(Vec<Vec<f32>>, usize, usize)> {
let b = 1usize;
let c = n_channels;
let t = n_samples;
anyhow::ensure!(
t.is_multiple_of(self.model_cfg.patch_size),
"n_samples ({t}) must be a multiple of patch_size ({})",
self.model_cfg.patch_size
);
let patch_size = self.model_cfg.patch_size;
let embed_dim = self.model_cfg.embed_dim;
let num_classes = self.model_cfg.num_classes;
let timing = std::env::var_os("LUMAMBA_TIMING").is_some();
let t_prep = std::time::Instant::now();
let mut sig = signal.to_vec();
if opts.normalize {
channel_wise_normalize(&mut sig, c, t);
}
let ch_emb = self.channel_emb_slice(channel_indices, b, c);
let (x_tok, dec_q) = prepare_tokens(
&sig,
chan_pos,
ch_emb.as_deref(),
b,
c,
t,
patch_size,
embed_dim,
&self.prepare_params,
);
let ms_prep = t_prep.elapsed().as_secs_f64() * 1000.0;
let spec = self.spec(b, c, t);
let queries = self.expand_queries(spec.bt);
let agg_query = if num_classes > 0 { Some(self.expand_agg_query(b)) } else { None };
let compiled = self.compiled_for(b, c, t);
let mut inputs: Vec<(&str, &[f32])> =
vec![("x_tokenized", x_tok.as_slice()), ("queries", queries.as_slice())];
if let Some(ref agg) = agg_query {
inputs.push(("agg_query", agg.as_slice()));
} else {
inputs.push(("decoder_queries", dec_q.as_slice()));
}
let t_run = std::time::Instant::now();
let outs = compiled.run(&inputs);
if timing {
eprintln!(
" [timing] host-prepare {ms_prep:.1} ms | graph-run {:.1} ms",
t_run.elapsed().as_secs_f64() * 1000.0
);
}
Ok((outs, c, t))
}
pub fn run_encoder(&mut self, h_in: &[f32], s: usize) -> anyhow::Result<Vec<f32>> {
let b = 1usize;
let hidden = self.model_cfg.hidden_dim();
anyhow::ensure!(
h_in.len() == b * s * hidden,
"h_in must be [1, {s}, {hidden}] = {} elems, got {}",
b * s * hidden,
h_in.len()
);
let t = s * self.model_cfg.patch_size;
let spec = self.spec(b, 1, t); let key = (s as u64) | (1u64 << 63);
if !self.encoder_cache.contains_key(&key) {
let graph = build_encoder_graph(&spec);
let mut compiled = self.session.compile(graph);
for (name, buf) in &self.forward_params {
if name.starts_with("mamba_blocks.") || name.starts_with("norm_layers.") {
compiled.set_param(name, &buf.data);
}
}
self.encoder_cache.insert(key, compiled);
}
let compiled = self.encoder_cache.get_mut(&key).expect("just inserted");
let outs = compiled.run(&[("h_in", h_in)]);
outs.into_iter()
.next()
.ok_or_else(|| anyhow::anyhow!("encoder graph produced no output"))
}
pub fn classify_epoch(
&mut self,
signal: &[f32],
chan_pos: &[f32],
channel_indices: Option<&[i32]>,
n_channels: usize,
n_samples: usize,
) -> anyhow::Result<Vec<f32>> {
match self.classifier_kind {
ClassifierKind::Luna => {
anyhow::ensure!(
self.model_cfg.num_classes > 0,
"LUNA classifier checkpoint requires num_classes>0 in config.json"
);
let emb = self.run_epoch_opts(
signal,
chan_pos,
channel_indices,
n_channels,
n_samples,
RunEpochOpts::default(),
)?;
Ok(emb.output)
}
ClassifierKind::Linear => {
let (latent, shape) =
self.encode(signal, chan_pos, channel_indices, n_channels, n_samples)?;
let (s, hidden) = (shape[0], shape[1]);
let w = self
.prepare_params
.get("classifier.fc1.weight")
.ok_or_else(|| anyhow::anyhow!("missing classifier.fc1.weight"))?;
let b = &self.prepare_params["classifier.fc1.bias"].data;
let nc = b.len();
let mut mean = vec![0f32; hidden];
for si in 0..s {
for d in 0..hidden {
mean[d] += latent[si * hidden + d];
}
}
for m in mean.iter_mut() {
*m /= s as f32;
}
let mut logits = vec![0f32; nc];
for (cc, lg) in logits.iter_mut().enumerate() {
let mut acc = b[cc];
for d in 0..hidden {
acc += mean[d] * w.data[d * nc + cc];
}
*lg = gelu_host(acc);
}
Ok(logits)
}
ClassifierKind::Mamba => anyhow::bail!(
"checkpoint uses a MambaClassifier head — not implemented in this port yet"
),
ClassifierKind::None => anyhow::bail!(
"checkpoint has no classifier head (pretrained encoder, num_classes=0); \
provide a fine-tuned checkpoint to evaluate"
),
}
}
#[cfg(feature = "validation")]
pub fn run_encoder_debug(&mut self, h_in: &[f32], s: usize) -> anyhow::Result<Vec<Vec<f32>>> {
let t = s * self.model_cfg.patch_size;
let spec = self.spec(1, 1, t);
let graph = build_encoder_debug_graph(&spec);
let mut compiled = self.session.compile(graph);
for (name, buf) in &self.forward_params {
if name.starts_with("mamba_blocks.0.mamba_fwd") {
compiled.set_param(name, &buf.data);
}
}
Ok(compiled.run(&[("h_in", h_in)]))
}
#[cfg(feature = "validation")]
pub fn run_encoder_debug2(&mut self, h_in: &[f32], s: usize) -> anyhow::Result<Vec<Vec<f32>>> {
let t = s * self.model_cfg.patch_size;
let spec = self.spec(1, 1, t);
let graph = build_encoder_debug2_graph(&spec);
let mut compiled = self.session.compile(graph);
for (name, buf) in &self.forward_params {
if name.starts_with("mamba_blocks.0.") {
compiled.set_param(name, &buf.data);
}
}
Ok(compiled.run(&[("h_in", h_in)]))
}
pub fn run_rlx_epoch(&mut self, ep: &super::io::RlxEpoch) -> anyhow::Result<EpochEmbedding> {
self.run_epoch(
&ep.signal,
&ep.chan_pos,
ep.channel_indices.as_deref(),
ep.n_channels,
ep.n_samples,
)
}
}