ruda_model/module/
precision_record.rs1use 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#[derive(Clone, Debug, Serialize, Deserialize)]
15pub struct ModuleDTypeRecord {
16 entries: Vec<(u64, bool, DType)>,
17}
18
19impl ModuleDTypeRecord {
20 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 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}