1use std::collections::HashMap;
56use trustformers_core::errors::{Result, TrustformersError};
57use trustformers_core::tensor::Tensor;
58
59pub const NAMED_KEY_PREFIX: &str = "n:";
61pub const INDEXED_KEY_PREFIX: &str = "p:";
63
64#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
69pub struct ParamId(usize);
70
71impl ParamId {
72 pub fn index(self) -> usize {
74 self.0
75 }
76}
77
78#[derive(Debug, Clone)]
80struct ParamEntry {
81 name: Option<String>,
83 numel: usize,
85 key: String,
87 addr: Option<usize>,
91}
92
93#[derive(Debug, Clone, Default)]
98pub struct ParamRegistry {
99 entries: Vec<ParamEntry>,
101 by_name: HashMap<String, usize>,
103 by_addr: HashMap<usize, usize>,
105 bind_cursor: usize,
107}
108
109impl ParamRegistry {
110 pub fn new() -> Self {
112 Self::default()
113 }
114
115 pub fn len(&self) -> usize {
117 self.entries.len()
118 }
119
120 pub fn is_empty(&self) -> bool {
122 self.entries.is_empty()
123 }
124
125 pub fn clear(&mut self) {
127 self.entries.clear();
128 self.by_name.clear();
129 self.by_addr.clear();
130 self.bind_cursor = 0;
131 }
132
133 pub fn key(&self, id: ParamId) -> Option<&str> {
135 self.entries.get(id.0).map(|e| e.key.as_str())
136 }
137
138 pub fn name(&self, id: ParamId) -> Option<&str> {
140 self.entries.get(id.0).and_then(|e| e.name.as_deref())
141 }
142
143 pub fn numel(&self, id: ParamId) -> Option<usize> {
145 self.entries.get(id.0).map(|e| e.numel)
146 }
147
148 pub fn id_for_named_tensor(&mut self, name: &str, tensor: &Tensor) -> Result<ParamId> {
153 let (addr, numel) = tensor_identity(tensor)?;
154 Ok(self.id_for_named_addr(name, addr, numel))
155 }
156
157 pub fn id_for_named_addr(&mut self, name: &str, addr: usize, numel: usize) -> ParamId {
159 if let Some(&index) = self.by_name.get(name) {
160 if let Some(entry) = self.entries.get_mut(index) {
162 if let Some(old) = entry.addr.replace(addr) {
163 if old != addr {
164 self.by_addr.remove(&old);
165 }
166 }
167 if entry.numel == 0 {
169 entry.numel = numel;
170 }
171 }
172 self.by_addr.insert(addr, index);
173 self.advance_bind_cursor();
174 return ParamId(index);
175 }
176
177 let index = self.entries.len();
178 self.entries.push(ParamEntry {
179 name: Some(name.to_string()),
180 numel,
181 key: format!("{NAMED_KEY_PREFIX}{name}"),
182 addr: Some(addr),
183 });
184 self.by_name.insert(name.to_string(), index);
185 self.by_addr.insert(addr, index);
186 self.advance_bind_cursor();
187 ParamId(index)
188 }
189
190 pub fn id_for_tensor(&mut self, tensor: &Tensor) -> Result<ParamId> {
201 let (addr, numel) = tensor_identity(tensor)?;
202 self.id_for_addr(addr, numel)
203 }
204
205 pub fn id_for_addr(&mut self, addr: usize, numel: usize) -> Result<ParamId> {
211 if let Some(&index) = self.by_addr.get(&addr) {
213 if self.entries.get(index).map(|e| e.numel) == Some(numel) {
214 return Ok(ParamId(index));
215 }
216 self.by_addr.remove(&addr);
219 if let Some(entry) = self.entries.get_mut(index) {
220 entry.addr = None;
221 }
222 }
223
224 self.advance_bind_cursor();
228 if let Some(entry) = self.entries.get_mut(self.bind_cursor) {
229 if entry.name.is_some() {
230 return Err(TrustformersError::invalid_input(format!(
233 "optimizer state slot {} was checkpointed under the name '{}' but is \
234 being resumed through the anonymous update path; use `update_named` \
235 so the name can be matched",
236 self.bind_cursor,
237 entry.name.as_deref().unwrap_or("<unknown>")
238 )));
239 }
240 if entry.numel != 0 && entry.numel != numel {
241 return Err(TrustformersError::invalid_input(format!(
242 "optimizer state slot {} holds {} elements but the parameter being \
243 bound to it has {}; parameters must be passed to `update()` in the \
244 same order as the run that wrote the checkpoint (or use \
245 `update_named`)",
246 self.bind_cursor, entry.numel, numel
247 )));
248 }
249 entry.numel = numel;
250 entry.addr = Some(addr);
251 let index = self.bind_cursor;
252 self.by_addr.insert(addr, index);
253 self.advance_bind_cursor();
254 return Ok(ParamId(index));
255 }
256
257 let index = self.entries.len();
259 self.entries.push(ParamEntry {
260 name: None,
261 numel,
262 key: format!("{INDEXED_KEY_PREFIX}{index}"),
263 addr: Some(addr),
264 });
265 self.by_addr.insert(addr, index);
266 self.advance_bind_cursor();
267 Ok(ParamId(index))
268 }
269
270 pub fn rebind(&mut self, id: ParamId, tensor: &Tensor) -> Result<()> {
282 let (addr, numel) = tensor_identity(tensor)?;
283 let entry = self.entries.get_mut(id.0).ok_or_else(|| {
284 TrustformersError::invalid_input(format!(
285 "cannot rebind unregistered parameter id {}",
286 id.0
287 ))
288 })?;
289 if let Some(old) = entry.addr.replace(addr) {
290 if old != addr {
291 self.by_addr.remove(&old);
292 }
293 }
294 entry.numel = numel;
295 self.by_addr.insert(addr, id.0);
296 self.advance_bind_cursor();
297 Ok(())
298 }
299
300 pub fn key_for_named_tensor(&mut self, name: &str, tensor: &Tensor) -> Result<String> {
306 let id = self.id_for_named_tensor(name, tensor)?;
307 Ok(self.key_string(id))
308 }
309
310 pub fn key_for_named_addr(&mut self, name: &str, addr: usize, numel: usize) -> String {
313 let id = self.id_for_named_addr(name, addr, numel);
314 self.key_string(id)
315 }
316
317 pub fn key_for_tensor(&mut self, tensor: &Tensor) -> Result<String> {
325 let id = self.id_for_tensor(tensor)?;
326 Ok(self.key_string(id))
327 }
328
329 pub fn key_for_addr(&mut self, addr: usize, numel: usize) -> Result<String> {
336 let id = self.id_for_addr(addr, numel)?;
337 Ok(self.key_string(id))
338 }
339
340 pub fn restore_key(&mut self, key: &str, numel: usize) -> Result<ParamId> {
350 if let Some(name) = key.strip_prefix(NAMED_KEY_PREFIX) {
351 if let Some(&index) = self.by_name.get(name) {
352 if let Some(entry) = self.entries.get_mut(index) {
353 if entry.numel == 0 {
354 entry.numel = numel;
355 }
356 }
357 return Ok(ParamId(index));
358 }
359 let index = self.entries.len();
360 self.entries.push(ParamEntry {
361 name: Some(name.to_string()),
362 numel,
363 key: key.to_string(),
364 addr: None,
365 });
366 self.by_name.insert(name.to_string(), index);
367 self.reset_bind_cursor();
368 return Ok(ParamId(index));
369 }
370
371 if let Some(raw_index) = key.strip_prefix(INDEXED_KEY_PREFIX) {
372 let index: usize = raw_index.parse().map_err(|_| {
373 TrustformersError::invalid_input(format!(
374 "malformed optimizer state key '{key}': '{raw_index}' is not an index"
375 ))
376 })?;
377 while self.entries.len() <= index {
378 let placeholder = self.entries.len();
379 self.entries.push(ParamEntry {
380 name: None,
381 numel: 0,
382 key: format!("{INDEXED_KEY_PREFIX}{placeholder}"),
383 addr: None,
384 });
385 }
386 if let Some(entry) = self.entries.get_mut(index) {
387 if entry.name.is_none() {
388 entry.numel = numel;
389 }
390 }
391 self.reset_bind_cursor();
392 return Ok(ParamId(index));
393 }
394
395 Err(TrustformersError::invalid_input(format!(
396 "unrecognised optimizer state key '{key}': expected a '{NAMED_KEY_PREFIX}' or \
397 '{INDEXED_KEY_PREFIX}' identity prefix"
398 )))
399 }
400
401 pub fn keys(&self) -> Vec<String> {
403 self.entries.iter().map(|e| e.key.clone()).collect()
404 }
405
406 fn key_string(&self, id: ParamId) -> String {
407 self.entries
408 .get(id.0)
409 .map(|e| e.key.clone())
410 .unwrap_or_else(|| format!("{INDEXED_KEY_PREFIX}{}", id.0))
411 }
412
413 fn advance_bind_cursor(&mut self) {
414 while self.entries.get(self.bind_cursor).is_some_and(|e| e.addr.is_some()) {
415 self.bind_cursor += 1;
416 }
417 }
418
419 fn reset_bind_cursor(&mut self) {
420 self.bind_cursor = 0;
421 self.advance_bind_cursor();
422 }
423}
424
425fn tensor_identity(tensor: &Tensor) -> Result<(usize, usize)> {
429 let numel: usize = tensor.shape().iter().product();
430 let addr = match tensor {
431 Tensor::F32(a) => a.as_ptr() as usize,
432 Tensor::F64(a) => a.as_ptr() as usize,
433 Tensor::F16(a) => a.as_ptr() as usize,
434 Tensor::BF16(a) => a.as_ptr() as usize,
435 Tensor::I64(a) => a.as_ptr() as usize,
436 Tensor::C32(a) => a.as_ptr() as usize,
437 Tensor::C64(a) => a.as_ptr() as usize,
438 Tensor::CF16(a) => a.as_ptr() as usize,
439 Tensor::CBF16(a) => a.as_ptr() as usize,
440 other => {
441 return Err(TrustformersError::tensor_op_error(
442 &format!(
443 "cannot derive a parameter identity for tensor dtype {:?}",
444 other.dtype()
445 ),
446 "ParamRegistry::id_for_tensor",
447 ))
448 },
449 };
450 Ok((addr, numel))
451}
452
453#[cfg(test)]
454mod tests {
455 use super::*;
456
457 fn tensor(len: usize) -> Tensor {
458 Tensor::from_vec(vec![0.0_f32; len], &[len]).expect("tensor")
459 }
460
461 #[test]
462 fn same_tensor_resolves_to_one_id() {
463 let mut registry = ParamRegistry::new();
464 let t = tensor(4);
465 let a = registry.id_for_tensor(&t).expect("first");
466 let b = registry.id_for_tensor(&t).expect("second");
467 assert_eq!(a, b);
468 assert_eq!(registry.len(), 1, "state must not grow per call");
469 }
470
471 #[test]
472 fn distinct_tensors_get_distinct_ids() {
473 let mut registry = ParamRegistry::new();
474 let t1 = tensor(4);
475 let t2 = tensor(8);
476 let a = registry.id_for_tensor(&t1).expect("t1");
477 let b = registry.id_for_tensor(&t2).expect("t2");
478 assert_ne!(a, b);
479 assert_eq!(registry.len(), 2);
480 }
481
482 #[test]
483 fn mutating_a_tensor_in_place_does_not_change_its_id() {
484 let mut registry = ParamRegistry::new();
488 let mut t = tensor(4);
489 let before = registry.id_for_tensor(&t).expect("before");
490 match &mut t {
491 Tensor::F32(array) => {
492 for value in array.iter_mut() {
493 *value = 9.0;
494 }
495 },
496 _ => panic!("expected F32"),
497 }
498 let after = registry.id_for_tensor(&t).expect("after");
499 assert_eq!(before, after);
500 assert_eq!(registry.len(), 1);
501 }
502
503 #[test]
504 fn rebind_tracks_a_reallocated_buffer() {
505 let mut registry = ParamRegistry::new();
507 let mut t = tensor(4);
508 let id = registry.id_for_tensor(&t).expect("register");
509 t.set_data_f32(&[9.0, 9.0, 9.0, 9.0]).expect("reallocate");
510 registry.rebind(id, &t).expect("rebind");
511 assert_eq!(registry.id_for_tensor(&t).expect("after"), id);
512 assert_eq!(registry.len(), 1, "rebinding must not append a slot");
513 }
514
515 #[test]
516 fn named_ids_are_address_independent() {
517 let mut registry = ParamRegistry::new();
518 let first = tensor(4);
519 let id1 = registry.id_for_named_tensor("w", &first).expect("first");
520 drop(first);
521 let second = tensor(4);
522 let id2 = registry.id_for_named_tensor("w", &second).expect("second");
523 assert_eq!(id1, id2, "a name must outlive the tensor allocation");
524 assert_eq!(registry.key(id1), Some("n:w"));
525 }
526
527 #[test]
528 fn keys_are_stable_and_prefixed() {
529 let mut registry = ParamRegistry::new();
530 let t = tensor(2);
531 assert_eq!(registry.key_for_tensor(&t).expect("key"), "p:0");
532 assert_eq!(
533 registry.key_for_named_tensor("bias", &t).expect("key"),
534 "n:bias"
535 );
536 }
537
538 #[test]
539 fn restored_anonymous_slots_are_claimed_in_order() {
540 let mut registry = ParamRegistry::new();
543 registry.restore_key("p:0", 4).expect("restore 0");
544 registry.restore_key("p:1", 8).expect("restore 1");
545
546 let t1 = tensor(4);
547 let t2 = tensor(8);
548 assert_eq!(registry.key_for_tensor(&t1).expect("bind 0"), "p:0");
549 assert_eq!(registry.key_for_tensor(&t2).expect("bind 1"), "p:1");
550 assert_eq!(registry.len(), 2, "resume must not append new slots");
551 }
552
553 #[test]
554 fn restored_named_slots_match_by_name_in_any_order() {
555 let mut registry = ParamRegistry::new();
556 registry.restore_key("n:a", 4).expect("restore a");
557 registry.restore_key("n:b", 8).expect("restore b");
558
559 let tb = tensor(8);
560 let ta = tensor(4);
561 assert_eq!(registry.key_for_named_tensor("b", &tb).expect("b"), "n:b");
563 assert_eq!(registry.key_for_named_tensor("a", &ta).expect("a"), "n:a");
564 assert_eq!(registry.len(), 2);
565 }
566
567 #[test]
568 fn order_mismatch_on_resume_is_an_error_not_silent_reset() {
569 let mut registry = ParamRegistry::new();
570 registry.restore_key("p:0", 4).expect("restore 0");
571 let wrong = tensor(9);
572 let err = registry.id_for_tensor(&wrong);
573 assert!(
574 err.is_err(),
575 "binding a 9-element tensor to a 4-element slot must be reported"
576 );
577 }
578
579 #[test]
580 fn anonymous_path_refuses_to_hijack_a_named_slot() {
581 let mut registry = ParamRegistry::new();
582 registry.restore_key("n:w", 4).expect("restore");
583 let t = tensor(4);
584 assert!(registry.id_for_tensor(&t).is_err());
585 }
586
587 #[test]
588 fn restore_rejects_unprefixed_keys() {
589 let mut registry = ParamRegistry::new();
590 assert!(registry.restore_key("0x7f9c2a001234", 4).is_err());
591 }
592
593 #[test]
594 fn clear_resets_everything() {
595 let mut registry = ParamRegistry::new();
596 let t = tensor(4);
597 registry.id_for_tensor(&t).expect("register");
598 registry.clear();
599 assert!(registry.is_empty());
600 assert_eq!(registry.key_for_tensor(&t).expect("re-register"), "p:0");
601 }
602}