1use super::GradientsParams;
2use alloc::{format, vec::Vec};
3use hashbrown::HashSet;
4use ruda_model::{
5 module::ParamId,
6 record::{PrecisionSettings, Record, RecorderError},
7 tensor::{DType, TensorData, TensorMetadata, TensorPrimitive, backend::Backend, try_read_sync},
8};
9use serde::{Deserialize, Serialize};
10use ruda_model::tensor::quantization::{QuantLevel, QuantParam, QuantScheme, QuantStore, QuantValue};
11
12#[derive(Clone, Debug, Serialize, Deserialize)]
19pub struct GradientsParamsRecord {
20 gradients: Vec<(u64, DType, TensorData)>,
21}
22
23impl<B: Backend> Record<B> for GradientsParamsRecord {
24 type Item<S: PrecisionSettings> = Self;
25
26 fn into_item<S: PrecisionSettings>(self) -> Self::Item<S> {
27 Self {
28 gradients: self
29 .gradients
30 .into_iter()
31 .map(|(id, dtype, data)| {
32 let data = if matches!(data.dtype, DType::QFloat(_)) {
33 data
34 } else {
35 data.convert::<S::FloatElem>()
36 };
37 (id, dtype, data)
38 })
39 .collect(),
40 }
41 }
42
43 fn from_item<S: PrecisionSettings>(item: Self::Item<S>, _device: &B::Device) -> Self {
44 item
45 }
46}
47
48impl GradientsParams {
49 pub fn try_to_record<B: Backend>(&self) -> Result<GradientsParamsRecord, RecorderError> {
54 try_read_sync(self.to_record_async::<B>()).ok_or_else(|| {
55 RecorderError::Unknown(
56 "Synchronous gradient read is unavailable; use to_record_async".into(),
57 )
58 })?
59 }
60
61 pub async fn to_record_async<B: Backend>(&self) -> Result<GradientsParamsRecord, RecorderError> {
65 let mut ids = self.container.ids().into_iter().copied().collect::<Vec<_>>();
66 ids.sort();
67 let mut gradients = Vec::with_capacity(ids.len());
68 for id in ids {
69 let primitive = self.container.get::<B>(&id).ok_or_else(|| {
70 RecorderError::Unknown(format!("Missing gradient for parameter {id}"))
71 })?;
72 let dtype = primitive.dtype();
73 let data = match primitive {
74 TensorPrimitive::Float(tensor) => B::float_into_data(tensor).await,
75 TensorPrimitive::QFloat(tensor) => B::q_into_data(tensor).await,
76 }
77 .map_err(|error| {
78 RecorderError::Unknown(format!("Reading gradient for parameter {id}: {error}"))
79 })?;
80 gradients.push((id.val(), dtype, data));
81 }
82 Ok(GradientsParamsRecord { gradients })
83 }
84
85 pub fn from_record<B: Backend>(
90 record: GradientsParamsRecord,
91 device: &B::Device,
92 ) -> Result<Self, RecorderError> {
93 let mut ids = HashSet::with_capacity(record.gradients.len());
94 for (id, dtype, data) in &record.gradients {
95 if !ids.insert(*id) {
96 return Err(RecorderError::Unknown(format!(
97 "Duplicate gradient parameter ID {id}"
98 )));
99 }
100 let compatible = match (dtype, data.dtype) {
101 (DType::QFloat(expected), DType::QFloat(actual)) => *expected == actual,
102 (DType::QFloat(_), _) | (_, DType::QFloat(_)) => false,
103 _ => dtype.is_float() && data.dtype.is_float(),
104 };
105 if !compatible {
106 return Err(RecorderError::Unknown(format!(
107 "Invalid gradient dtype for parameter {id}: {dtype:?}, stored {:?}",
108 data.dtype
109 )));
110 }
111 let num_elements = data.shape.iter().try_fold(1usize, |count, dim| {
112 count.checked_mul(*dim)
113 }).ok_or_else(|| RecorderError::Unknown(format!(
114 "Gradient shape overflows for parameter {id}: {:?}", data.shape
115 )))?;
116 if let DType::QFloat(scheme) = data.dtype {
117 validate_quantized_gradient(*id, data, scheme, num_elements)?;
118 }
119 if data.dtype.is_float() {
120 let stored_bytes = num_elements.checked_mul(data.dtype.size()).ok_or_else(|| {
121 RecorderError::Unknown(format!(
122 "Stored gradient byte count overflows for parameter {id}"
123 ))
124 })?;
125 if data.bytes.len() != stored_bytes {
126 return Err(RecorderError::Unknown(format!(
127 "Invalid gradient byte count for parameter {id}: expected {stored_bytes}, got {}",
128 data.bytes.len()
129 )));
130 }
131 let restored_bytes = num_elements.checked_mul(dtype.size()).ok_or_else(|| {
132 RecorderError::Unknown(format!(
133 "Restored gradient byte count overflows for parameter {id}"
134 ))
135 })?;
136 if restored_bytes > isize::MAX as usize {
137 return Err(RecorderError::Unknown(format!(
138 "Restored gradient allocation exceeds addressable size for parameter {id}"
139 )));
140 }
141 }
142 }
143
144 let mut gradients = Self::new();
145 for (id, dtype, data) in record.gradients {
146 let primitive = match dtype {
147 DType::QFloat(_) => TensorPrimitive::QFloat(B::q_from_data(data, device)),
148 _ => {
149 let tensor = B::float_from_data(data.convert_dtype(dtype), device);
150 let tensor = if tensor.dtype() == dtype {
151 tensor
152 } else {
153 B::float_cast(tensor, dtype.into())
154 };
155 TensorPrimitive::Float(tensor)
156 }
157 };
158 gradients
159 .container
160 .register::<B>(ParamId::from(id), primitive);
161 }
162 Ok(gradients)
163 }
164}
165
166fn validate_quantized_gradient(
167 id: u64,
168 data: &TensorData,
169 scheme: QuantScheme,
170 num_elements: usize,
171) -> Result<(), RecorderError> {
172 let invalid = |reason: &str| RecorderError::Unknown(format!(
173 "Invalid quantized gradient for parameter {id}: {reason}"
174 ));
175 let (values_count, value_size) = match scheme.store {
176 QuantStore::Native => (num_elements, 1usize),
177 QuantStore::PackedU32(dim) | QuantStore::PackedNative(dim) => {
178 let axis = data.shape.rank().checked_sub(dim).and_then(|rank| rank.checked_sub(1))
179 .ok_or_else(|| invalid("packing dimension exceeds tensor rank"))?;
180 let value_size = match scheme.store {
181 QuantStore::PackedU32(_) => 4usize,
182 QuantStore::PackedNative(_) if scheme.value == QuantValue::E2M1 => 1,
183 _ => return Err(invalid("value type does not support native packing")),
184 };
185 let num_quants = scheme.num_quants();
186 let count = data.shape.iter().enumerate().try_fold(1usize, |count, (dim, size)| {
187 let size = if dim == axis { size.div_ceil(num_quants) } else { *size };
188 count.checked_mul(size)
189 }).ok_or_else(|| invalid("packed value count overflows"))?;
190 (count, value_size)
191 }
192 };
193 let values_bytes = values_count.checked_mul(value_size)
194 .ok_or_else(|| invalid("packed value byte count overflows"))?;
195 let params_count = match scheme.level {
196 QuantLevel::Tensor => 1usize,
197 QuantLevel::Block(blocks) => {
198 let block_dims = blocks.to_dim_vec(data.shape.rank());
199 if block_dims.contains(&0) {
200 return Err(invalid("quantization block dimension is zero"));
201 }
202 data.shape.iter().zip(block_dims).try_fold(1usize, |count, (size, block)| {
203 count.checked_mul(size.div_ceil(block as usize))
204 }).ok_or_else(|| invalid("scale count overflows"))?
205 }
206 };
207 let param_size = match scheme.param {
208 QuantParam::F32 => 4usize,
209 QuantParam::F16 | QuantParam::BF16 => 2,
210 QuantParam::UE8M0 | QuantParam::UE4M3 => 1,
211 };
212 let required_bytes = params_count.checked_mul(param_size)
213 .and_then(|params_bytes| values_bytes.checked_add(params_bytes))
214 .ok_or_else(|| invalid("value and scale byte count overflows"))?;
215 if required_bytes > isize::MAX as usize {
216 return Err(invalid("value and scale allocation exceeds addressable size"));
217 }
218 let exact = scheme.param == QuantParam::F32;
219 if data.bytes.len() < required_bytes || (exact && data.bytes.len() != required_bytes) {
220 let qualifier = if exact { "" } else { "at least " };
221 return Err(RecorderError::Unknown(format!(
222 "Invalid quantized gradient byte count for parameter {id}: expected {qualifier}{required_bytes}, got {}",
223 data.bytes.len()
224 )));
225 }
226 Ok(())
227}