ruda-optim 0.21.31

Ruda optimizer updates, gradient state, clipping and learning-rate schedules.
Documentation
use super::*;
use crate::{AdamW, AdamWConfig};
use crate::data_parallel::zero::{Zero1, Zero1Record};
use ruda_model::record::{FullPrecisionSettings, Record};
use ruda_nn::modules::tensor_parallel::{ColumnParallelLinear, RowParallelLinear};

#[test]
fn zero1_adamw_state_shards_match_dense_updates_and_resume() {
    world(|rank, communicator| {
        let (session, mut replica) = DataParallel::<B>::initialize(communicator, model(rank), 0).unwrap();
        let mut zero = Zero1::new(session, &replica, AdamWConfig::new().build(), vec![0, 1]).unwrap();
        let mut reference = model(0);
        let mut reference_optimizer = AdamWConfig::new().init();
        for _ in 0..3 {
            let gradients = GradientsParams::from_grads(loss(&replica, rank).backward(), &replica);
            let result = zero.step(0.01, replica, gradients, if rank == 0 { 2 } else { 1 }, MissingGradientPolicy::Error).unwrap();
            assert_eq!(result.global_weight, 3);
            replica = result.model;
            let mean = (loss(&reference, 0) + loss(&reference, 1)) / 3.;
            let gradients = GradientsParams::from_grads(mean.backward(), &reference);
            reference = reference_optimizer.step(0.01, reference, gradients);
            replica.weight.val().to_data().assert_approx_eq::<f32>(&reference.weight.val().to_data(), Tolerance::absolute(1e-6));
            replica.bias.as_ref().unwrap().val().to_data().assert_approx_eq::<f32>(&reference.bias.as_ref().unwrap().val().to_data(), Tolerance::absolute(1e-6));
        }
        assert_eq!(zero.owned_parameter_count(), 1);
        assert_eq!(zero.state_parameter_count(), 1);
        let saved_model = replica.clone();
        let saved = zero.to_record().into_item::<FullPrecisionSettings>();
        let gradients = GradientsParams::from_grads(loss(&replica, rank).backward(), &replica);
        let expected = zero.step(0.01, replica, gradients, if rank == 0 { 2 } else { 1 }, MissingGradientPolicy::Error).unwrap().model;
        let saved = Zero1Record::<B, Linear<B>, AdamW>::from_item::<FullPrecisionSettings>(saved, &Default::default());
        zero.load_record(saved).unwrap();
        let gradients = GradientsParams::from_grads(loss(&saved_model, rank).backward(), &saved_model);
        let resumed = zero.step(0.01, saved_model, gradients, if rank == 0 { 2 } else { 1 }, MissingGradientPolicy::Error).unwrap().model;
        resumed.weight.val().to_data().assert_eq(&expected.weight.val().to_data(), false);
        resumed.bias.unwrap().val().to_data().assert_eq(&expected.bias.unwrap().val().to_data(), false);
    });
}

#[test]
fn zero1_ownership_mismatch_rejects_collectively() {
    world(|rank, communicator| {
        let (session, replica) = DataParallel::<B>::initialize(communicator, model(rank), 0).unwrap();
        let owners = if rank == 0 { vec![0, 1] } else { vec![1, 0] };
        assert!(Zero1::new(session, &replica, AdamWConfig::new().build(), owners).is_err());
    });
}

#[test]
fn tensor_parallel_column_row_chain_matches_dense_forward_and_every_gradient() {
    world(|rank, communicator| {
        let device = Default::default();
        let first = Linear::<B> {
            weight: Param::from_tensor(Tensor::from_floats([[1., 2., 3., 4.], [-1., 0.5, 2., -2.]], &device)),
            bias: Some(Param::from_tensor(Tensor::from_floats([0.1, 0.2, 0.3, 0.4], &device))),
        };
        let second = Linear::<B> {
            weight: Param::from_tensor(Tensor::from_floats([[0.5], [1.], [-1.], [0.25]], &device)),
            bias: Some(Param::from_tensor(Tensor::from_floats([0.25], &device))),
        };
        let range = rank as usize * 2..(rank as usize + 1) * 2;
        let column = ColumnParallelLinear::from_shard(Linear::<B> {
            weight: Param::from_tensor(first.weight.val().slice_dim(1, range.clone()).detach()),
            bias: Some(Param::from_tensor(first.bias.as_ref().unwrap().val().slice_dim(0, range.clone()).detach())),
        });
        let row = RowParallelLinear::from_shard(Linear::<B> {
            weight: Param::from_tensor(second.weight.val().slice_dim(0, range.clone()).detach()),
            bias: Some(Param::from_tensor(second.bias.as_ref().unwrap().val().detach())),
        });
        let input = Tensor::<B, 2>::from_floats([[1., 2.], [-1., 0.5]], &device).require_grad();
        let reference_input = input.clone().detach().require_grad();
        let output = column.forward(input.clone(), communicator.clone(), false).unwrap().square();
        let output = row.forward(output, communicator, true).unwrap();
        let reference = second.forward(first.forward(reference_input.clone()).square());
        output.to_data().assert_approx_eq::<f32>(&reference.to_data(), Tolerance::absolute(1e-5));
        let gradients = output.square().sum().backward();
        let expected = reference.square().sum().backward();
        let compare = |actual: Tensor<Host, 2>, expected: Tensor<Host, 2>| {
            actual.to_data().assert_approx_eq::<f32>(&expected.to_data(), Tolerance::absolute(1e-3));
        };
        compare(input.grad(&gradients).unwrap(), reference_input.grad(&expected).unwrap());
        compare(column.local.weight.val().grad(&gradients).unwrap(), first.weight.val().grad(&expected).unwrap().slice_dim(1, range.clone()));
        compare(row.local.weight.val().grad(&gradients).unwrap(), second.weight.val().grad(&expected).unwrap().slice_dim(0, range.clone()));
        column.local.bias.as_ref().unwrap().val().grad(&gradients).unwrap().to_data().assert_approx_eq::<f32>(
            &first.bias.as_ref().unwrap().val().grad(&expected).unwrap().slice_dim(0, range).to_data(), Tolerance::absolute(1e-3));
        row.local.bias.as_ref().unwrap().val().grad(&gradients).unwrap().to_data().assert_approx_eq::<f32>(
            &second.bias.as_ref().unwrap().val().grad(&expected).unwrap().to_data(), Tolerance::absolute(1e-3));
    });
}

#[test]
fn tensor_parallel_gather_scatter_preserves_replicated_loss_derivative() {
    world(|_, communicator| {
        use ruda_autodiff::tensor_parallel::{gather_from_region, scatter_to_region};
        let input = Tensor::<B, 2>::from_floats([[1., 2., 3., 4.]], &Default::default()).require_grad();
        let shard = scatter_to_region(input.clone(), communicator.clone(), 1).unwrap();
        let output = gather_from_region(shard, communicator, 1).unwrap();
        output.to_data().assert_eq(&input.to_data(), false);
        let gradients = output.square().sum().backward();
        input.grad(&gradients).unwrap().to_data().assert_eq(&TensorData::from([[2., 4., 6., 8.]]), false);
    });
}