Skip to main content

ruda_model/module/
precision_record.rs

1use super::{Module, ModuleMapper, ModuleVisitor, Param, ParamId, precision::DtypeMapper};
2use crate::record::{PrecisionSettings, Record, RecorderError};
3use alloc::{format, vec::Vec};
4use hashbrown::{HashMap, HashSet};
5use ruda_tensor::{DType, FloatDType, api::Tensor, backend::Backend};
6use serde::{Deserialize, Serialize};
7
8/// Floating-parameter storage dtypes, separate from recorder value precision.
9///
10/// Capture alongside the original module record and apply after loading it.
11/// Parameter IDs and trainable/frozen configuration must match that record.
12/// Recorder settings still determine saved value precision; this metadata
13/// does not recover values narrowed during serialization.
14#[derive(Clone, Debug, Serialize, Deserialize)]
15pub struct ModuleDTypeRecord {
16    entries: Vec<(u64, bool, DType)>,
17}
18
19impl ModuleDTypeRecord {
20    /// Capture floating-parameter dtype metadata without reading tensor values.
21    pub fn capture<B: Backend, M: Module<B>>(module: &M) -> Result<Self, RecorderError> {
22        let mut visitor = Capture {
23            entries: HashMap::new(),
24            error: None,
25        };
26        module.visit(&mut visitor);
27        if let Some(error) = visitor.error {
28            return Err(error);
29        }
30        let mut entries: Vec<_> = visitor
31            .entries
32            .into_iter()
33            .map(|((id, trainable), dtype)| (id.val(), trainable, dtype))
34            .collect();
35        entries.sort_by_key(|(id, trainable, _)| (*id, *trainable));
36        Ok(Self { entries })
37    }
38
39    /// Restore per-parameter storage dtypes using the shared leaf-preserving mapper.
40    ///
41    /// Tied parameters keep one converted leaf, frozen aliases remain independent,
42    /// and already-recorded quantized tensors are checked rather than requantized.
43    /// Integer and Bool parameters are not modified.
44    pub fn apply<B: Backend, M: Module<B>>(self, module: M) -> Result<M, RecorderError> {
45        let mut targets = HashMap::new();
46        for (id, trainable, dtype) in self.entries {
47            if !matches!(
48                dtype,
49                DType::F64
50                    | DType::F32
51                    | DType::Flex32
52                    | DType::F16
53                    | DType::BF16
54                    | DType::QFloat(_)
55            ) {
56                return Err(invalid("non-floating dtype in module dtype record"));
57            }
58            if targets
59                .insert((ParamId::from(id), trainable), dtype)
60                .is_some()
61            {
62                return Err(invalid("duplicate parameter in module dtype record"));
63            }
64        }
65        let mut mapper = Apply {
66            targets,
67            seen: HashSet::new(),
68            mappers: HashMap::new(),
69            error: None,
70        };
71        let module = module.map(&mut mapper);
72        if let Some(error) = mapper.error {
73            return Err(error);
74        }
75        if mapper.seen.len() != mapper.targets.len() {
76            return Err(invalid("module dtype record includes an unknown parameter"));
77        }
78        Ok(module)
79    }
80}
81
82impl<B: Backend> Record<B> for ModuleDTypeRecord {
83    type Item<P: PrecisionSettings> = Self;
84    fn into_item<P: PrecisionSettings>(self) -> Self {
85        self
86    }
87    fn from_item<P: PrecisionSettings>(item: Self, _device: &B::Device) -> Self {
88        item
89    }
90}
91
92fn invalid(reason: &str) -> RecorderError {
93    RecorderError::Unknown(format!("Invalid module dtype record: {reason}"))
94}
95
96struct Capture {
97    entries: HashMap<(ParamId, bool), DType>,
98    error: Option<RecorderError>,
99}
100impl<B: Backend> ModuleVisitor<B> for Capture {
101    fn visit_float<const D: usize>(&mut self, param: &Param<Tensor<B, D>>) {
102        let tensor = param.val();
103        let key = (param.id, tensor.is_require_grad());
104        if let Some(previous) = self.entries.insert(key, tensor.dtype())
105            && previous != tensor.dtype()
106        {
107            self.error = Some(invalid("tied parameter has inconsistent storage dtypes"));
108        }
109    }
110}
111
112struct Apply {
113    targets: HashMap<(ParamId, bool), DType>,
114    seen: HashSet<(ParamId, bool)>,
115    mappers: HashMap<FloatDType, DtypeMapper>,
116    error: Option<RecorderError>,
117}
118impl<B: Backend> ModuleMapper<B> for Apply {
119    fn map_float<const D: usize>(&mut self, param: Param<Tensor<B, D>>) -> Param<Tensor<B, D>> {
120        let key = (param.id, param.val().is_require_grad());
121        let Some(&dtype) = self.targets.get(&key) else {
122            self.error = Some(invalid("parameter ID or trainable configuration differs"));
123            return param;
124        };
125        self.seen.insert(key);
126        if matches!(dtype, DType::QFloat(_)) {
127            if param.val().dtype() != dtype {
128                self.error = Some(invalid("quantized parameter scheme differs"));
129            }
130            return param;
131        }
132        let dtype = FloatDType::from(dtype);
133        let mapper = self
134            .mappers
135            .entry(dtype)
136            .or_insert_with(|| DtypeMapper::new(dtype));
137        <DtypeMapper as ModuleMapper<B>>::map_float(mapper, param)
138    }
139}