use std::cell::RefCell;
use std::f32::consts::FRAC_PI_2;
use crate::ops;
use crate::shortlist::Shortlist;
use crate::spm::SpmVocab;
use crate::weights::{Config, Weights};
const EPS: f32 = 1e-6;
pub struct Engine {
weights: Weights,
src_vocab: SpmVocab,
trg_vocab: SpmVocab,
config: Config,
pe_freq: Vec<f32>,
pe_offs: Vec<f32>,
shortlist: Option<Shortlist>,
shared_vocab: bool,
#[cfg(feature = "threads")]
threads: usize,
}
#[cfg(feature = "threads")]
const _: fn() = || {
fn assert_sync<T: Sync>() {}
assert_sync::<Engine>();
};
thread_local! {
static LOGITS_SCRATCH: RefCell<Vec<f32>> = const { RefCell::new(Vec::new()) };
}
pub struct Timing {
pub encode_ms: f64,
pub first_token_ms: f64,
pub decode_ms: f64,
pub out_tokens: usize,
}
pub struct BlockTiming {
pub encode_ms: f64,
pub first_token_ms: f64,
pub decode_ms: f64,
pub sentences: usize,
pub src_tokens: usize,
pub tokens: usize,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Phase {
EncodeStart,
DecodeStart,
FirstToken,
DecodeEnd,
}
pub struct BlockCounts {
pub sentences: usize,
pub src_tokens: usize,
pub tokens: usize,
}
pub struct BatchedContext {
pub data: Vec<f32>,
pub batch: usize,
pub seq: usize,
pub dim: usize,
pub lens: Vec<usize>,
}
impl BatchedContext {
pub fn sentence(&self, b: usize) -> &[f32] {
let stride = self.seq * self.dim;
&self.data[b * stride..b * stride + self.lens[b] * self.dim]
}
}
impl Engine {
pub fn load(
model_path: impl AsRef<std::path::Path>,
src_vocab_path: impl AsRef<std::path::Path>,
trg_vocab_path: impl AsRef<std::path::Path>,
) -> Result<Engine, String> {
let shared = src_vocab_path.as_ref() == trg_vocab_path.as_ref();
let weights = Weights::load(model_path)?;
let src_vocab = SpmVocab::load(src_vocab_path).map_err(|e| e.to_string())?;
let trg_vocab = SpmVocab::load(trg_vocab_path).map_err(|e| e.to_string())?;
let mut engine = Engine::new(weights, src_vocab, trg_vocab);
engine.shared_vocab = shared;
Ok(engine)
}
#[cfg(feature = "mmap")]
pub fn load_mmapped(
model_path: impl AsRef<std::path::Path>,
src_vocab_path: impl AsRef<std::path::Path>,
trg_vocab_path: impl AsRef<std::path::Path>,
) -> Result<Engine, String> {
let shared = src_vocab_path.as_ref() == trg_vocab_path.as_ref();
let weights = Weights::load_mmapped(model_path)?;
let src_vocab = SpmVocab::load(src_vocab_path).map_err(|e| e.to_string())?;
let trg_vocab = SpmVocab::load(trg_vocab_path).map_err(|e| e.to_string())?;
let mut engine = Engine::new(weights, src_vocab, trg_vocab);
engine.shared_vocab = shared;
Ok(engine)
}
pub fn from_bytes(model: &[u8], src_vocab: &[u8], trg_vocab: &[u8]) -> Result<Engine, String> {
let shared = src_vocab == trg_vocab;
let weights = Weights::from_bytes(model)?;
let src_vocab = SpmVocab::from_bytes(src_vocab);
let trg_vocab = SpmVocab::from_bytes(trg_vocab);
let mut engine = Engine::new(weights, src_vocab, trg_vocab);
engine.shared_vocab = shared;
Ok(engine)
}
pub fn with_shortlist(mut self, shortlist: Shortlist) -> Engine {
self.shortlist = Some(shortlist);
self
}
pub fn with_shortlist_bytes(self, shortlist: &[u8]) -> Engine {
self.with_shortlist(Shortlist::from_bytes(shortlist))
}
pub fn new(weights: Weights, src_vocab: SpmVocab, trg_vocab: SpmVocab) -> Engine {
let config = weights.config();
let d = config.dim_emb;
let t = d / 2;
let mut pe_freq = vec![0.0f32; d];
let mut pe_offs = vec![0.0f32; d];
for c in 0..d {
pe_freq[c] = 1e-4f32.powf((c % t) as f32 / (t as f32 - 1.0));
pe_offs[c] = (c / t) as f32 * FRAC_PI_2;
}
Engine {
weights,
src_vocab,
trg_vocab,
config,
pe_freq,
pe_offs,
shortlist: None,
shared_vocab: true,
#[cfg(feature = "threads")]
threads: 1,
}
}
#[cfg(feature = "threads")]
pub fn with_threads(mut self, n: usize) -> Engine {
self.threads = match n {
0 => std::thread::available_parallelism()
.map(|v| v.get())
.unwrap_or(1),
n => n,
};
self
}
pub fn src_ids(&self, text: &str) -> Vec<u32> {
self.src_vocab.encode_with_eos(text)
}
pub fn translate(&self, text: &str) -> String {
let src_ids = self.src_vocab.encode_with_eos(text);
let out_ids = self.greedy(&src_ids);
self.trg_vocab.decode(&out_ids)
}
pub fn translate_segmented(&self, text: &str, seg: &dyn crate::segment::Segmenter) -> String {
let spans = seg.sentences(text);
let outputs: Vec<String> = spans
.iter()
.map(|s| self.translate_within_window(s.of(text)))
.collect();
crate::segment::reassemble(text, &spans, &outputs)
}
pub fn translate_long(&self, text: &str) -> String {
#[cfg(feature = "icu-segmenter")]
{
self.translate_segmented(text, &crate::segment::IcuSegmenter::new())
}
#[cfg(not(feature = "icu-segmenter"))]
{
self.translate_segmented(text, &crate::segment::BasicSegmenter)
}
}
fn translate_within_window(&self, sentence: &str) -> String {
let mut ids = self.src_vocab.encode(sentence);
let eos = self.src_vocab.eos_id();
if ids.len() <= crate::segment::MAX_SOURCE_TOKENS {
ids.push(eos);
return self.trg_vocab.decode(&self.greedy(&ids));
}
let mut out = String::new();
for slice in ids.chunks(crate::segment::MAX_SOURCE_TOKENS) {
let mut piece = slice.to_vec();
piece.push(eos);
out.push_str(&self.trg_vocab.decode(&self.greedy(&piece)));
}
out
}
pub fn translate_timed(&self, text: &str) -> (String, Timing) {
use std::time::Instant;
let d = self.config.dim_emb;
let src_ids = self.src_vocab.encode_with_eos(text);
let seq = src_ids.len();
let t_enc = Instant::now();
let context = self.encode(&src_ids);
let encode_ms = t_enc.elapsed().as_secs_f64() * 1e3;
let eos = self.trg_vocab.eos_id();
let mut cells = vec![vec![0.0f32; d]; self.config.dec_depth + 1];
let max_len = ((2.0 * seq as f32).ceil() as usize + 4).min(256);
let candidates = self
.shortlist
.as_ref()
.map(|s| s.candidates(&src_ids, self.shared_vocab));
let t_dec = Instant::now();
let mut first_token_ms = 0.0;
let mut out = Vec::new();
let mut prev = eos;
for step in 0..max_len {
let top = self.decode_step(prev, step, &context, seq, &mut cells);
let next = self.project_argmax(&top, candidates.as_deref());
if step == 0 {
first_token_ms = t_dec.elapsed().as_secs_f64() * 1e3;
}
if next == eos {
break;
}
out.push(next);
prev = next;
}
let decode_ms = t_dec.elapsed().as_secs_f64() * 1e3;
let timing = Timing {
encode_ms,
first_token_ms,
decode_ms,
out_tokens: out.len(),
};
(self.trg_vocab.decode(&out), timing)
}
pub fn greedy(&self, src_ids: &[u32]) -> Vec<u32> {
let d = self.config.dim_emb;
let seq = src_ids.len();
let context = self.encode(src_ids);
let eos = self.trg_vocab.eos_id();
let mut cells = vec![vec![0.0f32; d]; self.config.dec_depth + 1];
let max_len = ((2.0 * seq as f32).ceil() as usize + 4).min(256);
let candidates = self
.shortlist
.as_ref()
.map(|s| s.candidates(src_ids, self.shared_vocab));
let mut out = Vec::new();
let mut prev = eos; for step in 0..max_len {
let top = self.decode_step(prev, step, &context, seq, &mut cells);
let next = self.project_argmax(&top, candidates.as_deref());
if next == eos {
break;
}
out.push(next);
prev = next;
}
out
}
fn select_active(
&self,
active: &[usize],
tops: &[f32],
cands: &[Option<Vec<u32>>],
eos: u32,
prev: &mut [u32],
out: &mut [Vec<u32>],
done: &mut [bool],
) {
let d = self.config.dim_emb;
let n = active.len();
if self.shortlist.is_none() {
let vocab = self.weights.output_vocab();
LOGITS_SCRATCH.with_borrow_mut(|logits| {
self.weights.full_logits_batch_into(tops, n, logits);
for (i, &b) in active.iter().enumerate() {
let next = argmax(&logits[i * vocab..(i + 1) * vocab]);
if next == eos {
done[b] = true;
} else {
out[b].push(next);
prev[b] = next;
}
}
});
} else {
for (i, &b) in active.iter().enumerate() {
let next = self.project_argmax(&tops[i * d..(i + 1) * d], cands[b].as_deref());
if next == eos {
done[b] = true;
} else {
out[b].push(next);
prev[b] = next;
}
}
}
}
fn project_argmax(&self, h: &[f32], candidates: Option<&[u32]>) -> u32 {
match (candidates, self.weights.output_wemb_qmult()) {
(Some(cands), Some(qwemb)) => self.project_int8(h, cands, qwemb),
(Some(cands), None) => argmax_restricted(&self.project(h), cands),
(None, _) => argmax(&self.project(h)),
}
}
fn project_int8(&self, h: &[f32], candidates: &[u32], qwemb: f32) -> u32 {
let d = self.config.dim_emb;
let n = candidates.len();
let qa = self.weights.output_qa();
let unquant = 1.0 / (qa * qwemb);
let bias_full = self
.weights
.f32("decoder_ff_logit_out_b")
.unwrap_or_else(|| vec![0.0; self.weights.output_vocab()]);
let mut b_transposed = vec![0i8; n * d];
let mut raw_bias = vec![0.0f32; n];
for (j, &c) in candidates.iter().enumerate() {
self.weights
.output_wemb_int8_row(c, &mut b_transposed[j * d..(j + 1) * d]);
raw_bias[j] = bias_full[c as usize];
}
let prepared = ops::prepare_bias(&b_transposed, n, d, &raw_bias, unquant);
let a = ops::prepare_a(h, qa);
let logits = ops::intgemm_affine(&a, 1, d, &b_transposed, n, unquant, &prepared);
let mut best = 0usize;
for j in 1..n {
if logits[j] > logits[best] {
best = j;
}
}
candidates[best]
}
pub fn encode(&self, src_ids: &[u32]) -> Vec<f32> {
let seq = src_ids.len();
let mut x = self.embed(src_ids, 0, Side::Source);
for layer in 1..=self.config.enc_depth {
x = self.encoder_layer(layer, &x, seq);
}
x
}
fn encoder_layer(&self, layer: usize, x: &[f32], seq: usize) -> Vec<f32> {
let p = format!("encoder_l{layer}");
let attn = self.multihead(&format!("{p}_self"), x, x, seq, seq);
let x = self.postnorm(&attn, x, seq, &format!("{p}_self_Wo"));
self.ffn(&format!("{p}_ffn"), &x, seq)
}
pub fn encode_batch(&self, sentences: &[Vec<u32>]) -> BatchedContext {
let d = self.config.dim_emb;
let batch = sentences.len();
let seq = sentences.iter().map(Vec::len).max().unwrap_or(0);
let lens: Vec<usize> = sentences.iter().map(Vec::len).collect();
let mut x = vec![0.0f32; batch * seq * d];
for (b, ids) in sentences.iter().enumerate() {
let base = b * seq * d;
self.embed_into(ids, 0, Side::Source, &mut x[base..base + ids.len() * d]);
}
for layer in 1..=self.config.enc_depth {
x = self.encoder_layer_batched(layer, &x, batch, seq, &lens);
}
BatchedContext {
data: x,
batch,
seq,
dim: d,
lens,
}
}
fn encoder_layer_batched(
&self,
layer: usize,
x: &[f32],
batch: usize,
seq: usize,
lens: &[usize],
) -> Vec<f32> {
let p = format!("encoder_l{layer}");
let rows = batch * seq;
let attn = self.multihead_batched(&format!("{p}_self"), x, x, batch, seq, seq, lens);
let x = self.postnorm(&attn, x, rows, &format!("{p}_self_Wo"));
self.ffn(&format!("{p}_ffn"), &x, rows)
}
pub fn decode_step(
&self,
prev_id: u32,
pos: usize,
context: &[f32],
seq: usize,
cells: &mut [Vec<f32>],
) -> Vec<f32> {
let mut x = self.embed(&[prev_id], pos, Side::Target);
for layer in 1..=self.config.dec_depth {
let p = format!("decoder_l{layer}");
let cand = self.weights.affine(&format!("{p}_rnn_W"), &x, 1, None); let gate =
self.weights
.affine(&format!("{p}_rnn_Wf"), &x, 1, Some(&format!("{p}_rnn_bf"))); let c = ops::highway(&cells[layer], &cand, &gate);
cells[layer] = c.clone();
let h = ops::relu(&c);
let x_self = self.postnorm(&h, &x, 1, &format!("{p}_rnn_ffn"));
let attn = self.multihead(&format!("{p}_context"), &x_self, context, 1, seq);
let x_ctx = self.postnorm(&attn, &x_self, 1, &format!("{p}_context_Wo"));
x = self.ffn(&format!("{p}_ffn"), &x_ctx, 1);
}
x
}
fn decode_step_batch(
&self,
active: &[usize],
prev: &[u32],
pos: usize,
ctx: &BatchedContext,
cross_kv: &[(Vec<f32>, Vec<f32>)],
cells: &mut [Vec<f32>],
) -> Vec<f32> {
let d = self.config.dim_emb;
let n = active.len();
let mut x = vec![0.0f32; n * d];
for (i, &b) in active.iter().enumerate() {
self.embed_into(&[prev[b]], pos, Side::Target, &mut x[i * d..(i + 1) * d]);
}
for layer in 1..=self.config.dec_depth {
let p = format!("decoder_l{layer}");
let cand = self.weights.affine(&format!("{p}_rnn_W"), &x, n, None);
let gate =
self.weights
.affine(&format!("{p}_rnn_Wf"), &x, n, Some(&format!("{p}_rnn_bf")));
let mut cell_prev = vec![0.0f32; n * d];
for (i, &b) in active.iter().enumerate() {
cell_prev[i * d..(i + 1) * d].copy_from_slice(&cells[layer][b * d..(b + 1) * d]);
}
let c = ops::highway(&cell_prev, &cand, &gate);
for (i, &b) in active.iter().enumerate() {
cells[layer][b * d..(b + 1) * d].copy_from_slice(&c[i * d..(i + 1) * d]);
}
let h = ops::relu(&c);
let x_self = self.postnorm(&h, &x, n, &format!("{p}_rnn_ffn"));
let (k, v) = &cross_kv[layer - 1];
let attn = self.attend_cross(
&format!("{p}_context"),
&x_self,
k,
v,
active,
ctx.seq,
&ctx.lens,
);
let x_ctx = self.postnorm(&attn, &x_self, n, &format!("{p}_context_Wo"));
x = self.ffn(&format!("{p}_ffn"), &x_ctx, n);
}
x
}
#[allow(clippy::too_many_arguments)]
fn attend_cross(
&self,
prefix: &str,
q_in: &[f32],
k: &[f32],
v: &[f32],
active: &[usize],
kv_len: usize,
kv_lens: &[usize],
) -> Vec<f32> {
let d = self.config.dim_emb;
let h = self.config.heads;
let dk = d / h;
let scale = 1.0 / (dk as f32).sqrt();
let n = active.len();
let q = self.weights.affine(
&format!("{prefix}_Wq"),
q_in,
n,
Some(&format!("{prefix}_bq")),
);
let mut joined = vec![0.0f32; n * d];
let mut scores = vec![0.0f32; kv_len];
for (i, &b) in active.iter().enumerate() {
let klen = kv_lens[b];
for head in 0..h {
let off = head * dk;
let qh = &q[i * d + off..i * d + off + dk];
for j in 0..kv_len {
if j < klen {
let kh = &k[(b * kv_len + j) * d + off..(b * kv_len + j) * d + off + dk];
scores[j] = qh.iter().zip(kh).map(|(&a, &b)| a * b).sum::<f32>() * scale;
} else {
scores[j] = f32::NEG_INFINITY;
}
}
ops::softmax_in_place(&mut scores, 1, kv_len);
let out = &mut joined[i * d + off..i * d + off + dk];
for (j, &w) in scores.iter().enumerate() {
if w == 0.0 {
continue;
}
let vh = &v[(b * kv_len + j) * d + off..(b * kv_len + j) * d + off + dk];
for c in 0..dk {
out[c] += w * vh[c];
}
}
}
}
self.weights.affine(
&format!("{prefix}_Wo"),
&joined,
n,
Some(&format!("{prefix}_bo")),
)
}
fn cross_attn_kv(&self, ctx: &BatchedContext) -> Vec<(Vec<f32>, Vec<f32>)> {
(1..=self.config.dec_depth)
.map(|layer| {
self.project_kv(
&format!("decoder_l{layer}_context"),
&ctx.data,
ctx.batch * ctx.seq,
)
})
.collect()
}
pub fn greedy_batch(&self, sentences: &[Vec<u32>]) -> Vec<Vec<u32>> {
#[cfg(feature = "threads")]
{
if self.threads > 1 && sentences.len() > 1 {
let n = self.threads.min(sentences.len());
let chunk = sentences.len().div_ceil(n);
let mut out: Vec<Vec<u32>> = Vec::with_capacity(sentences.len());
std::thread::scope(|s| {
let handles: Vec<_> = sentences
.chunks(chunk)
.map(|c| s.spawn(move || self.greedy_batch_serial(c)))
.collect();
for h in handles {
out.extend(h.join().expect("greedy_batch worker panicked"));
}
});
return out;
}
}
self.greedy_batch_serial(sentences)
}
fn greedy_batch_serial(&self, sentences: &[Vec<u32>]) -> Vec<Vec<u32>> {
let d = self.config.dim_emb;
let batch = sentences.len();
let ctx = self.encode_batch(sentences);
let eos = self.trg_vocab.eos_id();
let max_len: Vec<usize> = sentences
.iter()
.map(|s| ((2.0 * s.len() as f32).ceil() as usize + 4).min(256))
.collect();
let cap = max_len.iter().copied().max().unwrap_or(0);
let cands: Vec<Option<Vec<u32>>> = sentences
.iter()
.map(|s| {
self.shortlist
.as_ref()
.map(|sl| sl.candidates(s, self.shared_vocab))
})
.collect();
let mut cells = vec![vec![0.0f32; batch * d]; self.config.dec_depth + 1];
let mut prev = vec![eos; batch];
let mut out = vec![Vec::new(); batch];
let mut done = vec![false; batch];
let cross_kv = self.cross_attn_kv(&ctx);
for step in 0..cap {
let active = active_rows(&done, &max_len, step);
if active.is_empty() {
break;
}
let tops = self.decode_step_batch(&active, &prev, step, &ctx, &cross_kv, &mut cells);
self.select_active(&active, &tops, &cands, eos, &mut prev, &mut out, &mut done);
}
out
}
pub fn translate_batch(&self, texts: &[&str]) -> Vec<String> {
let ids: Vec<Vec<u32>> = texts
.iter()
.map(|t| self.src_vocab.encode_with_eos(t))
.collect();
self.greedy_batch(&ids)
.iter()
.map(|o| self.trg_vocab.decode(o))
.collect()
}
pub fn translate_batch_timed(&self, texts: &[&str]) -> (Vec<String>, BlockTiming) {
use std::time::Instant;
let d = self.config.dim_emb;
let sentences: Vec<Vec<u32>> = texts
.iter()
.map(|t| self.src_vocab.encode_with_eos(t))
.collect();
let batch = sentences.len();
let t_enc = Instant::now();
let ctx = self.encode_batch(&sentences);
let encode_ms = t_enc.elapsed().as_secs_f64() * 1e3;
let eos = self.trg_vocab.eos_id();
let max_len: Vec<usize> = sentences
.iter()
.map(|s| ((2.0 * s.len() as f32).ceil() as usize + 4).min(256))
.collect();
let cap = max_len.iter().copied().max().unwrap_or(0);
let cands: Vec<Option<Vec<u32>>> = sentences
.iter()
.map(|s| {
self.shortlist
.as_ref()
.map(|sl| sl.candidates(s, self.shared_vocab))
})
.collect();
let mut cells = vec![vec![0.0f32; batch * d]; self.config.dec_depth + 1];
let mut prev = vec![eos; batch];
let mut out = vec![Vec::new(); batch];
let mut done = vec![false; batch];
let t_dec = Instant::now();
let cross_kv = self.cross_attn_kv(&ctx);
let mut first_token_ms = 0.0;
for step in 0..cap {
let active = active_rows(&done, &max_len, step);
if active.is_empty() {
break;
}
let tops = self.decode_step_batch(&active, &prev, step, &ctx, &cross_kv, &mut cells);
if step == 0 {
first_token_ms = t_dec.elapsed().as_secs_f64() * 1e3;
}
self.select_active(&active, &tops, &cands, eos, &mut prev, &mut out, &mut done);
}
let decode_ms = t_dec.elapsed().as_secs_f64() * 1e3;
let timing = BlockTiming {
encode_ms,
first_token_ms,
decode_ms,
sentences: batch,
src_tokens: sentences.iter().map(Vec::len).sum(),
tokens: out.iter().map(Vec::len).sum(),
};
(
out.iter().map(|o| self.trg_vocab.decode(o)).collect(),
timing,
)
}
pub fn translate_batch_phased(
&self,
texts: &[&str],
mut on_phase: impl FnMut(Phase),
) -> (Vec<String>, BlockCounts) {
let d = self.config.dim_emb;
let sentences: Vec<Vec<u32>> = texts
.iter()
.map(|t| self.src_vocab.encode_with_eos(t))
.collect();
let batch = sentences.len();
on_phase(Phase::EncodeStart);
let ctx = self.encode_batch(&sentences);
on_phase(Phase::DecodeStart);
let eos = self.trg_vocab.eos_id();
let max_len: Vec<usize> = sentences
.iter()
.map(|s| ((2.0 * s.len() as f32).ceil() as usize + 4).min(256))
.collect();
let cap = max_len.iter().copied().max().unwrap_or(0);
let cands: Vec<Option<Vec<u32>>> = sentences
.iter()
.map(|s| {
self.shortlist
.as_ref()
.map(|sl| sl.candidates(s, self.shared_vocab))
})
.collect();
let mut cells = vec![vec![0.0f32; batch * d]; self.config.dec_depth + 1];
let mut prev = vec![eos; batch];
let mut out = vec![Vec::new(); batch];
let mut done = vec![false; batch];
let cross_kv = self.cross_attn_kv(&ctx);
for step in 0..cap {
let active = active_rows(&done, &max_len, step);
if active.is_empty() {
break;
}
let tops = self.decode_step_batch(&active, &prev, step, &ctx, &cross_kv, &mut cells);
if step == 0 {
on_phase(Phase::FirstToken);
}
self.select_active(&active, &tops, &cands, eos, &mut prev, &mut out, &mut done);
}
on_phase(Phase::DecodeEnd);
let counts = BlockCounts {
sentences: batch,
src_tokens: sentences.iter().map(Vec::len).sum(),
tokens: out.iter().map(Vec::len).sum(),
};
(
out.iter().map(|o| self.trg_vocab.decode(o)).collect(),
counts,
)
}
pub fn project(&self, h: &[f32]) -> Vec<f32> {
self.weights.full_logits(h)
}
fn multihead(
&self,
prefix: &str,
q_in: &[f32],
kv_in: &[f32],
q_len: usize,
kv_len: usize,
) -> Vec<f32> {
let d = self.config.dim_emb;
let h = self.config.heads;
let dk = d / h;
let scale = 1.0 / (dk as f32).sqrt();
let q = self.weights.affine(
&format!("{prefix}_Wq"),
q_in,
q_len,
Some(&format!("{prefix}_bq")),
);
let k = self.weights.affine(
&format!("{prefix}_Wk"),
kv_in,
kv_len,
Some(&format!("{prefix}_bk")),
);
let v = self.weights.affine(
&format!("{prefix}_Wv"),
kv_in,
kv_len,
Some(&format!("{prefix}_bv")),
);
let mut joined = vec![0.0f32; q_len * d];
let mut scores = vec![0.0f32; kv_len];
for head in 0..h {
let off = head * dk;
for i in 0..q_len {
let qh = &q[i * d + off..i * d + off + dk];
for j in 0..kv_len {
let kh = &k[j * d + off..j * d + off + dk];
let dot: f32 = qh.iter().zip(kh).map(|(&a, &b)| a * b).sum();
scores[j] = dot * scale;
}
ops::softmax_in_place(&mut scores, 1, kv_len);
let out = &mut joined[i * d + off..i * d + off + dk];
for (j, &w) in scores.iter().enumerate() {
let vh = &v[j * d + off..j * d + off + dk];
for c in 0..dk {
out[c] += w * vh[c];
}
}
}
}
self.weights.affine(
&format!("{prefix}_Wo"),
&joined,
q_len,
Some(&format!("{prefix}_bo")),
)
}
#[allow(clippy::too_many_arguments)]
fn multihead_batched(
&self,
prefix: &str,
q_in: &[f32],
kv_in: &[f32],
batch: usize,
q_len: usize,
kv_len: usize,
kv_lens: &[usize],
) -> Vec<f32> {
let (k, v) = self.project_kv(prefix, kv_in, batch * kv_len);
self.attend_batched(prefix, q_in, &k, &v, batch, q_len, kv_len, kv_lens)
}
fn project_kv(&self, prefix: &str, kv_in: &[f32], rows_kv: usize) -> (Vec<f32>, Vec<f32>) {
let k = self.weights.affine(
&format!("{prefix}_Wk"),
kv_in,
rows_kv,
Some(&format!("{prefix}_bk")),
);
let v = self.weights.affine(
&format!("{prefix}_Wv"),
kv_in,
rows_kv,
Some(&format!("{prefix}_bv")),
);
(k, v)
}
#[allow(clippy::too_many_arguments)]
fn attend_batched(
&self,
prefix: &str,
q_in: &[f32],
k: &[f32],
v: &[f32],
batch: usize,
q_len: usize,
kv_len: usize,
kv_lens: &[usize],
) -> Vec<f32> {
let d = self.config.dim_emb;
let h = self.config.heads;
let dk = d / h;
let scale = 1.0 / (dk as f32).sqrt();
let rows_q = batch * q_len;
let q = self.weights.affine(
&format!("{prefix}_Wq"),
q_in,
rows_q,
Some(&format!("{prefix}_bq")),
);
let mut joined = vec![0.0f32; rows_q * d];
let mut scores = vec![0.0f32; kv_len];
for b in 0..batch {
let klen = kv_lens[b];
for head in 0..h {
let off = head * dk;
for i in 0..q_len {
let qh = &q[(b * q_len + i) * d + off..(b * q_len + i) * d + off + dk];
for j in 0..kv_len {
if j < klen {
let kh =
&k[(b * kv_len + j) * d + off..(b * kv_len + j) * d + off + dk];
scores[j] =
qh.iter().zip(kh).map(|(&a, &b)| a * b).sum::<f32>() * scale;
} else {
scores[j] = f32::NEG_INFINITY;
}
}
ops::softmax_in_place(&mut scores, 1, kv_len);
let out =
&mut joined[(b * q_len + i) * d + off..(b * q_len + i) * d + off + dk];
for (j, &w) in scores.iter().enumerate() {
if w == 0.0 {
continue;
}
let vh = &v[(b * kv_len + j) * d + off..(b * kv_len + j) * d + off + dk];
for c in 0..dk {
out[c] += w * vh[c];
}
}
}
}
}
self.weights.affine(
&format!("{prefix}_Wo"),
&joined,
rows_q,
Some(&format!("{prefix}_bo")),
)
}
fn ffn(&self, prefix: &str, x: &[f32], seq: usize) -> Vec<f32> {
let hidden = self.weights.affine(
&format!("{prefix}_W1"),
x,
seq,
Some(&format!("{prefix}_b1")),
);
let hidden = ops::relu(&hidden);
let inner = self.config.dim_ffn;
let rows = hidden.len() / inner;
debug_assert_eq!(rows, seq);
let out = self.weights.affine(
&format!("{prefix}_W2"),
&hidden,
seq,
Some(&format!("{prefix}_b2")),
);
self.postnorm(&out, x, seq, &format!("{prefix}_ffn"))
}
fn postnorm(&self, branch: &[f32], residual: &[f32], rows: usize, ln: &str) -> Vec<f32> {
let d = self.config.dim_emb;
let sum: Vec<f32> = branch.iter().zip(residual).map(|(&a, &b)| a + b).collect();
let (gamma, beta) = self
.weights
.layer_norm(ln)
.unwrap_or_else(|| panic!("missing {ln}_ln_scale"));
ops::layer_normalization(&sum, gamma, beta, rows, d, EPS)
}
fn embed(&self, ids: &[u32], start: usize, side: Side) -> Vec<f32> {
let mut out = vec![0.0f32; ids.len() * self.config.dim_emb];
self.embed_into(ids, start, side, &mut out);
out
}
fn embed_into(&self, ids: &[u32], start: usize, side: Side, out: &mut [f32]) {
let d = self.config.dim_emb;
let scale = (d as f32).sqrt();
for (t, &id) in ids.iter().enumerate() {
let dst = &mut out[t * d..(t + 1) * d];
match side {
Side::Source => self.weights.src_embed_row_into(id, dst),
Side::Target => self.weights.trg_embed_row_into(id, dst),
}
let pos = (start + t) as f32;
for c in 0..d {
dst[c] = scale * dst[c] + (pos * self.pe_freq[c] + self.pe_offs[c]).sin();
}
}
}
}
pub enum Translation {
Direct(Engine),
Pivot {
pivot: String,
first: Engine,
second: Engine,
},
}
impl Translation {
pub fn translate(&self, text: &str) -> String {
match self {
Translation::Direct(engine) => engine.translate(text),
Translation::Pivot { first, second, .. } => second.translate(&first.translate(text)),
}
}
pub fn translate_long(&self, text: &str) -> String {
match self {
Translation::Direct(engine) => engine.translate_long(text),
Translation::Pivot { first, second, .. } => {
second.translate_long(&first.translate_long(text))
}
}
}
pub fn pivot(&self) -> Option<&str> {
match self {
Translation::Direct(_) => None,
Translation::Pivot { pivot, .. } => Some(pivot),
}
}
}
#[derive(Clone, Copy)]
enum Side {
Source,
Target,
}
fn active_rows(done: &[bool], max_len: &[usize], step: usize) -> Vec<usize> {
(0..done.len())
.filter(|&b| !done[b] && step < max_len[b])
.collect()
}
fn argmax(v: &[f32]) -> u32 {
let mut best = 0usize;
for i in 1..v.len() {
if v[i] > v[best] {
best = i;
}
}
best as u32
}
fn argmax_restricted(logits: &[f32], candidates: &[u32]) -> u32 {
let mut best = candidates[0];
let mut best_val = logits[best as usize];
for &c in &candidates[1..] {
let val = logits[c as usize];
if val > best_val {
best_val = val;
best = c;
}
}
best
}