use super::attention_gqa::{add_positional_encoding, slice_pe};
#[allow(clippy::wildcard_imports)]
use super::*;
impl TransformerDecoderLayer {
#[must_use]
pub fn new(d_model: usize, nhead: usize, dim_feedforward: usize) -> Self {
Self {
self_attn: MultiHeadAttention::new(d_model, nhead),
cross_attn: MultiHeadAttention::new(d_model, nhead),
linear1: Linear::new(d_model, dim_feedforward),
linear2: Linear::new(dim_feedforward, d_model),
norm1: LayerNorm::new(&[d_model]),
norm2: LayerNorm::new(&[d_model]),
norm3: LayerNorm::new(&[d_model]),
dropout: Dropout::new(0.1),
dropout1: Dropout::new(0.1),
dropout2: Dropout::new(0.1),
dropout3: Dropout::new(0.1),
d_model,
training: true,
}
}
pub fn forward_with_memory(
&self,
tgt: &Tensor,
memory: &Tensor,
tgt_mask: Option<&Tensor>,
memory_mask: Option<&Tensor>,
) -> Tensor {
let tgt_norm = self.norm1.forward(tgt);
let (attn_out, _) = self.self_attn.forward_self(&tgt_norm, tgt_mask);
let attn_out = self.dropout1.forward(&attn_out);
let tgt = tgt.add(&attn_out);
let tgt_norm = self.norm2.forward(&tgt);
let (cross_out, _) = self
.cross_attn
.forward_qkv(&tgt_norm, memory, memory, memory_mask);
let cross_out = self.dropout2.forward(&cross_out);
let tgt = tgt.add(&cross_out);
let tgt_norm = self.norm3.forward(&tgt);
let ff_out = self.linear1.forward(&tgt_norm);
let ff_out = ff_out.gelu();
let ff_out = self.dropout.forward(&ff_out);
let ff_out = self.linear2.forward(&ff_out);
let ff_out = self.dropout3.forward(&ff_out);
tgt.add(&ff_out)
}
}
impl Module for TransformerDecoderLayer {
fn forward(&self, input: &Tensor) -> Tensor {
self.forward_with_memory(input, input, None, None)
}
fn parameters(&self) -> Vec<&Tensor> {
let mut params = self.self_attn.parameters();
params.extend(self.cross_attn.parameters());
params.extend(self.linear1.parameters());
params.extend(self.linear2.parameters());
params.extend(self.norm1.parameters());
params.extend(self.norm2.parameters());
params.extend(self.norm3.parameters());
params
}
fn parameters_mut(&mut self) -> Vec<&mut Tensor> {
let mut params = self.self_attn.parameters_mut();
params.extend(self.cross_attn.parameters_mut());
params.extend(self.linear1.parameters_mut());
params.extend(self.linear2.parameters_mut());
params.extend(self.norm1.parameters_mut());
params.extend(self.norm2.parameters_mut());
params.extend(self.norm3.parameters_mut());
params
}
fn train(&mut self) {
self.training = true;
self.self_attn.train();
self.cross_attn.train();
self.dropout.train();
self.dropout1.train();
self.dropout2.train();
self.dropout3.train();
}
fn eval(&mut self) {
self.training = false;
self.self_attn.eval();
self.cross_attn.eval();
self.dropout.eval();
self.dropout1.eval();
self.dropout2.eval();
self.dropout3.eval();
}
fn training(&self) -> bool {
self.training
}
}
impl std::fmt::Debug for TransformerDecoderLayer {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("TransformerDecoderLayer")
.field("d_model", &self.d_model)
.field("self_attn", &self.self_attn)
.field("cross_attn", &self.cross_attn)
.finish_non_exhaustive()
}
}
#[derive(Debug)]
pub struct PositionalEncoding {
d_model: usize,
max_len: usize,
dropout: Dropout,
pe: Tensor,
training: bool,
}
impl PositionalEncoding {
#[must_use]
pub fn new(d_model: usize, max_len: usize) -> Self {
let pe = compute_positional_encoding(d_model, max_len);
Self {
d_model,
max_len,
dropout: Dropout::new(0.1),
pe,
training: true,
}
}
pub fn with_dropout(mut self, dropout: f32) -> Self {
self.dropout = Dropout::new(dropout);
self
}
}
impl Module for PositionalEncoding {
fn forward(&self, input: &Tensor) -> Tensor {
let seq_len = input.shape()[1];
assert!(
seq_len <= self.max_len,
"Sequence length {seq_len} exceeds max_len {}",
self.max_len
);
let pe_slice = slice_pe(&self.pe, seq_len, self.d_model);
let output = add_positional_encoding(input, &pe_slice);
self.dropout.forward(&output)
}
fn train(&mut self) {
self.training = true;
self.dropout.train();
}
fn eval(&mut self) {
self.training = false;
self.dropout.eval();
}
fn training(&self) -> bool {
self.training
}
}
pub(crate) fn transpose_last_two(x: &Tensor) -> Tensor {
let shape = x.shape();
let ndim = shape.len();
if ndim < 2 {
return x.clone();
}
let last = shape[ndim - 1];
let second_last = shape[ndim - 2];
let mut new_shape = shape.to_vec();
new_shape[ndim - 2] = last;
new_shape[ndim - 1] = second_last;
let batch_size: usize = shape[..ndim - 2].iter().product();
let matrix_size = last * second_last;
let mut output = vec![0.0; x.data().len()];
const TILE: usize = 32;
let src = x.data();
for b in 0..batch_size {
let offset = b * matrix_size;
for i0 in (0..second_last).step_by(TILE) {
let i_end = (i0 + TILE).min(second_last);
for j0 in (0..last).step_by(TILE) {
let j_end = (j0 + TILE).min(last);
for i in i0..i_end {
let src_base = offset + i * last;
for j in j0..j_end {
output[offset + j * second_last + i] = src[src_base + j];
}
}
}
}
}
let mut result = Tensor::from_vec(output, &new_shape);
record_attention_grad(x, &mut result, || {
std::sync::Arc::new(crate::autograd::grad_fn::TransposeLastTwoBackward {
input_shape: shape.to_vec(),
})
});
result
}
fn record_attention_grad<F>(input: &Tensor, result: &mut Tensor, make: F)
where
F: FnOnce() -> std::sync::Arc<dyn crate::autograd::GradFn>,
{
use crate::autograd::{is_grad_enabled, with_graph};
if is_grad_enabled() && input.requires_grad_enabled() {
result.requires_grad_(true);
let grad_fn = make();
result.set_grad_fn(grad_fn.clone());
with_graph(|graph| {
graph.register_tensor(input.clone());
graph.record(result.id(), grad_fn, vec![input.id()]);
});
}
}
fn record_attention_grad2<F>(a: &Tensor, b: &Tensor, result: &mut Tensor, make: F)
where
F: FnOnce() -> std::sync::Arc<dyn crate::autograd::GradFn>,
{
use crate::autograd::{is_grad_enabled, with_graph};
if is_grad_enabled() && (a.requires_grad_enabled() || b.requires_grad_enabled()) {
result.requires_grad_(true);
let grad_fn = make();
result.set_grad_fn(grad_fn.clone());
with_graph(|graph| {
graph.register_tensor(a.clone());
graph.register_tensor(b.clone());
graph.record(result.id(), grad_fn, vec![a.id(), b.id()]);
});
}
}
#[allow(clippy::expect_used)]
pub(crate) fn matmul_batched(a: &Tensor, b: &Tensor) -> Tensor {
let a_shape = a.shape();
let b_shape = b.shape();
if a_shape.len() == 4 && b_shape.len() == 4 {
let (batch, heads, m, k1) = (a_shape[0], a_shape[1], a_shape[2], a_shape[3]);
let k2 = b_shape[2];
let n = b_shape[3];
assert_eq!(k1, k2, "Inner dimensions must match for matmul");
let output = Matrix::batched_matmul_4d(a.data(), b.data(), batch, heads, m, k1, n)
.expect("batched_matmul_4d failed: dimensions validated but operation failed");
let mut result = Tensor::from_vec(output, &[batch, heads, m, n]);
record_attention_grad2(a, b, &mut result, || {
std::sync::Arc::new(crate::autograd::grad_fn::BatchedMatmul4dBackward {
a: a.clone(),
b: b.clone(),
})
});
result
} else {
a.matmul(b)
}
}
pub(super) fn scale_tensor(x: &Tensor, scale: f32) -> Tensor {
x.mul_scalar(scale)
}
fn broadcast_mask_to(target_shape: &[usize], mask: &Tensor) -> Tensor {
let mask_shape = mask.shape();
let mask_data = mask.data();
let rank = target_shape.len();
let m_rank = mask_shape.len();
debug_assert!(
m_rank <= rank,
"add_mask: mask rank {m_rank} exceeds scores rank {rank} — not broadcastable"
);
let mut m_strides = vec![0usize; m_rank];
let mut acc = 1usize;
for (stride, &extent) in m_strides.iter_mut().zip(mask_shape.iter()).rev() {
*stride = acc;
acc *= extent;
}
let offset = rank.saturating_sub(m_rank);
let skip = m_rank.saturating_sub(rank);
let total: usize = target_shape.iter().product();
let mut out = vec![0.0f32; total];
let mut idx = vec![0usize; rank];
for slot in &mut out {
let mut m_off = 0usize;
for (d, (&extent, &stride)) in mask_shape
.iter()
.zip(m_strides.iter())
.enumerate()
.skip(skip)
{
let t_d = d - skip + offset;
let i = if extent == 1 {
0
} else if extent == target_shape[t_d] {
idx[t_d]
} else {
debug_assert!(
false,
"add_mask: mask dim {d} (extent {extent}) is not broadcastable \
against scores dim {t_d} (extent {})",
target_shape[t_d]
);
idx[t_d].min(extent - 1)
};
m_off += i * stride;
}
*slot = mask_data[m_off];
for d in (0..rank).rev() {
idx[d] += 1;
if idx[d] < target_shape[d] {
break;
}
idx[d] = 0;
}
}
Tensor::from_vec(out, target_shape)
}
#[provable_contracts_macros::contract(
"setfit-encoder-conformance-v1",
equation = "apply_additive_mask"
)]
pub(super) fn add_mask(scores: &Tensor, mask: &Tensor) -> Tensor {
contract_pre_apply_additive_mask!(scores.data());
let result = if scores.shape() == mask.shape() {
scores.add(mask)
} else {
let expanded = broadcast_mask_to(scores.shape(), mask);
scores.add(&expanded)
};
contract_post_apply_additive_mask!(result.data());
result
}
pub(super) fn softmax_last_dim(x: &Tensor) -> Tensor {
let mut result = crate::nn::functional::softmax(x, -1);
let out_for_grad = Tensor::from_vec(result.data().to_vec(), result.shape());
record_attention_grad(x, &mut result, move || {
std::sync::Arc::new(crate::autograd::grad_fn::SoftmaxLastDimBackward {
output: out_for_grad,
})
});
result
}
pub(super) fn apply_dropout(x: &Tensor, p: f32) -> Tensor {
crate::nn::functional::dropout(x, p, true)
}
pub(super) fn apply_dropout_seeded(x: &Tensor, p: f32, seed: Option<u64>) -> Tensor {
match seed {
None => apply_dropout(x, p),
Some(seed) => {
use crate::nn::module::Module as _;
crate::nn::Dropout::with_seed(p, seed).forward(x)
}
}
}
pub(super) fn reshape_for_attention(
x: &Tensor,
batch: usize,
seq_len: usize,
num_heads: usize,
head_dim: usize,
) -> Tensor {
let mut output = vec![0.0; batch * num_heads * seq_len * head_dim];
for b in 0..batch {
for s in 0..seq_len {
for h in 0..num_heads {
for d in 0..head_dim {
let in_idx = b * seq_len * (num_heads * head_dim)
+ s * (num_heads * head_dim)
+ h * head_dim
+ d;
let out_idx = b * num_heads * seq_len * head_dim
+ h * seq_len * head_dim
+ s * head_dim
+ d;
output[out_idx] = x.data()[in_idx];
}
}
}
}
let mut result = Tensor::from_vec(output, &[batch, num_heads, seq_len, head_dim]);
record_attention_grad(x, &mut result, || {
std::sync::Arc::new(crate::autograd::grad_fn::ReshapeForAttentionBackward {
batch,
seq_len,
num_heads,
head_dim,
})
});
result
}
pub(crate) fn reshape_from_attention(
x: &Tensor,
batch: usize,
seq_len: usize,
embed_dim: usize,
) -> Tensor {
let num_heads = x.shape()[1];
let head_dim = x.shape()[3];
let mut output = vec![0.0; batch * seq_len * embed_dim];
for b in 0..batch {
for s in 0..seq_len {
for h in 0..num_heads {
for d in 0..head_dim {
let in_idx = b * num_heads * seq_len * head_dim
+ h * seq_len * head_dim
+ s * head_dim
+ d;
let out_idx = b * seq_len * embed_dim + s * embed_dim + h * head_dim + d;
output[out_idx] = x.data()[in_idx];
}
}
}
}
let mut result = Tensor::from_vec(output, &[batch, seq_len, embed_dim]);
record_attention_grad(x, &mut result, || {
std::sync::Arc::new(crate::autograd::grad_fn::ReshapeFromAttentionBackward {
batch,
seq_len,
num_heads,
head_dim,
})
});
result
}
fn compute_positional_encoding(d_model: usize, max_len: usize) -> Tensor {
let mut pe = vec![0.0; max_len * d_model];
for pos in 0..max_len {
for i in 0..d_model / 2 {
let angle = pos as f32 / 10000_f32.powf(2.0 * i as f32 / d_model as f32);
pe[pos * d_model + 2 * i] = angle.sin();
pe[pos * d_model + 2 * i + 1] = angle.cos();
}
}
Tensor::new(&pe, &[max_len, d_model])
}