use crate::config::{DepFormerConfig, TransformerConfig};
use crate::sampling::LogitsProcessor;
use anyhow::{Context, Result};
use ndarray::ArrayView1;
use rayon::prelude::*;
use rlx_ir::infer::GraphExt;
use rlx_ir::op::MaskKind;
use rlx_ir::{DType, Graph, NodeId, Op, Shape};
use rlx_runtime::{Device, Session};
use std::collections::HashMap;
const RMS_EPS: f32 = 1e-8;
#[derive(Debug, Clone, Copy)]
pub struct HeliumDims {
pub d_model: usize,
pub n_heads: usize,
pub head_dim: usize,
pub n_layers: usize,
pub ffn: usize, pub vocab_out: usize,
pub rope_theta: f32,
}
impl HeliumDims {
pub fn from_cfg(cfg: &TransformerConfig, vocab_out: usize) -> Self {
Self {
d_model: cfg.d_model,
n_heads: cfg.num_heads,
head_dim: cfg.d_model / cfg.num_heads,
n_layers: cfg.num_layers,
ffn: cfg.swiglu_hidden(),
vocab_out,
rope_theta: cfg.max_period as f32,
}
}
}
fn p(li: usize, name: &str) -> String {
format!("L{li}.{name}")
}
fn apply_rope(g: &mut Graph, x: NodeId, cos: NodeId, sin: NodeId, half: usize) -> NodeId {
let x1 = g.narrow_(x, 3, 0, half);
let x2 = g.narrow_(x, 3, half, half);
let x1c = g.mul(x1, cos);
let x2s = g.mul(x2, sin);
let x2c = g.mul(x2, cos);
let x1s = g.mul(x1, sin);
let r1 = g.sub(x1c, x2s);
let r2 = g.add(x2c, x1s);
g.concat_(vec![r1, r2], 3)
}
fn transpose(data: &[f32], rows: usize, cols: usize) -> Vec<f32> {
let mut out = vec![0.0f32; rows * cols];
out.par_chunks_mut(rows)
.enumerate()
.for_each(|(c, out_row)| {
for (r, slot) in out_row.iter_mut().enumerate() {
*slot = data[r * cols + c];
}
});
out
}
fn rms(g: &mut Graph, x: NodeId, name: &str, d: usize, shape: Shape, zero_beta: NodeId) -> NodeId {
let w = g.param(name, Shape::new(&[d], DType::F32));
g.add_node(
Op::RmsNorm {
axis: -1,
eps: RMS_EPS,
},
vec![x, w, zero_beta],
shape,
)
}
fn qkv_proj(
g: &mut Graph,
n1: NodeId,
li: usize,
d: usize,
heads: &[i64],
) -> (NodeId, NodeId, NodeId) {
let qw = g.param(p(li, "q"), Shape::new(&[d, d], DType::F32));
let kw = g.param(p(li, "k"), Shape::new(&[d, d], DType::F32));
let vw = g.param(p(li, "v"), Shape::new(&[d, d], DType::F32));
let q = g.mm(n1, qw);
let k = g.mm(n1, kw);
let v = g.mm(n1, vw);
let q = g.reshape_(q, heads.to_vec());
let k = g.reshape_(k, heads.to_vec());
let v = g.reshape_(v, heads.to_vec());
(q, k, v)
}
fn swiglu_block(
g: &mut Graph,
x: NodeId,
li: usize,
d: usize,
ffn: usize,
shape: Shape,
zero_beta: NodeId,
) -> NodeId {
let n2 = rms(g, x, &p(li, "n2"), d, shape, zero_beta);
let gw = g.param(p(li, "gate"), Shape::new(&[d, ffn], DType::F32));
let uw = g.param(p(li, "up"), Shape::new(&[d, ffn], DType::F32));
let gate = g.mm(n2, gw);
let up = g.mm(n2, uw);
let gate = g.silu(gate);
let h = g.mul(gate, up);
let dw = g.param(p(li, "down"), Shape::new(&[ffn, d], DType::F32));
let mlp = g.mm(h, dw);
g.add(x, mlp)
}
fn for_each_transformer_param(
weights: &HashMap<String, (Vec<f32>, Vec<usize>)>,
base: &str,
d: usize,
ffn: usize,
n_layers: usize,
mut emit: impl FnMut(String, Vec<f32>),
) -> Result<()> {
let get = |k: &str| -> Result<&Vec<f32>> {
weights
.get(k)
.map(|(v, _)| v)
.with_context(|| format!("missing weight {k}"))
};
for li in 0..n_layers {
let pre = format!("{base}transformer.layers.{li}");
let inproj = get(&format!("{pre}.self_attn.in_proj_weight"))?;
emit(p(li, "q"), transpose(&inproj[0..d * d], d, d));
emit(p(li, "k"), transpose(&inproj[d * d..2 * d * d], d, d));
emit(p(li, "v"), transpose(&inproj[2 * d * d..3 * d * d], d, d));
emit(
p(li, "o"),
transpose(get(&format!("{pre}.self_attn.out_proj.weight"))?, d, d),
);
emit(p(li, "n1"), get(&format!("{pre}.norm1.alpha"))?.clone());
emit(p(li, "n2"), get(&format!("{pre}.norm2.alpha"))?.clone());
let gate_up = get(&format!("{pre}.gating.linear_in.weight"))?;
emit(p(li, "gate"), transpose(&gate_up[0..ffn * d], ffn, d));
emit(
p(li, "up"),
transpose(&gate_up[ffn * d..2 * ffn * d], ffn, d),
);
emit(
p(li, "down"),
transpose(get(&format!("{pre}.gating.linear_out.weight"))?, d, ffn),
);
}
Ok(())
}
pub fn build_temporal_graph(dims: &HeliumDims, seq: usize) -> Graph {
let HeliumDims {
d_model: d,
n_heads: nh,
head_dim: hd,
n_layers,
ffn,
vocab_out,
..
} = *dims;
let mut g = Graph::new("moshi_temporal");
let bsd = Shape::new(&[1, seq, d], DType::F32);
let half = hd / 2;
let mut x = g.input("inputs_embeds", bsd.clone());
let cos = g.input("rope_cos", Shape::new(&[seq, half], DType::F32));
let sin = g.input("rope_sin", Shape::new(&[seq, half], DType::F32));
let cos4 = g.reshape_(cos, vec![1, seq as i64, 1, half as i64]);
let sin4 = g.reshape_(sin, vec![1, seq as i64, 1, half as i64]);
let zero_beta = g.param("zero_beta", Shape::new(&[d], DType::F32));
let heads = [1i64, seq as i64, nh as i64, hd as i64];
for li in 0..n_layers {
let n1 = rms(&mut g, x, &p(li, "n1"), d, bsd.clone(), zero_beta);
let (q, k, v) = qkv_proj(&mut g, n1, li, d, &heads);
let q = apply_rope(&mut g, q, cos4, sin4, half);
let k = apply_rope(&mut g, k, cos4, sin4, half);
let attn = g.attention_kind(
q,
k,
v,
nh,
hd,
MaskKind::Causal,
Shape::new(&[1, seq, nh, hd], DType::F32),
);
let attn = g.reshape_(attn, vec![1i64, seq as i64, d as i64]);
let ow = g.param(p(li, "o"), Shape::new(&[d, d], DType::F32));
let attn = g.mm(attn, ow);
x = g.add(x, attn);
x = swiglu_block(&mut g, x, li, d, ffn, bsd.clone(), zero_beta);
}
let xn = rms(&mut g, x, "out_norm", d, bsd, zero_beta);
let tl = g.param("text_linear", Shape::new(&[d, vocab_out], DType::F32));
let logits = g.mm(xn, tl);
g.set_outputs(vec![logits]);
g
}
pub fn build_temporal_decode_graph(dims: &HeliumDims, past: usize) -> Graph {
let HeliumDims {
d_model: d,
n_heads: nh,
head_dim: hd,
n_layers,
ffn,
vocab_out,
..
} = *dims;
let half = hd / 2;
let p1d = Shape::new(&[1, 1, d], DType::F32);
let mut g = Graph::new("moshi_temporal_decode");
let mut x = g.input("inputs_embeds", p1d.clone());
let cos = g.input("rope_cos", Shape::new(&[1, half], DType::F32));
let sin = g.input("rope_sin", Shape::new(&[1, half], DType::F32));
let cos4 = g.reshape_(cos, vec![1, 1, 1, half as i64]);
let sin4 = g.reshape_(sin, vec![1, 1, 1, half as i64]);
let zero_beta = g.param("zero_beta", Shape::new(&[d], DType::F32));
let kv_shape = Shape::new(&[1, past, nh, hd], DType::F32);
let heads1 = vec![1i64, 1, nh as i64, hd as i64];
let mut kv_outputs = Vec::with_capacity(2 * n_layers);
for li in 0..n_layers {
let n1 = rms(&mut g, x, &p(li, "n1"), d, p1d.clone(), zero_beta);
let (q, k, v) = qkv_proj(&mut g, n1, li, d, &heads1);
let q = apply_rope(&mut g, q, cos4, sin4, half);
let k = apply_rope(&mut g, k, cos4, sin4, half);
let (k_full, v_full) = if past == 0 {
(k, v)
} else {
let past_k = g.input(format!("past_k_{li}"), kv_shape.clone());
let past_v = g.input(format!("past_v_{li}"), kv_shape.clone());
(g.concat_(vec![past_k, k], 1), g.concat_(vec![past_v, v], 1))
};
kv_outputs.push(k_full);
kv_outputs.push(v_full);
let attn = g.attention_kind(
q,
k_full,
v_full,
nh,
hd,
MaskKind::Causal,
Shape::new(&[1, 1, nh, hd], DType::F32),
);
let attn = g.reshape_(attn, vec![1i64, 1, d as i64]);
let ow = g.param(p(li, "o"), Shape::new(&[d, d], DType::F32));
let attn = g.mm(attn, ow);
x = g.add(x, attn);
x = swiglu_block(&mut g, x, li, d, ffn, p1d.clone(), zero_beta);
}
let hidden = x;
let xn = rms(&mut g, x, "out_norm", d, p1d, zero_beta);
let tl = g.param("text_linear", Shape::new(&[d, vocab_out], DType::F32));
let logits = g.mm(xn, tl);
let mut outs = vec![logits, hidden];
outs.extend(kv_outputs);
g.set_outputs(outs);
g
}
pub fn bucket_decode_mask(past_seq: usize, upper: usize) -> Vec<f32> {
(0..=upper)
.map(|i| if i < past_seq || i == upper { 1.0 } else { 0.0 })
.collect()
}
pub fn build_temporal_decode_graph_bucketed(dims: &HeliumDims, upper: usize) -> Graph {
let HeliumDims {
d_model: d,
n_heads: nh,
head_dim: hd,
n_layers,
ffn,
vocab_out,
..
} = *dims;
let half = hd / 2;
let p1d = Shape::new(&[1, 1, d], DType::F32);
let mut g = Graph::new("moshi_temporal_decode_bucketed");
let mut x = g.input("inputs_embeds", p1d.clone());
let cos = g.input("rope_cos", Shape::new(&[1, half], DType::F32));
let sin = g.input("rope_sin", Shape::new(&[1, half], DType::F32));
let cos4 = g.reshape_(cos, vec![1, 1, 1, half as i64]);
let sin4 = g.reshape_(sin, vec![1, 1, 1, half as i64]);
let mask = g.input("attn_mask", Shape::new(&[1, upper + 1], DType::F32));
let zero_beta = g.param("zero_beta", Shape::new(&[d], DType::F32));
let one = g.param("kv_one", Shape::new(&[1], DType::F32));
let kv_shape = Shape::new(&[1, upper, nh, hd], DType::F32);
let heads1 = vec![1i64, 1, nh as i64, hd as i64];
let mut kv_outputs = Vec::with_capacity(2 * n_layers);
for li in 0..n_layers {
let n1 = rms(&mut g, x, &p(li, "n1"), d, p1d.clone(), zero_beta);
let (q, k, v) = qkv_proj(&mut g, n1, li, d, &heads1);
let q = apply_rope(&mut g, q, cos4, sin4, half);
let k = apply_rope(&mut g, k, cos4, sin4, half);
let new_k = g.mul(k, one);
let new_v = g.mul(v, one);
kv_outputs.push(new_k);
kv_outputs.push(new_v);
let past_k = g.input(format!("past_k_{li}"), kv_shape.clone());
let past_v = g.input(format!("past_v_{li}"), kv_shape.clone());
let k_full = g.concat_(vec![past_k, k], 1);
let v_full = g.concat_(vec![past_v, v], 1);
let attn = g.attention_(q, k_full, v_full, mask, nh, hd);
let attn = g.reshape_(attn, vec![1i64, 1, d as i64]);
let ow = g.param(p(li, "o"), Shape::new(&[d, d], DType::F32));
let attn = g.mm(attn, ow);
x = g.add(x, attn);
x = swiglu_block(&mut g, x, li, d, ffn, p1d.clone(), zero_beta);
}
let hidden = x;
let xn = rms(&mut g, x, "out_norm", d, p1d, zero_beta);
let tl = g.param("text_linear", Shape::new(&[d, vocab_out], DType::F32));
let logits = g.mm(xn, tl);
let mut outs = vec![logits, hidden];
outs.extend(kv_outputs);
g.set_outputs(outs);
g
}
pub fn temporal_params(
dims: &HeliumDims,
weights: &HashMap<String, (Vec<f32>, Vec<usize>)>,
) -> Result<HashMap<String, Vec<f32>>> {
let HeliumDims {
d_model: d,
ffn,
n_layers,
..
} = *dims;
let get = |k: &str| -> Result<&Vec<f32>> {
weights
.get(k)
.map(|(v, _)| v)
.with_context(|| format!("missing weight {k}"))
};
let mut out = HashMap::new();
out.insert("zero_beta".to_string(), vec![0.0f32; d]);
for_each_transformer_param(weights, "", d, ffn, n_layers, |name, data| {
out.insert(name, data);
})?;
out.insert("out_norm".to_string(), get("out_norm.alpha")?.clone());
out.insert(
"text_linear".to_string(),
transpose(get("text_linear.weight")?, dims.vocab_out, d),
);
Ok(out)
}
pub fn set_temporal_params(
compiled: &mut rlx_runtime::CompiledGraph,
dims: &HeliumDims,
weights: &HashMap<String, (Vec<f32>, Vec<usize>)>,
) -> Result<()> {
let HeliumDims {
d_model: d,
ffn,
n_layers,
vocab_out,
..
} = *dims;
let get = |k: &str| -> Result<&Vec<f32>> {
weights
.get(k)
.map(|(v, _)| v)
.with_context(|| format!("missing weight {k}"))
};
compiled.set_param("zero_beta", &vec![0.0f32; d]);
compiled.set_param("kv_one", &[1.0]);
for_each_transformer_param(weights, "", d, ffn, n_layers, |name, data| {
compiled.set_param(&name, &data);
})?;
compiled.set_param("out_norm", get("out_norm.alpha")?);
compiled.set_param(
"text_linear",
&transpose(get("text_linear.weight")?, vocab_out, d),
);
Ok(())
}
#[derive(Debug, Clone, Copy)]
pub struct DepDims {
pub d_main: usize, pub d_dep: usize, pub n_heads: usize,
pub head_dim: usize,
pub n_layers: usize,
pub ffn: usize,
pub audio_vocab: usize,
pub num_slices: usize,
}
impl DepDims {
pub fn from_cfg(cfg: &DepFormerConfig, d_main: usize, audio_vocab: usize) -> Self {
let t = &cfg.transformer;
Self {
d_main,
d_dep: t.d_model,
n_heads: t.num_heads,
head_dim: t.d_model / t.num_heads,
n_layers: t.num_layers,
ffn: t.swiglu_hidden(),
audio_vocab,
num_slices: cfg.num_slices,
}
}
}
pub fn build_depformer_slice_graph(dd: &DepDims, past: usize) -> Graph {
let DepDims {
d_main,
d_dep: d,
n_heads: nh,
head_dim: hd,
n_layers,
ffn,
audio_vocab,
..
} = *dd;
let p1d = Shape::new(&[1, 1, d], DType::F32);
let mut g = Graph::new("moshi_depformer_slice");
let hidden = g.input("temporal_hidden", Shape::new(&[1, 1, d_main], DType::F32));
let emb_vec = g.input("emb_vec", p1d.clone());
let zero_beta = g.param("zero_beta", Shape::new(&[d], DType::F32));
let one = g.param("kv_one", Shape::new(&[1], DType::F32));
let kv_shape = Shape::new(&[1, past, nh, hd], DType::F32);
let heads1 = vec![1i64, 1, nh as i64, hd as i64];
let lin_in = g.param("linear_in", Shape::new(&[d_main, d], DType::F32));
let proj = g.mm(hidden, lin_in);
let mut x = g.add(proj, emb_vec);
let mut kv_outputs = Vec::with_capacity(2 * n_layers);
for li in 0..n_layers {
let n1 = rms(&mut g, x, &p(li, "n1"), d, p1d.clone(), zero_beta);
let (q, k, v) = qkv_proj(&mut g, n1, li, d, &heads1);
let (k_full, v_full) = if past == 0 {
(k, v)
} else {
let past_k = g.input(format!("past_k_{li}"), kv_shape.clone());
let past_v = g.input(format!("past_v_{li}"), kv_shape.clone());
(g.concat_(vec![past_k, k], 1), g.concat_(vec![past_v, v], 1))
};
let k_out = g.mul(k_full, one);
let v_out = g.mul(v_full, one);
kv_outputs.push(k_out);
kv_outputs.push(v_out);
let attn = g.attention_kind(
q,
k_full,
v_full,
nh,
hd,
MaskKind::Causal,
Shape::new(&[1, 1, nh, hd], DType::F32),
);
let attn = g.reshape_(attn, vec![1i64, 1, d as i64]);
let ow = g.param(p(li, "o"), Shape::new(&[d, d], DType::F32));
let attn = g.mm(attn, ow);
x = g.add(x, attn);
x = swiglu_block(&mut g, x, li, d, ffn, p1d.clone(), zero_beta);
}
let lo = g.param("linear_out", Shape::new(&[d, audio_vocab], DType::F32));
let logits = g.mm(x, lo);
let mut outs = vec![logits];
outs.extend(kv_outputs);
g.set_outputs(outs);
g
}
pub fn depformer_slice_params(
dd: &DepDims,
si: usize,
weights: &HashMap<String, (Vec<f32>, Vec<usize>)>,
) -> Result<HashMap<String, Vec<f32>>> {
let DepDims {
d_main,
d_dep: d,
ffn,
n_layers,
audio_vocab,
..
} = *dd;
let get = |k: &str| -> Result<&Vec<f32>> {
weights
.get(k)
.map(|(v, _)| v)
.with_context(|| format!("missing weight {k}"))
};
let mut out = HashMap::new();
out.insert("zero_beta".to_string(), vec![0.0f32; d]);
out.insert("kv_one".to_string(), vec![1.0f32]);
out.insert(
"linear_in".to_string(),
transpose(get(&format!("depformer.{si}.linear_in.weight"))?, d, d_main),
);
out.insert(
"linear_out".to_string(),
transpose(
get(&format!("depformer.{si}.linear_out.weight"))?,
audio_vocab,
d,
),
);
for_each_transformer_param(
weights,
&format!("depformer.{si}."),
d,
ffn,
n_layers,
|name, data| {
out.insert(name, data);
},
)?;
Ok(out)
}
pub fn compile_depformer_slice(
dd: &DepDims,
weights: &HashMap<String, (Vec<f32>, Vec<usize>)>,
si: usize,
past: usize,
device: Device,
) -> Result<rlx_runtime::CompiledGraph> {
let mut c = Session::new(device).compile(build_depformer_slice_graph(dd, past));
for (name, data) in &depformer_slice_params(dd, si, weights)? {
c.set_param(name, data);
}
Ok(c)
}
pub fn depformer_slice_run(
compiled: &mut rlx_runtime::CompiledGraph,
dd: &DepDims,
weights: &HashMap<String, (Vec<f32>, Vec<usize>)>,
temporal_hidden: &[f32],
si: usize,
last_token: Option<u32>,
past_kv: &[(Vec<f32>, Vec<f32>)],
) -> Result<(Vec<f32>, Vec<(Vec<f32>, Vec<f32>)>)> {
let emb_vec = match last_token {
Some(t) => {
let (data, shape) = weights
.get(&format!("depformer.{si}.emb.weight"))
.with_context(|| format!("missing depformer.{si}.emb.weight"))?;
let row = shape[1];
data[t as usize * row..(t as usize + 1) * row].to_vec()
}
None => vec![0.0f32; dd.d_dep],
};
let mut inputs: Vec<(String, &[f32])> = vec![
("temporal_hidden".to_string(), temporal_hidden),
("emb_vec".to_string(), emb_vec.as_slice()),
];
if !past_kv.is_empty() {
for (li, (k, v)) in past_kv.iter().enumerate() {
inputs.push((format!("past_k_{li}"), k.as_slice()));
inputs.push((format!("past_v_{li}"), v.as_slice()));
}
}
let refs: Vec<(&str, &[f32])> = inputs.iter().map(|(n, dt)| (n.as_str(), *dt)).collect();
let mut it = compiled.run(&refs).into_iter();
let logits = it.next().context("depformer slice produced no logits")?;
let mut new_kv = Vec::with_capacity(dd.n_layers);
for _ in 0..dd.n_layers {
let k = it.next().context("missing new_k")?;
let v = it.next().context("missing new_v")?;
new_kv.push((k, v));
}
Ok((logits, new_kv))
}
pub fn depformer_run_slice(
dd: &DepDims,
weights: &HashMap<String, (Vec<f32>, Vec<usize>)>,
temporal_hidden: &[f32],
si: usize,
last_token: Option<u32>,
past_kv: &[(Vec<f32>, Vec<f32>)],
device: Device,
) -> Result<(Vec<f32>, Vec<(Vec<f32>, Vec<f32>)>)> {
let past = if past_kv.is_empty() {
0
} else {
past_kv[0].0.len() / (dd.n_heads * dd.head_dim)
};
let mut compiled = compile_depformer_slice(dd, weights, si, past, device)?;
depformer_slice_run(
&mut compiled,
dd,
weights,
temporal_hidden,
si,
last_token,
past_kv,
)
}
pub fn depformer_sample_rlx(
dd: &DepDims,
weights: &HashMap<String, (Vec<f32>, Vec<usize>)>,
temporal_hidden: &[f32],
text_token: Option<u32>,
forced: &[Option<u32>],
lp: &mut LogitsProcessor,
device: Device,
) -> Result<Vec<u32>> {
let mut tokens = Vec::with_capacity(dd.num_slices);
let mut last_token = text_token;
let mut past_kv: Vec<(Vec<f32>, Vec<f32>)> = Vec::new();
for si in 0..dd.num_slices {
let (logits, new_kv) = depformer_run_slice(
dd,
weights,
temporal_hidden,
si,
last_token,
&past_kv,
device,
)?;
past_kv = new_kv;
let token = lp.sample(ArrayView1::from(&logits))?;
tokens.push(token);
let next = forced.get(si).copied().flatten().unwrap_or(token);
last_token = Some(next);
}
Ok(tokens)
}
pub fn depformer_forced_logits_rlx(
dd: &DepDims,
weights: &HashMap<String, (Vec<f32>, Vec<usize>)>,
temporal_hidden: &[f32],
text_token: Option<u32>,
forced: &[u32],
inherit: bool,
device: Device,
) -> Result<Vec<Vec<f32>>> {
let mut out = Vec::with_capacity(dd.num_slices);
let mut last_token = text_token;
let mut past_kv: Vec<(Vec<f32>, Vec<f32>)> = Vec::new();
for si in 0..dd.num_slices {
let (logits, new_kv) = depformer_run_slice(
dd,
weights,
temporal_hidden,
si,
last_token,
&past_kv,
device,
)?;
past_kv = if inherit { new_kv } else { Vec::new() };
out.push(logits);
last_token = forced.get(si).copied();
}
Ok(out)
}
pub fn rope_tables(dims: &HeliumDims, seq: usize) -> (Vec<f32>, Vec<f32>) {
let half = dims.head_dim / 2;
let mut cos = vec![0.0f32; seq * half];
let mut sin = vec![0.0f32; seq * half];
for pos in 0..seq {
for i in 0..half {
let inv_freq = 1.0f32 / dims.rope_theta.powf(i as f32 / half as f32);
let f = pos as f32 * inv_freq;
cos[pos * half + i] = f.cos();
sin[pos * half + i] = f.sin();
}
}
(cos, sin)
}
pub fn temporal_logits_rlx(
dims: &HeliumDims,
weights: &HashMap<String, (Vec<f32>, Vec<usize>)>,
inputs_embeds: &[f32],
seq: usize,
device: Device,
) -> Result<Vec<f32>> {
let graph = build_temporal_graph(dims, seq);
let params = temporal_params(dims, weights)?;
let mut compiled = Session::new(device).compile(graph);
for (name, data) in ¶ms {
compiled.set_param(name, data);
}
let (cos, sin) = rope_tables(dims, seq);
let outs = compiled.run(&[
("inputs_embeds", inputs_embeds),
("rope_cos", cos.as_slice()),
("rope_sin", sin.as_slice()),
]);
outs.into_iter()
.next()
.context("temporal graph produced no output")
}
pub fn temporal_decode_step_rlx(
dims: &HeliumDims,
weights: &HashMap<String, (Vec<f32>, Vec<usize>)>,
inputs_embeds: &[f32],
past_kv: &[(Vec<f32>, Vec<f32>)],
pos: usize,
device: Device,
) -> Result<(Vec<f32>, Vec<f32>, Vec<(Vec<f32>, Vec<f32>)>)> {
let graph = build_temporal_decode_graph(dims, pos);
let params = temporal_params(dims, weights)?;
let mut compiled = Session::new(device).compile(graph);
for (name, data) in ¶ms {
compiled.set_param(name, data);
}
let half = dims.head_dim / 2;
let mut cos = vec![0.0f32; half];
let mut sin = vec![0.0f32; half];
for i in 0..half {
let inv_freq = 1.0f32 / dims.rope_theta.powf(i as f32 / half as f32);
let f = pos as f32 * inv_freq;
cos[i] = f.cos();
sin[i] = f.sin();
}
let mut inputs: Vec<(String, &[f32])> = vec![
("inputs_embeds".to_string(), inputs_embeds),
("rope_cos".to_string(), cos.as_slice()),
("rope_sin".to_string(), sin.as_slice()),
];
if pos > 0 {
for (li, (k, v)) in past_kv.iter().enumerate() {
inputs.push((format!("past_k_{li}"), k.as_slice()));
inputs.push((format!("past_v_{li}"), v.as_slice()));
}
}
let refs: Vec<(&str, &[f32])> = inputs.iter().map(|(n, d)| (n.as_str(), *d)).collect();
let mut it = compiled.run(&refs).into_iter();
let logits = it.next().context("decode produced no logits")?;
let hidden = it.next().context("decode produced no hidden")?;
let mut new_kv = Vec::with_capacity(dims.n_layers);
for _ in 0..dims.n_layers {
let k = it.next().context("missing new_k")?;
let v = it.next().context("missing new_v")?;
new_kv.push((k, v));
}
Ok((logits, hidden, new_kv))
}
pub fn temporal_decode_bucketed_rlx(
dims: &HeliumDims,
weights: &HashMap<String, (Vec<f32>, Vec<usize>)>,
inputs_embeds: &[f32],
real_past_kv: &[(Vec<f32>, Vec<f32>)],
past_seq: usize,
upper: usize,
device: Device,
) -> Result<(Vec<f32>, Vec<f32>, Vec<(Vec<f32>, Vec<f32>)>)> {
let mut compiled =
Session::new(device).compile(build_temporal_decode_graph_bucketed(dims, upper));
for (name, data) in &temporal_params(dims, weights)? {
compiled.set_param(name, data);
}
compiled.set_param("kv_one", &[1.0]);
decode_bucketed_run(
&mut compiled,
dims,
inputs_embeds,
real_past_kv,
past_seq,
upper,
)
}
pub fn decode_bucketed_run(
compiled: &mut rlx_runtime::CompiledGraph,
dims: &HeliumDims,
inputs_embeds: &[f32],
real_past_kv: &[(Vec<f32>, Vec<f32>)],
past_seq: usize,
upper: usize,
) -> Result<(Vec<f32>, Vec<f32>, Vec<(Vec<f32>, Vec<f32>)>)> {
let half = dims.head_dim / 2;
let mut cos = vec![0.0f32; half];
let mut sin = vec![0.0f32; half];
for i in 0..half {
let inv_freq = 1.0f32 / dims.rope_theta.powf(i as f32 / half as f32);
let f = past_seq as f32 * inv_freq;
cos[i] = f.cos();
sin[i] = f.sin();
}
let mask = bucket_decode_mask(past_seq, upper);
let kvw = dims.n_heads * dims.head_dim;
let pad_len = upper * kvw;
let padded: Vec<(Vec<f32>, Vec<f32>)> = (0..dims.n_layers)
.map(|li| {
let (rk, rv) = real_past_kv
.get(li)
.map(|(k, v)| (k.as_slice(), v.as_slice()))
.unwrap_or((&[], &[]));
let mut pk = rk.to_vec();
pk.resize(pad_len, 0.0);
let mut pv = rv.to_vec();
pv.resize(pad_len, 0.0);
(pk, pv)
})
.collect();
let mut inputs: Vec<(String, &[f32])> = vec![
("inputs_embeds".to_string(), inputs_embeds),
("rope_cos".to_string(), cos.as_slice()),
("rope_sin".to_string(), sin.as_slice()),
("attn_mask".to_string(), mask.as_slice()),
];
for (li, (k, v)) in padded.iter().enumerate() {
inputs.push((format!("past_k_{li}"), k.as_slice()));
inputs.push((format!("past_v_{li}"), v.as_slice()));
}
let refs: Vec<(&str, &[f32])> = inputs.iter().map(|(n, d)| (n.as_str(), *d)).collect();
let mut it = compiled.run(&refs).into_iter();
let logits = it.next().context("bucketed decode produced no logits")?;
let hidden = it.next().context("bucketed decode produced no hidden")?;
let mut new_kv = Vec::with_capacity(dims.n_layers);
for _ in 0..dims.n_layers {
let k = it.next().context("missing new_k")?;
let v = it.next().context("missing new_v")?;
new_kv.push((k, v));
}
Ok((logits, hidden, new_kv))
}