ruda-model 0.21.25

Ruda model configuration, module parameters, records and data loading.
Documentation
use super::{ModuleMapper, Param, ParamId};
use ruda_tensor::{FloatDType, api::Tensor, backend::Backend, tensor::TensorContainer};

pub(super) struct DtypeMapper {
    dtype: FloatDType,
    converted: TensorContainer<(ParamId, bool)>,
}

impl DtypeMapper {
    pub(super) fn new(dtype: FloatDType) -> Self {
        Self {
            dtype,
            converted: TensorContainer::new(),
        }
    }
}

impl<B: Backend> ModuleMapper<B> for DtypeMapper {
    fn map_float<const D: usize>(&mut self, param: Param<Tensor<B, D>>) -> Param<Tensor<B, D>> {
        let (id, tensor, mapper) = param.consume();
        let requires_grad = tensor.is_require_grad();
        let key = (id, requires_grad);
        if let Some(converted) = self.converted.get::<B>(&key) {
            return Param::from_mapped_value(id, Tensor::<B, D>::from_primitive(converted), mapper);
        }
        let target_dtype: ruda_tensor::DType = self.dtype.into();
        let tensor = if tensor.dtype() == target_dtype {
            tensor
        } else {
            tensor
                .cast(self.dtype)
                .detach()
                .set_require_grad(requires_grad)
        };
        self.converted
            .register::<B>(key, tensor.clone().into_primitive());
        Param::from_mapped_value(id, tensor, mapper)
    }
}

#[cfg(test)]
mod tests {
    use super::*;
    use crate::{TestAutodiffBackend, module::Module};

    #[derive(Module, Debug)]
    struct SharedModule<B: Backend> {
        first: Param<Tensor<B, 1>>,
        alias: Param<Tensor<B, 1>>,
        frozen_alias: Param<Tensor<B, 1>>,
        frozen: Param<Tensor<B, 1>>,
    }

    #[test]
    fn dtype_conversion_preserves_shared_leaf_gradients_and_frozen_params() {
        let device = Default::default();
        type B = TestAutodiffBackend;
        for dtype in [FloatDType::F16, FloatDType::BF16] {
            let original = Tensor::<B, 1>::ones([2], &device).require_grad();
            let first = Param::from_tensor(original.clone());
            let id = first.id;
            let frozen = Param::from_tensor(Tensor::<B, 1>::ones([2], &device)).no_grad();
            let frozen_id = frozen.id;
            let model = SharedModule {
                first: first.clone(),
                frozen_alias: first.clone().no_grad(),
                alias: first,
                frozen,
            }
            .to_dtype(dtype);
            assert_eq!(model.first.id, id);
            assert_eq!(model.alias.id, id);
            assert_eq!(model.frozen_alias.id, id);
            assert_eq!(model.frozen.id, frozen_id);
            assert_eq!(model.first.val().dtype(), dtype.into());
            assert!(model.first.val().is_require_grad());
            assert!(!model.frozen.val().is_require_grad());
            assert!(!model.frozen_alias.val().is_require_grad());
            let input = Tensor::<B, 1>::ones([2], &device).cast(dtype);
            let output = input.clone() * model.first.val()
                + input * model.alias.val()
                + model.frozen.val()
                + model.frozen_alias.val();
            let gradients = output.sum().backward();
            for param in [&model.first, &model.alias] {
                let gradient = param
                    .val()
                    .grad(&gradients)
                    .expect("missing converted parameter gradient");
                assert_eq!(gradient.dtype(), dtype.into());
                assert_eq!(
                    gradient
                        .cast(FloatDType::F32)
                        .into_data()
                        .to_vec::<f32>()
                        .unwrap(),
                    alloc::vec![2.; 2]
                );
            }
            assert!(model.frozen.val().grad(&gradients).is_none());
            assert!(model.frozen_alias.val().grad(&gradients).is_none());
            assert!(original.grad(&gradients).is_none());
        }
    }

    #[test]
    fn dtype_conversion_keeps_diverged_frozen_alias_values_and_record_restore() {
        use crate::{
            module::ModuleDTypeRecord,
            record::{BinBytesRecorder, FullPrecisionSettings, Recorder},
        };
        let device = Default::default();
        type B = TestAutodiffBackend;
        for dtype in [FloatDType::F16, FloatDType::BF16] {
            let first = Param::from_tensor(Tensor::<B, 1>::from_floats([1., 3.], &device));
            let frozen_alias = first.clone().no_grad();
            let first = first.map(|tensor| (tensor + 1.).detach().require_grad());
            let module = SharedModule {
                first: first.clone(),
                alias: first,
                frozen_alias,
                frozen: Param::from_tensor(Tensor::<B, 1>::ones([2], &device)).no_grad(),
            }
            .to_dtype(dtype);
            assert_eq!(module.first.id, module.frozen_alias.id);
            let recorder = BinBytesRecorder::<FullPrecisionSettings>::default();
            let profile = ModuleDTypeRecord::capture(&module).unwrap();
            let bytes = <BinBytesRecorder<FullPrecisionSettings> as Recorder<B>>::record(
                &recorder,
                (module.clone().into_record(), profile),
                (),
            )
            .unwrap();
            let (record, profile): (<SharedModule<B> as Module<B>>::Record, ModuleDTypeRecord) =
                <BinBytesRecorder<FullPrecisionSettings> as Recorder<B>>::load(
                    &recorder, bytes, &device,
                )
                .unwrap();
            let restored = profile
                .apply(
                    module
                        .clone()
                        .to_dtype(FloatDType::F32)
                        .load_record(record)
                        .fork(&device),
                )
                .unwrap();
            for model in [module, restored] {
                assert_eq!(model.first.val().dtype(), dtype.into());
                assert_eq!(model.frozen_alias.val().dtype(), dtype.into());
                assert_eq!(
                    model
                        .first
                        .val()
                        .cast(FloatDType::F32)
                        .into_data()
                        .to_vec::<f32>()
                        .unwrap(),
                    alloc::vec![2., 4.]
                );
                assert_eq!(
                    model
                        .frozen_alias
                        .val()
                        .cast(FloatDType::F32)
                        .into_data()
                        .to_vec::<f32>()
                        .unwrap(),
                    alloc::vec![1., 3.]
                );
                assert!(!model.frozen_alias.val().is_require_grad());
                let gradients = (model.first.val() + model.alias.val() + model.frozen_alias.val())
                    .sum()
                    .backward();
                assert_eq!(
                    model
                        .first
                        .val()
                        .grad(&gradients)
                        .unwrap()
                        .cast(FloatDType::F32)
                        .into_data()
                        .to_vec::<f32>()
                        .unwrap(),
                    alloc::vec![2.; 2]
                );
                assert!(model.frozen_alias.val().grad(&gradients).is_none());
            }
        }
    }
}