use crate::tensor::Tensor;
use crate::v1::config::V1TransformerConfig;
use crate::weights::{get, param_to_tensor, WeightStore};
fn alibi_slopes(heads: usize) -> Vec<f32> {
fn slopes_power_of_2(n: usize) -> Vec<f32> {
let start = 2f32.powf(-2f32.powf(-((n as f32).log2() - 3.0)));
let ratio = start;
(0..n).map(|i| start * ratio.powi(i as i32)).collect()
}
if heads.is_power_of_two() {
return slopes_power_of_2(heads);
}
let closest = 2usize.pow((heads as f32).log2().floor() as u32);
let mut out = slopes_power_of_2(closest);
let extra: Vec<f32> = slopes_power_of_2(2 * closest)
.into_iter()
.step_by(2)
.take(heads - closest)
.collect();
out.extend(extra);
out
}
pub struct ScaleNorm {
pub g: f32,
pub dim: usize,
}
impl ScaleNorm {
pub fn forward(&self, x: &Tensor) -> Tensor {
let d = self.dim;
let scale = (d as f32).sqrt();
let gamma = self.g + 1.0;
let (b, n, _) = (x.shape[0], x.shape[1], x.shape[2]);
let mut out = vec![0.0f32; x.data.len()];
for bi in 0..b {
for ni in 0..n {
let base = (bi * n + ni) * d;
let mut norm = 0.0f32;
for i in 0..d {
norm += x.data[base + i] * x.data[base + i];
}
norm = norm.sqrt().max(1e-12);
for i in 0..d {
out[base + i] = (x.data[base + i] / norm) * scale * gamma;
}
}
}
Tensor::from_vec(out, x.shape.clone())
}
}
struct AttnLayer {
pre_norm: ScaleNorm,
q_w: Tensor,
k_w: Tensor,
v_w: Tensor,
residual_scale: Option<Tensor>,
}
struct FfLayer {
pre_norm: ScaleNorm,
w1: Tensor,
b1: Tensor,
w2: Tensor,
b2: Tensor,
residual_scale: Option<Tensor>,
}
pub struct SentenceTransformer {
pub dim: usize,
pub heads: usize,
pub head_dim: usize,
pub alibi_heads: usize,
pub alibi_slopes: Vec<f32>,
attn_layers: Vec<AttnLayer>,
ff_layers: Vec<FfLayer>,
final_norm: ScaleNorm,
}
impl SentenceTransformer {
pub fn from_config_and_weights(
cfg: &V1TransformerConfig,
dim: usize,
store: &mut WeightStore,
prefix: &str,
) -> anyhow::Result<Self> {
anyhow::ensure!(
!cfg.rotary_pos_emb,
"rotary_pos_emb is not supported in the Rust V1 transformer yet"
);
anyhow::ensure!(dim % cfg.heads == 0, "dim must divide heads");
let head_dim = dim / cfg.heads;
let alibi_heads = cfg.heads;
let mut attn_layers = Vec::new();
let mut ff_layers = Vec::new();
for layer in 0..cfg.depth {
let attn_i = layer * 2;
let ff_i = layer * 2 + 1;
let ap = format!("{prefix}layers.{attn_i}.");
let fp = format!("{prefix}layers.{ff_i}.");
attn_layers.push(AttnLayer {
pre_norm: ScaleNorm {
g: get(store, &format!("{ap}0.0.g"))
.map(|p| p.data[0])
.unwrap_or(1.0),
dim,
},
q_w: param_to_tensor(get(store, &format!("{ap}1.to_q.weight"))?),
k_w: param_to_tensor(get(store, &format!("{ap}1.to_k.weight"))?),
v_w: param_to_tensor(get(store, &format!("{ap}1.to_v.weight"))?),
residual_scale: get(store, &format!("{ap}2.residual_scale"))
.ok()
.map(param_to_tensor),
});
ff_layers.push(FfLayer {
pre_norm: ScaleNorm {
g: get(store, &format!("{fp}0.0.g"))
.map(|p| p.data[0])
.unwrap_or(1.0),
dim,
},
w1: param_to_tensor(get(store, &format!("{fp}1.ff.0.0.weight"))?),
b1: param_to_tensor(get(store, &format!("{fp}1.ff.0.0.bias"))?),
w2: param_to_tensor(get(store, &format!("{fp}1.ff.2.weight"))?),
b2: param_to_tensor(get(store, &format!("{fp}1.ff.2.bias"))?),
residual_scale: get(store, &format!("{fp}2.residual_scale"))
.ok()
.map(param_to_tensor),
});
}
let final_g = get(store, &format!("{prefix}final_norm.g"))
.map(|p| p.data[0])
.unwrap_or(1.0);
Ok(Self {
dim,
heads: cfg.heads,
head_dim,
alibi_heads,
alibi_slopes: alibi_slopes(alibi_heads),
attn_layers,
ff_layers,
final_norm: ScaleNorm { g: final_g, dim },
})
}
fn alibi_bias(&self, n: usize) -> Vec<f32> {
let h = self.heads;
let mut bias = vec![0.0f32; h * n * n];
for hi in 0..h {
let slope = self.alibi_slopes[hi.min(self.alibi_slopes.len() - 1)];
for i in 0..n {
for j in 0..n {
bias[hi * n * n + i * n + j] = -((j as f32 - i as f32).abs()) * slope;
}
}
}
bias
}
fn self_attention(&self, layer: &AttnLayer, x: &Tensor, mask: &[bool]) -> Tensor {
let (b, n, d) = (x.shape[0], x.shape[1], x.shape[2]);
let h = self.heads;
let hd = self.head_dim;
let xn = layer.pre_norm.forward(x);
let q = xn.linear(&layer.q_w, None);
let k = xn.linear(&layer.k_w, None);
let v = xn.linear(&layer.v_w, None);
let alibi = self.alibi_bias(n);
let scale = (hd as f32).sqrt();
let mut out = vec![0.0f32; b * n * d];
for bi in 0..b {
for ni in 0..n {
for hi in 0..h {
let mut scores = vec![0.0f32; n];
let q_base = (bi * n + ni) * d + hi * hd;
for kj in 0..n {
if !mask[kj] {
scores[kj] = f32::NEG_INFINITY;
continue;
}
let mut dot = 0.0f32;
let k_base = (bi * n + kj) * d + hi * hd;
for di in 0..hd {
dot += q.data[q_base + di] * k.data[k_base + di];
}
scores[kj] = dot / scale + alibi[hi * n * n + ni * n + kj];
}
let max_s = scores.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let mut denom = 0.0f32;
for kj in 0..n {
if scores[kj].is_finite() {
scores[kj] = (scores[kj] - max_s).exp();
denom += scores[kj];
} else {
scores[kj] = 0.0;
}
}
if denom > 0.0 {
for s in &mut scores {
*s /= denom;
}
}
let out_base = (bi * n + ni) * d + hi * hd;
for di in 0..hd {
let mut sum = 0.0f32;
for kj in 0..n {
let v_base = (bi * n + kj) * d + hi * hd;
sum += scores[kj] * v.data[v_base + di];
}
out[out_base + di] = sum;
}
}
}
}
Tensor::from_vec(out, vec![b, n, d])
}
fn feed_forward(&self, layer: &FfLayer, x: &Tensor) -> Tensor {
let xn = layer.pre_norm.forward(x);
let mut h = xn.linear(&layer.w1, Some(&layer.b1)).gelu();
h = h.linear(&layer.w2, Some(&layer.b2));
h
}
fn apply_residual(&self, branch: &Tensor, residual: &Tensor, scale: Option<&Tensor>) -> Tensor {
let mut out = branch.data.clone();
if let Some(s) = scale {
for i in 0..out.len() {
out[i] += residual.data[i] * s.data[i % s.data.len()];
}
} else {
for (o, &r) in out.iter_mut().zip(residual.data.iter()) {
*o += r;
}
}
Tensor::from_vec(out, branch.shape.clone())
}
pub fn forward(&self, x: &Tensor, mask: &[bool]) -> Tensor {
assert_eq!(x.ndim(), 3);
assert_eq!(x.shape[2], self.dim);
assert_eq!(mask.len(), x.shape[1]);
let mut h = x.clone();
for (attn, ff) in self.attn_layers.iter().zip(self.ff_layers.iter()) {
let residual = h.clone();
let attn_out = self.self_attention(attn, &h, mask);
h = self.apply_residual(&attn_out, &residual, attn.residual_scale.as_ref());
let residual = h.clone();
let ff_out = self.feed_forward(ff, &h);
h = self.apply_residual(&ff_out, &residual, ff.residual_scale.as_ref());
}
self.final_norm.forward(&h)
}
}