use torsh_core::{Result as TorshResult, TorshError};
use torsh_tensor::Tensor;
use super::core::NormOrd;
pub fn chain_matmul(matrices: &[Tensor]) -> TorshResult<Tensor> {
if matrices.is_empty() {
return Err(TorshError::invalid_argument_with_context(
"chain_matmul requires at least one matrix",
"chain_matmul",
));
}
if matrices.len() == 1 {
return Ok(matrices[0].clone());
}
for (i, mat) in matrices.iter().enumerate() {
if mat.shape().ndim() != 2 {
return Err(TorshError::invalid_argument_with_context(
&format!("Matrix {} is not 2D", i),
"chain_matmul",
));
}
}
for i in 0..matrices.len() - 1 {
let m1_cols = matrices[i].shape().dims()[1];
let m2_rows = matrices[i + 1].shape().dims()[0];
if m1_cols != m2_rows {
return Err(TorshError::invalid_argument_with_context(
&format!(
"Matrix dimensions incompatible: [{} x {}] @ [{} x {}]",
matrices[i].shape().dims()[0],
m1_cols,
m2_rows,
matrices[i + 1].shape().dims()[1]
),
"chain_matmul",
));
}
}
let mut dim_seq = Vec::with_capacity(matrices.len() + 1);
dim_seq.push(matrices[0].shape().dims()[0]);
for mat in matrices {
dim_seq.push(mat.shape().dims()[1]);
}
let (_cost, split) = matrix_chain_order(&dim_seq);
multiply_chain(matrices, &split, 0, matrices.len() - 1)
}
fn matrix_chain_order(dim_seq: &[usize]) -> (u128, Vec<Vec<usize>>) {
let n = dim_seq.len() - 1;
let mut m = vec![vec![0u128; n]; n];
let mut split = vec![vec![0usize; n]; n];
for chain_len in 2..=n {
for i in 0..=(n - chain_len) {
let j = i + chain_len - 1;
m[i][j] = u128::MAX;
for k in i..j {
let cost = m[i][k]
+ m[k + 1][j]
+ (dim_seq[i] as u128) * (dim_seq[k + 1] as u128) * (dim_seq[j + 1] as u128);
if cost < m[i][j] {
m[i][j] = cost;
split[i][j] = k;
}
}
}
}
(m[0][n - 1], split)
}
fn multiply_chain(
matrices: &[Tensor],
split: &[Vec<usize>],
i: usize,
j: usize,
) -> TorshResult<Tensor> {
if i == j {
return Ok(matrices[i].clone());
}
let k = split[i][j];
let left = multiply_chain(matrices, split, i, k)?;
let right = multiply_chain(matrices, split, k + 1, j)?;
left.matmul(&right)
}
pub fn matrix_chain_optimal_cost(dim_seq: &[usize]) -> u128 {
if dim_seq.len() < 2 {
return 0;
}
matrix_chain_order(dim_seq).0
}
pub fn norm(
tensor: &Tensor,
ord: Option<NormOrd>,
dim: Option<Vec<isize>>,
keepdim: bool,
) -> TorshResult<Tensor> {
let ord = ord.unwrap_or(NormOrd::Fro);
match ord {
NormOrd::Fro => {
let squared = tensor.pow(2.0)?;
let sum = if let Some(dims) = dim {
let mut result = squared;
for &d in dims.iter() {
result = result.sum_dim(&[d as i32], keepdim)?;
}
result
} else {
squared.sum()?
};
sum.sqrt()
}
NormOrd::Nuclear => {
crate::reduction::norm_nuclear(tensor)
}
NormOrd::Inf => {
if let Some(dims) = dim {
let abs_tensor = tensor.abs()?;
let mut result = abs_tensor;
for &d in dims.iter() {
result = result.sum_dim(&[d as i32], keepdim)?;
}
result.max(None, false)
} else {
tensor.abs()?.max(None, false)
}
}
NormOrd::NegInf => {
if let Some(dims) = dim {
let abs_tensor = tensor.abs()?;
let mut result = abs_tensor;
for &d in dims.iter() {
result = result.sum_dim(&[d as i32], keepdim)?;
}
result.min()
} else {
tensor.abs()?.min()
}
}
NormOrd::P(p) => {
let abs_p = tensor.abs()?.pow(p)?;
let sum = if let Some(dims) = dim {
let mut result = abs_p;
for &d in dims.iter() {
result = result.sum_dim(&[d as i32], keepdim)?;
}
result
} else {
abs_p.sum()?
};
sum.pow(1.0 / p)
}
}
}
pub fn bmm(input: &Tensor, mat2: &Tensor) -> TorshResult<Tensor> {
if input.shape().ndim() != 3 || mat2.shape().ndim() != 3 {
return Err(TorshError::invalid_argument_with_context(
"Batch matrix multiplication requires 3D tensors (batch, rows, cols)",
"bmm",
));
}
let input_binding = input.shape();
let input_dims = input_binding.dims();
let mat2_binding = mat2.shape();
let mat2_dims = mat2_binding.dims();
if input_dims[0] != mat2_dims[0] {
return Err(TorshError::invalid_argument_with_context(
&format!(
"Batch sizes don't match: {} vs {}",
input_dims[0], mat2_dims[0]
),
"bmm",
));
}
if input_dims[2] != mat2_dims[1] {
return Err(TorshError::invalid_argument_with_context(
&format!(
"Matrix dimensions incompatible: [{} x {}] @ [{} x {}]",
input_dims[1], input_dims[2], mat2_dims[1], mat2_dims[2]
),
"bmm",
));
}
let batch_size = input_dims[0];
let out_rows = input_dims[1];
let out_cols = mat2_dims[2];
let mut result_data = vec![0.0f32; batch_size * out_rows * out_cols];
let input_data = input.to_vec()?;
let mat2_data = mat2.to_vec()?;
for b in 0..batch_size {
for i in 0..out_rows {
for j in 0..out_cols {
let mut sum = 0.0f32;
for k in 0..input_dims[2] {
let input_idx = b * input_dims[1] * input_dims[2] + i * input_dims[2] + k;
let mat2_idx = b * mat2_dims[1] * mat2_dims[2] + k * mat2_dims[2] + j;
sum += input_data[input_idx] * mat2_data[mat2_idx];
}
let result_idx = b * out_rows * out_cols + i * out_cols + j;
result_data[result_idx] = sum;
}
}
}
Tensor::from_data(
result_data,
vec![batch_size, out_rows, out_cols],
input.device(),
)
}
pub fn baddbmm(
input: &Tensor,
batch1: &Tensor,
batch2: &Tensor,
beta: f32,
alpha: f32,
) -> TorshResult<Tensor> {
let mm_result = bmm(batch1, batch2)?;
let scaled_input = input.mul_scalar(beta)?;
let scaled_mm = mm_result.mul_scalar(alpha)?;
scaled_input.add_op(&scaled_mm)
}
#[cfg(test)]
mod tests {
use super::*;
use torsh_core::device::DeviceType;
fn pattern_matrix(rows: usize, cols: usize, seed: usize) -> TorshResult<Tensor> {
let mut data = Vec::with_capacity(rows * cols);
for i in 0..rows {
for j in 0..cols {
data.push(((i + j + seed) % 3) as f32);
}
}
Tensor::from_data(data, vec![rows, cols], DeviceType::Cpu)
}
#[test]
fn test_matrix_chain_optimal_cost_known_values() {
assert_eq!(matrix_chain_optimal_cost(&[10, 100, 5, 50]), 7500);
assert_eq!(matrix_chain_optimal_cost(&[50, 5, 100, 10]), 7500);
assert_eq!(
matrix_chain_optimal_cost(&[30, 35, 15, 5, 10, 20, 25]),
15125
);
assert_eq!(matrix_chain_optimal_cost(&[7, 3]), 0); assert_eq!(matrix_chain_optimal_cost(&[42]), 0); assert_eq!(matrix_chain_optimal_cost(&[]), 0); }
#[test]
fn test_chain_matmul_uses_optimal_order_and_matches_reference() -> TorshResult<()> {
let a1 = pattern_matrix(50, 5, 0)?;
let a2 = pattern_matrix(5, 100, 1)?;
let a3 = pattern_matrix(100, 10, 2)?;
assert_eq!(matrix_chain_optimal_cost(&[50, 5, 100, 10]), 7500);
let result = chain_matmul(&[a1.clone(), a2.clone(), a3.clone()])?;
assert_eq!(result.shape().dims(), &[50, 10]);
let reference = a1.matmul(&a2)?.matmul(&a3)?;
assert_eq!(reference.shape().dims(), &[50, 10]);
let result_data = result.to_vec()?;
let reference_data = reference.to_vec()?;
assert_eq!(result_data.len(), reference_data.len());
assert!(result_data.iter().any(|&v| v > 0.0));
for (idx, (&got, &want)) in result_data.iter().zip(reference_data.iter()).enumerate() {
assert!(
(got - want).abs() < 1e-3,
"element {idx}: chain_matmul={got}, reference={want}"
);
}
Ok(())
}
#[test]
fn test_chain_matmul_four_matrices_matches_reference() -> TorshResult<()> {
let dims = [4usize, 5, 3, 6, 2];
let mut mats = Vec::new();
for w in dims.windows(2).enumerate() {
let (seed, pair) = w;
let rows = pair[0];
let cols = pair[1];
let mut data = Vec::with_capacity(rows * cols);
for idx in 0..rows * cols {
data.push(((idx + seed) % 2) as f32);
}
mats.push(Tensor::from_data(data, vec![rows, cols], DeviceType::Cpu)?);
}
let result = chain_matmul(&mats)?;
assert_eq!(result.shape().dims(), &[4, 2]);
let reference = mats[0]
.matmul(&mats[1])?
.matmul(&mats[2])?
.matmul(&mats[3])?;
let result_data = result.to_vec()?;
let reference_data = reference.to_vec()?;
for (&got, &want) in result_data.iter().zip(reference_data.iter()) {
assert!((got - want).abs() < 1e-3, "chain={got}, reference={want}");
}
Ok(())
}
#[test]
fn test_chain_matmul_single_matrix_is_identity() -> TorshResult<()> {
let a = pattern_matrix(3, 4, 0)?;
let result = chain_matmul(std::slice::from_ref(&a))?;
assert_eq!(result.to_vec()?, a.to_vec()?);
Ok(())
}
}