use crate::Engine;
use crate::model::Model;
impl Model {
pub fn forward(
&self,
e: &Engine,
tokens: &[u32],
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
let cfg = &self.cfg;
let n_embd = cfg.n_embd as usize;
let n_head = cfg.n_head as usize;
let n_head_kv = cfg.n_head_kv as usize;
let head_dim = cfg.head_dim_k as usize;
let t = tokens.len();
let eps = cfg.rms_eps;
let scale = 1.0 / (head_dim as f32).sqrt();
let pos: Vec<i32> = (0..t as i32).collect();
let pos_d = e.htod_i32(&pos)?;
let mut x = self.embed_tokens(e, tokens)?;
let max_block = self.max_moe_block();
for (il, layer) in self.layers.iter().enumerate() {
let mut h = e.zeros(t * n_embd)?;
e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
let q_out = layer.wq.out_features(); let k_out = layer.wk.out_features();
let mut q = e.matmul(&layer.wq, &h, t)?;
let mut k = e.matmul(&layer.wk, &h, t)?;
let v = e.matmul(&layer.wv, &h, t)?;
if let Some(qn) = &layer.q_norm {
let mut qn_out = e.zeros(t * q_out)?;
e.rms_norm(&q, qn.float_data(), &mut qn_out, head_dim, n_head * t, eps)?;
q = qn_out;
}
if let Some(kn) = &layer.k_norm {
let mut kn_out = e.zeros(t * k_out)?;
e.rms_norm(
&k,
kn.float_data(),
&mut kn_out,
head_dim,
n_head_kv * t,
eps,
)?;
k = kn_out;
}
e.rope_neox(
&mut q,
&pos_d,
head_dim,
cfg.rope_dim_count as usize,
n_head,
t,
cfg.rope_freq_base,
1.0,
)?;
e.rope_neox(
&mut k,
&pos_d,
head_dim,
cfg.rope_dim_count as usize,
n_head_kv,
t,
cfg.rope_freq_base,
1.0,
)?;
let mut attn = e.zeros(t * q_out)?;
e.sdpa_naive(
&q, &k, &v, &mut attn, head_dim, n_head, n_head_kv, t, t, scale, true,
)?;
let o = e.matmul(&layer.wo, &attn, t)?;
let mut x1 = e.zeros(t * n_embd)?;
e.add(&x, &o, &mut x1, t * n_embd)?;
let mut z = e.zeros(t * n_embd)?;
e.rms_norm(&x1, layer.ffn_norm.float_data(), &mut z, n_embd, t, eps)?;
let down = match &layer.ffn {
crate::hybrid::Ffn::Dense {
ffn_gate,
ffn_up,
ffn_down,
} => {
let n_ff = ffn_gate.out_features();
let gate = e.matmul(ffn_gate, &z, t)?;
let up = e.matmul(ffn_up, &z, t)?;
let mut act = e.zeros(t * n_ff)?;
crate::hybrid::HybridModel::ffn_act(e, cfg, &gate, &up, &mut act, t * n_ff)?;
e.matmul(ffn_down, &act, t)?
}
crate::hybrid::Ffn::Moe(m) => {
crate::hybrid::HybridModel::moe_ffn(e, m, &z, t, cfg, il as u16, max_block)?
}
};
let mut x2 = e.zeros(t * n_embd)?;
e.add(&x1, &down, &mut x2, t * n_embd)?;
x = x2;
}
let mut hn = e.zeros(t * n_embd)?;
e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
let logits = e.matmul(&self.output, &hn, t)?;
let host = e.dtoh(&logits)?;
Ok(host)
}
pub fn forward_last(
&self,
e: &Engine,
tokens: &[u32],
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
let all = self.forward(e, tokens)?;
let n_vocab = self.output.out_features();
let t = tokens.len();
Ok(all[(t - 1) * n_vocab..t * n_vocab].to_vec())
}
}
pub fn argmax(logits: &[f32]) -> usize {
let mut best = 0;
let mut bv = f32::NEG_INFINITY;
for (i, &v) in logits.iter().enumerate() {
if v > bv {
bv = v;
best = i;
}
}
best
}
pub fn top2(logits: &[f32]) -> (usize, f32, usize, f32) {
let (mut i1, mut v1, mut i2, mut v2) = (0usize, f32::NEG_INFINITY, 0usize, f32::NEG_INFINITY);
for (i, &v) in logits.iter().enumerate() {
if v > v1 {
i2 = i1;
v2 = v1;
i1 = i;
v1 = v;
} else if v > v2 {
i2 = i;
v2 = v;
}
}
(i1, v1, i2, v2)
}
#[derive(Debug, PartialEq, Eq, Clone, Copy)]
pub enum PrimeGateClass {
Match,
NearTieFlip,
Structured,
}
pub struct PrimeGateVerdict {
pub tw_argmax: usize,
pub bp_argmax: usize,
pub tw_margin: f32,
pub bp_margin: f32,
pub maxdiff: f32,
pub class: PrimeGateClass,
}
fn env_f32(key: &str, default: f32) -> f32 {
std::env::var(key)
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(default)
}
pub fn prime_gate_verdict(tokenwise: &[f32], batched: &[f32]) -> PrimeGateVerdict {
let (t1, tv1, _, tv2) = top2(tokenwise);
let (b1, bv1, _, bv2) = top2(batched);
let maxdiff = tokenwise
.iter()
.zip(batched)
.map(|(a, b)| (a - b).abs())
.fold(0.0f32, f32::max);
let maxdiff_bound = env_f32("MEMRA_PRIME_GATE_MAXDIFF", 8.0);
let margin_bound = env_f32("MEMRA_PRIME_GATE_MARGIN", 1.0);
let tw_margin = tv1 - tv2;
let class = if !maxdiff.is_finite() || maxdiff > maxdiff_bound {
PrimeGateClass::Structured
} else if t1 == b1 {
PrimeGateClass::Match
} else if tw_margin <= margin_bound {
PrimeGateClass::NearTieFlip
} else {
PrimeGateClass::Structured
};
PrimeGateVerdict {
tw_argmax: t1,
bp_argmax: b1,
tw_margin,
bp_margin: bv1 - bv2,
maxdiff,
class,
}
}