use super::*;
pub fn triangular_solve(
a: &dyn SparseTensor,
b: &Tensor,
upper: bool,
transpose: bool,
) -> TorshResult<Tensor> {
utils::validate_square(a)?;
let n = a.shape().dims()[0];
if b.shape().dims()[0] != n {
return Err(TorshError::InvalidArgument(format!(
"Dimension mismatch: matrix size {} but RHS size {}",
n,
b.shape().dims()[0]
)));
}
let is_vector = b.shape().ndim() == 1;
let nrhs = if is_vector { 1 } else { b.shape().dims()[1] };
let a_csr = a.to_csr()?;
let x = if is_vector {
zeros::<f32>(&[n])?
} else {
zeros::<f32>(&[n, nrhs])?
};
match (upper, transpose) {
(false, false) => {
for i in 0..n {
for j in 0..nrhs {
let mut sum = if is_vector {
b.get(&[i])?
} else {
b.get(&[i, j])?
};
let (cols, vals) = a_csr.get_row(i)?;
for (k, &col) in cols.iter().enumerate() {
if col < i {
let x_val = if is_vector {
x.get(&[col])?
} else {
x.get(&[col, j])?
};
sum -= vals[k] * x_val;
} else if col == i {
if vals[k].abs() < f32::EPSILON {
return Err(TorshError::ComputeError(
"Singular matrix: zero diagonal element".to_string(),
));
}
sum /= vals[k];
break;
}
}
if is_vector {
x.set(&[i], sum)?;
} else {
x.set(&[i, j], sum)?;
}
}
}
}
(true, false) => {
for i in (0..n).rev() {
for j in 0..nrhs {
let mut sum = if is_vector {
b.get(&[i])?
} else {
b.get(&[i, j])?
};
let (cols, vals) = a_csr.get_row(i)?;
for (k, &col) in cols.iter().enumerate() {
if col > i {
let x_val = if is_vector {
x.get(&[col])?
} else {
x.get(&[col, j])?
};
sum -= vals[k] * x_val;
} else if col == i {
if vals[k].abs() < f32::EPSILON {
return Err(TorshError::ComputeError(
"Singular matrix: zero diagonal element".to_string(),
));
}
sum /= vals[k];
}
}
if is_vector {
x.set(&[i], sum)?;
} else {
x.set(&[i, j], sum)?;
}
}
}
}
(false, true) => {
let a_csc = a.to_csc()?;
for i in (0..n).rev() {
for j in 0..nrhs {
let mut sum = if is_vector {
b.get(&[i])?
} else {
b.get(&[i, j])?
};
let (rows, vals) = a_csc.get_col(i)?;
for (k, &row) in rows.iter().enumerate() {
if row > i {
let x_val = if is_vector {
x.get(&[row])?
} else {
x.get(&[row, j])?
};
sum -= vals[k] * x_val;
} else if row == i {
if vals[k].abs() < f32::EPSILON {
return Err(TorshError::ComputeError(
"Singular matrix: zero diagonal element".to_string(),
));
}
sum /= vals[k];
}
}
if is_vector {
x.set(&[i], sum)?;
} else {
x.set(&[i, j], sum)?;
}
}
}
}
(true, true) => {
let a_csc = a.to_csc()?;
for i in 0..n {
for j in 0..nrhs {
let mut sum = if is_vector {
b.get(&[i])?
} else {
b.get(&[i, j])?
};
let (rows, vals) = a_csc.get_col(i)?;
for (k, &row) in rows.iter().enumerate() {
if row < i {
let x_val = if is_vector {
x.get(&[row])?
} else {
x.get(&[row, j])?
};
sum -= vals[k] * x_val;
} else if row == i {
if vals[k].abs() < f32::EPSILON {
return Err(TorshError::ComputeError(
"Singular matrix: zero diagonal element".to_string(),
));
}
sum /= vals[k];
break;
}
}
if is_vector {
x.set(&[i], sum)?;
} else {
x.set(&[i, j], sum)?;
}
}
}
}
}
Ok(x)
}
pub fn addcmul(
input: &dyn SparseTensor,
tensor1: &dyn SparseTensor,
tensor2: &dyn SparseTensor,
value: f32,
) -> TorshResult<CooTensor> {
utils::validate_same_shape(input, tensor1)?;
utils::validate_same_shape(input, tensor2)?;
let input_coo = utils::to_coo_safe(input)?;
let tensor1_coo = utils::to_coo_safe(tensor1)?;
let tensor2_coo = utils::to_coo_safe(tensor2)?;
let input_map = utils::create_position_map(&input_coo);
let tensor1_map = utils::create_position_map(&tensor1_coo);
let tensor2_map = utils::create_position_map(&tensor2_coo);
let mut result_map: HashMap<(usize, usize), f32> = HashMap::new();
for ((row, col), val) in input_map {
result_map.insert((row, col), val);
}
for ((row, col), val1) in tensor1_map {
if let Some(&val2) = tensor2_map.get(&(row, col)) {
let product = value * val1 * val2;
*result_map.entry((row, col)).or_insert(0.0) += product;
}
}
let (row_indices, col_indices, values): (Vec<_>, Vec<_>, Vec<_>) = result_map
.into_iter()
.filter(|(_, v)| v.abs() > f32::EPSILON)
.fold(
(Vec::new(), Vec::new(), Vec::new()),
|(mut rows, mut cols, mut vals), ((r, c), v)| {
rows.push(r);
cols.push(c);
vals.push(v);
(rows, cols, vals)
},
);
CooTensor::new(row_indices, col_indices, values, input.shape().clone())
}
pub fn addcdiv(
input: &dyn SparseTensor,
tensor1: &dyn SparseTensor,
tensor2: &dyn SparseTensor,
value: f32,
) -> TorshResult<CooTensor> {
utils::validate_same_shape(input, tensor1)?;
utils::validate_same_shape(input, tensor2)?;
let input_coo = utils::to_coo_safe(input)?;
let tensor1_coo = utils::to_coo_safe(tensor1)?;
let tensor2_coo = utils::to_coo_safe(tensor2)?;
let input_map = utils::create_position_map(&input_coo);
let tensor1_map = utils::create_position_map(&tensor1_coo);
let tensor2_map = utils::create_position_map(&tensor2_coo);
let mut result_map: HashMap<(usize, usize), f32> = HashMap::new();
for ((row, col), val) in input_map {
result_map.insert((row, col), val);
}
for ((row, col), val1) in tensor1_map {
if let Some(&val2) = tensor2_map.get(&(row, col)) {
if val2.abs() > f32::EPSILON {
let quotient = value * val1 / val2;
*result_map.entry((row, col)).or_insert(0.0) += quotient;
}
}
}
let (row_indices, col_indices, values): (Vec<_>, Vec<_>, Vec<_>) = result_map
.into_iter()
.filter(|(_, v)| v.abs() > f32::EPSILON)
.fold(
(Vec::new(), Vec::new(), Vec::new()),
|(mut rows, mut cols, mut vals), ((r, c), v)| {
rows.push(r);
cols.push(c);
vals.push(v);
(rows, cols, vals)
},
);
CooTensor::new(row_indices, col_indices, values, input.shape().clone())
}
pub fn masked_fill<F>(
tensor: &dyn SparseTensor,
condition: F,
fill_value: f32,
) -> TorshResult<CooTensor>
where
F: Fn(f32) -> bool,
{
let coo = utils::to_coo_safe(tensor)?;
let triplets: Vec<_> = coo
.triplets()
.into_iter()
.map(|(r, c, v)| {
if condition(v) {
(r, c, fill_value)
} else {
(r, c, v)
}
})
.collect();
let (row_indices, col_indices, values) =
utils::extract_filtered_triplets(triplets, f32::EPSILON);
CooTensor::new(row_indices, col_indices, values, tensor.shape().clone())
}
pub fn clamp(
tensor: &dyn SparseTensor,
min: Option<f32>,
max: Option<f32>,
) -> TorshResult<CooTensor> {
let coo = utils::to_coo_safe(tensor)?;
let triplets: Vec<_> = coo
.triplets()
.into_iter()
.map(|(r, c, mut v)| {
if let Some(min_val) = min {
v = v.max(min_val);
}
if let Some(max_val) = max {
v = v.min(max_val);
}
(r, c, v)
})
.collect();
let (row_indices, col_indices, values) =
utils::extract_filtered_triplets(triplets, f32::EPSILON);
CooTensor::new(row_indices, col_indices, values, tensor.shape().clone())
}
pub fn abs(tensor: &dyn SparseTensor) -> TorshResult<CooTensor> {
let coo = utils::to_coo_safe(tensor)?;
let triplets: Vec<_> = coo
.triplets()
.into_iter()
.map(|(r, c, v)| (r, c, v.abs()))
.collect();
let (row_indices, col_indices, values) = utils::extract_filtered_triplets(triplets, 0.0);
CooTensor::new(row_indices, col_indices, values, tensor.shape().clone())
}
pub fn sign(tensor: &dyn SparseTensor) -> TorshResult<CooTensor> {
let coo = utils::to_coo_safe(tensor)?;
let triplets: Vec<_> = coo
.triplets()
.into_iter()
.map(|(r, c, v)| {
let sign_val = if v > 0.0 {
1.0
} else if v < 0.0 {
-1.0
} else {
0.0
};
(r, c, sign_val)
})
.collect();
let (row_indices, col_indices, values) =
utils::extract_filtered_triplets(triplets, f32::EPSILON);
CooTensor::new(row_indices, col_indices, values, tensor.shape().clone())
}
pub fn pow(tensor: &dyn SparseTensor, exponent: f32) -> TorshResult<CooTensor> {
let coo = utils::to_coo_safe(tensor)?;
let triplets: Vec<_> = coo
.triplets()
.into_iter()
.map(|(r, c, v)| (r, c, v.powf(exponent)))
.collect();
let (row_indices, col_indices, values) =
utils::extract_filtered_triplets(triplets, f32::EPSILON);
CooTensor::new(row_indices, col_indices, values, tensor.shape().clone())
}
pub fn square(tensor: &dyn SparseTensor) -> TorshResult<CooTensor> {
pow(tensor, 2.0)
}
pub fn sqrt(tensor: &dyn SparseTensor) -> TorshResult<CooTensor> {
let coo = utils::to_coo_safe(tensor)?;
let triplets: Vec<_> = coo
.triplets()
.into_iter()
.map(|(r, c, v)| (r, c, v.sqrt()))
.collect();
let (row_indices, col_indices, values) =
utils::extract_filtered_triplets(triplets, f32::EPSILON);
CooTensor::new(row_indices, col_indices, values, tensor.shape().clone())
}