Skip to main content

ruda_optim/optim/grads/
record.rs

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/// Host-owned gradient state, keyed by the original model parameter IDs.
13///
14/// This record can be stored alongside model, optimizer and scheduler records.
15/// Recorder precision settings apply to floating-point values; the original
16/// gradient dtype and shape are restored when loading. Use settings that do not
17/// narrow the gradient values when exact continuation is required.
18#[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    /// Snapshot all gradients without clearing them, including accumulated gradients.
50    ///
51    /// `B` must be the backend used to register the gradients (normally the
52    /// autodiff backend's inner backend). Device reads complete before returning.
53    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    /// Asynchronously snapshot all gradients without clearing the container.
62    ///
63    /// `B` must match the backend used to register the gradients.
64    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    /// Restore recorded gradients to a device without summing or clearing entries.
86    ///
87    /// Restore the model's parameter IDs from the same checkpoint before using
88    /// these gradients. `B` is the gradient backend, normally the inner backend.
89    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}