use torsh_core::error::Result;
use torsh_tensor::Tensor;
fn with_grad_tracking_disabled<F>(param: &mut Tensor, f: F) -> Result<()>
where
F: FnOnce(&mut Tensor) -> Result<()>,
{
if !param.requires_grad() {
return f(param);
}
let placeholder = Tensor::zeros(&[1], param.device())?;
let taken = std::mem::replace(param, placeholder);
let mut scratch = taken.requires_grad_(false);
let outcome = f(&mut scratch);
*param = scratch.requires_grad_(true);
outcome
}
pub fn sub_assign(param: &mut Tensor, update: &Tensor) -> Result<()> {
with_grad_tracking_disabled(param, |p| p.sub_(update).map(|_| ()))
}
pub fn add_assign(param: &mut Tensor, update: &Tensor) -> Result<()> {
with_grad_tracking_disabled(param, |p| p.add_(update).map(|_| ()))
}
pub fn scale(param: &mut Tensor, factor: f32) -> Result<()> {
with_grad_tracking_disabled(param, |p| p.mul_scalar_(factor))
}
pub fn assign(param: &mut Tensor, src: &Tensor) -> Result<()> {
if param.shape().dims() != src.shape().dims() {
return Err(torsh_core::error::TorshError::ShapeMismatch {
expected: param.shape().dims().to_vec(),
got: src.shape().dims().to_vec(),
});
}
with_grad_tracking_disabled(param, |p| {
p.mul_scalar_(0.0)?;
p.add_(src).map(|_| ())
})
}
#[cfg(test)]
mod tests {
use super::*;
use torsh_core::device::DeviceType;
fn tensor(data: Vec<f32>) -> Tensor {
let len = data.len();
Tensor::from_data(data, vec![len], DeviceType::Cpu).expect("tensor creation")
}
#[test]
fn sub_assign_preserves_gradient_and_flag() -> Result<()> {
let mut param = tensor(vec![1.0, 2.0, 3.0]).requires_grad_(true);
param.set_grad(Some(tensor(vec![7.0, 7.0, 7.0])));
sub_assign(&mut param, &tensor(vec![0.5, 0.5, 0.5]))?;
assert_eq!(param.to_vec()?, vec![0.5, 1.5, 2.5]);
assert!(param.requires_grad());
assert!(param.has_grad());
assert_eq!(
param.grad().expect("gradient preserved").to_vec()?,
vec![7.0, 7.0, 7.0]
);
Ok(())
}
#[test]
fn repeated_updates_accumulate() -> Result<()> {
let mut param = tensor(vec![0.0]).requires_grad_(true);
for _ in 0..4 {
sub_assign(&mut param, &tensor(vec![0.25]))?;
}
assert_eq!(param.to_vec()?, vec![-1.0]);
Ok(())
}
#[test]
fn assign_overwrites_values() -> Result<()> {
let mut param = tensor(vec![1.0, 2.0]).requires_grad_(true);
assign(&mut param, &tensor(vec![-3.0, 9.0]))?;
assert_eq!(param.to_vec()?, vec![-3.0, 9.0]);
assert!(param.requires_grad());
Ok(())
}
#[test]
fn assign_is_exact_across_magnitudes() -> Result<()> {
let mut param = tensor(vec![1000.0, -1e6, 3.0]).requires_grad_(true);
assign(&mut param, &tensor(vec![0.001, 1e-7, -2.5]))?;
assert_eq!(param.to_vec()?, vec![0.001, 1e-7, -2.5]);
Ok(())
}
#[test]
fn assign_leaves_snapshots_untouched() -> Result<()> {
let mut param = tensor(vec![1.0, 2.0]).requires_grad_(true);
let snapshot = param.clone();
assign(&mut param, &tensor(vec![9.0, 9.0]))?;
assert_eq!(param.to_vec()?, vec![9.0, 9.0]);
assert_eq!(snapshot.to_vec()?, vec![1.0, 2.0]);
Ok(())
}
#[test]
fn assign_rejects_shape_mismatch() {
let mut param = tensor(vec![1.0, 2.0]);
assert!(assign(&mut param, &tensor(vec![1.0])).is_err());
}
#[test]
fn add_assign_and_scale_round_trip() -> Result<()> {
let mut param = tensor(vec![1.0, -1.0]).requires_grad_(true);
add_assign(&mut param, &tensor(vec![1.0, 1.0]))?;
assert_eq!(param.to_vec()?, vec![2.0, 0.0]);
scale(&mut param, 3.0)?;
assert_eq!(param.to_vec()?, vec![6.0, 0.0]);
assert!(param.requires_grad());
Ok(())
}
}