feat(batch29): GPU GEMM dispatch, weight serialization, training metrics
CI / Build (ubuntu-latest) (push) Failing after 7s
CI / Format Check (push) Failing after 14s
Documentation / Build User Guide (push) Successful in 10s
CI / Clippy Check (push) Failing after 40s
CI / Build (macos-latest) (push) Failing after 44s
CI / Test (macos-latest) (push) Has been skipped
CI / Test (ubuntu-latest) (push) Has been skipped
CI / Python Bindings (maturin) (macos-latest) (push) Has been skipped
CI / Python Bindings (maturin) (ubuntu-latest) (push) Has been skipped
Documentation / Build API Documentation (push) Failing after 59s
Performance Benchmarks / Run Benchmarks (push) Successful in 2m23s
CI / WASM Build + Size Check (push) Has been skipped
CI / Distributed Training Tests (push) Has been skipped
CI / Build CPU-Only (Explicit) (push) Failing after 3m22s
CI / CI Success (push) Failing after 0s
CI / Build (ubuntu-latest) (push) Failing after 7s
CI / Format Check (push) Failing after 14s
Documentation / Build User Guide (push) Successful in 10s
CI / Clippy Check (push) Failing after 40s
CI / Build (macos-latest) (push) Failing after 44s
CI / Test (macos-latest) (push) Has been skipped
CI / Test (ubuntu-latest) (push) Has been skipped
CI / Python Bindings (maturin) (macos-latest) (push) Has been skipped
CI / Python Bindings (maturin) (ubuntu-latest) (push) Has been skipped
Documentation / Build API Documentation (push) Failing after 59s
Performance Benchmarks / Run Benchmarks (push) Successful in 2m23s
CI / WASM Build + Size Check (push) Has been skipped
CI / Distributed Training Tests (push) Has been skipped
CI / Build CPU-Only (Explicit) (push) Failing after 3m22s
CI / CI Success (push) Failing after 0s
Batch 29a — GpuViTEncoder cudarc round-trip + 6 new tests (18 total): - GpuWeightBuffers: CudaSlice<f32> for patch_embed/proj_w/per-block qkv+ffn - cuda() constructor: CudaContext::new() + stream.clone_htod() weight upload - encode(): GPU htod→dtoh round-trip when context+weights present; CPU fallback - warmup(): touches proj_w buffer via dtoh; has_gpu_weights(), gpu_buffer_count() - JepaTrainerV2 encoder field visibility: ViTBlock+CpuViTEncoder pub(crate) - JepaEncoder trait: as_any()/as_any_mut() for downcasting; impl on all encoders Batch 29b — Binary weight serialization (jepa_checkpoint.rs, 18 tests): - Format: b"JEPA" magic + version u32 + fields + step u64 + checksum u32 - serialize/deserialize_checkpoint(): pure binary, no deps - save/load_checkpoint(): file I/O wrappers with CheckpointError enum - encoder_to_fields() / apply_fields_to_encoder(): CpuViTEncoder ↔ WeightField - JepaTrainerV2::context_encoder_as_cpu[_mut]() via Any downcast - JepaCheckpoint::save_with_trainer(): writes JSON summary + .jepa binary - run_jepa_training(): auto-resume from config.resume_from checkpoint path Batch 29c — Training metrics logger (jepa_metrics.rs, 20 tests + 3 runner): - StepMetrics, WindowMetrics, TrainingSummaryReport types - JepaMetricsLogger: EMA loss (α=0.02), loss_trend() linear regression, eta_seconds(), progress_line() with [====>.....] bar and ETA - to_csv() / save_csv() export; training_summary() → TrainingSummaryReport - run_jepa_training() wired: delegates all logging to metrics_logger.progress_line() - JepaRunConfig: +metrics_csv_path (saved at end if set) Co-Authored-By: Claude Sonnet 4.6 <[email protected]>
This commit is contained in:
co-authored by
Claude Sonnet 4.6
parent
0b06c0fa81
commit
37db99107f
@@ -0,0 +1,729 @@
|
||||
//! JEPA model checkpoint serialization.
|
||||
//!
|
||||
//! Binary format: simple, self-describing, no external dependencies.
|
||||
//!
|
||||
//! File layout:
|
||||
//! ```text
|
||||
//! [0..4] magic: b"JEPA"
|
||||
//! [4..8] version: u32 LE = 1
|
||||
//! [8..12] field_count: u32 LE
|
||||
//! For each field:
|
||||
//! [N..N+4] name_len: u32 LE
|
||||
//! [N+4..] name: UTF-8 bytes (name_len bytes)
|
||||
//! [..] value_count: u32 LE
|
||||
//! [..] values: f32 LE (value_count × 4 bytes)
|
||||
//! [end-8] step: u64 LE
|
||||
//! [end-4] checksum: u32 LE (sum of all value bytes, wrapping)
|
||||
//! ```
|
||||
|
||||
use super::jepa_vit::CpuViTEncoder;
|
||||
|
||||
// ============================================================================
|
||||
// Public types
|
||||
// ============================================================================
|
||||
|
||||
/// Named weight tensor saved to a JEPA checkpoint.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct WeightField {
|
||||
pub name: String,
|
||||
pub values: Vec<f32>,
|
||||
}
|
||||
|
||||
/// A complete JEPA model checkpoint.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct JepaModelCheckpoint {
|
||||
pub step: usize,
|
||||
pub fields: Vec<WeightField>,
|
||||
/// JSON-encoded JepaRunConfig fields for reference.
|
||||
pub config_summary: String,
|
||||
pub mean_loss: f32,
|
||||
pub ema_tau: f32,
|
||||
}
|
||||
|
||||
/// Save/load errors.
|
||||
#[derive(Debug)]
|
||||
pub enum CheckpointError {
|
||||
Io(String),
|
||||
InvalidMagic,
|
||||
VersionMismatch { found: u32, expected: u32 },
|
||||
ChecksumMismatch,
|
||||
MissingField(String),
|
||||
WrongSize { field: String, expected: usize, found: usize },
|
||||
}
|
||||
|
||||
impl std::fmt::Display for CheckpointError {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
match self {
|
||||
CheckpointError::Io(msg) => write!(f, "I/O error: {msg}"),
|
||||
CheckpointError::InvalidMagic => write!(f, "invalid magic bytes (expected b\"JEPA\")"),
|
||||
CheckpointError::VersionMismatch { found, expected } => {
|
||||
write!(f, "version mismatch: found {found}, expected {expected}")
|
||||
}
|
||||
CheckpointError::ChecksumMismatch => write!(f, "checksum mismatch: data may be corrupt"),
|
||||
CheckpointError::MissingField(name) => write!(f, "missing field: {name}"),
|
||||
CheckpointError::WrongSize { field, expected, found } => {
|
||||
write!(f, "field '{field}' has wrong size: expected {expected}, found {found}")
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// Binary format helpers
|
||||
// ============================================================================
|
||||
|
||||
const MAGIC: &[u8; 4] = b"JEPA";
|
||||
const VERSION: u32 = 1;
|
||||
|
||||
/// Read a u32 from `data` starting at `*cursor` (little-endian). Advances cursor.
|
||||
fn read_u32(data: &[u8], cursor: &mut usize) -> Option<u32> {
|
||||
if *cursor + 4 > data.len() {
|
||||
return None;
|
||||
}
|
||||
let v = u32::from_le_bytes([
|
||||
data[*cursor],
|
||||
data[*cursor + 1],
|
||||
data[*cursor + 2],
|
||||
data[*cursor + 3],
|
||||
]);
|
||||
*cursor += 4;
|
||||
Some(v)
|
||||
}
|
||||
|
||||
/// Read a u64 from `data` starting at `*cursor` (little-endian). Advances cursor.
|
||||
fn read_u64(data: &[u8], cursor: &mut usize) -> Option<u64> {
|
||||
if *cursor + 8 > data.len() {
|
||||
return None;
|
||||
}
|
||||
let v = u64::from_le_bytes([
|
||||
data[*cursor],
|
||||
data[*cursor + 1],
|
||||
data[*cursor + 2],
|
||||
data[*cursor + 3],
|
||||
data[*cursor + 4],
|
||||
data[*cursor + 5],
|
||||
data[*cursor + 6],
|
||||
data[*cursor + 7],
|
||||
]);
|
||||
*cursor += 8;
|
||||
Some(v)
|
||||
}
|
||||
|
||||
/// Write a u32 in little-endian to `buf`.
|
||||
#[inline]
|
||||
fn push_u32(buf: &mut Vec<u8>, v: u32) {
|
||||
buf.extend_from_slice(&v.to_le_bytes());
|
||||
}
|
||||
|
||||
/// Write a u64 in little-endian to `buf`.
|
||||
#[inline]
|
||||
fn push_u64(buf: &mut Vec<u8>, v: u64) {
|
||||
buf.extend_from_slice(&v.to_le_bytes());
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// Serialize / Deserialize
|
||||
// ============================================================================
|
||||
|
||||
/// Serialize a [`JepaModelCheckpoint`] to binary bytes.
|
||||
///
|
||||
/// Format:
|
||||
/// 1. `b"JEPA"` magic (4 bytes)
|
||||
/// 2. version = 1 as u32 LE (4 bytes)
|
||||
/// 3. field_count as u32 LE (4 bytes)
|
||||
/// 4. For each WeightField: name_len u32 LE, name bytes, value_count u32 LE, f32 values
|
||||
/// 5. step as u64 LE (8 bytes)
|
||||
/// 6. checksum: wrapping sum of all f32 raw bytes, as u32 LE (4 bytes)
|
||||
pub fn serialize_checkpoint(ckpt: &JepaModelCheckpoint) -> Vec<u8> {
|
||||
let mut buf = Vec::new();
|
||||
|
||||
// Magic
|
||||
buf.extend_from_slice(MAGIC);
|
||||
|
||||
// Version
|
||||
push_u32(&mut buf, VERSION);
|
||||
|
||||
// Field count
|
||||
push_u32(&mut buf, ckpt.fields.len() as u32);
|
||||
|
||||
// Accumulate checksum across all value bytes
|
||||
let mut checksum = 0u32;
|
||||
|
||||
for field in &ckpt.fields {
|
||||
// name_len + name bytes
|
||||
let name_bytes = field.name.as_bytes();
|
||||
push_u32(&mut buf, name_bytes.len() as u32);
|
||||
buf.extend_from_slice(name_bytes);
|
||||
|
||||
// value_count + f32 values
|
||||
push_u32(&mut buf, field.values.len() as u32);
|
||||
for &v in &field.values {
|
||||
let bytes = v.to_le_bytes();
|
||||
// Accumulate each byte into checksum
|
||||
for b in bytes {
|
||||
checksum = checksum.wrapping_add(b as u32);
|
||||
}
|
||||
buf.extend_from_slice(&bytes);
|
||||
}
|
||||
}
|
||||
|
||||
// Step
|
||||
push_u64(&mut buf, ckpt.step as u64);
|
||||
|
||||
// Checksum
|
||||
push_u32(&mut buf, checksum);
|
||||
|
||||
buf
|
||||
}
|
||||
|
||||
/// Deserialize a [`JepaModelCheckpoint`] from binary bytes.
|
||||
///
|
||||
/// Returns `Err` on format violations (magic, version, checksum).
|
||||
pub fn deserialize_checkpoint(data: &[u8]) -> Result<JepaModelCheckpoint, CheckpointError> {
|
||||
let mut cursor = 0usize;
|
||||
|
||||
// Magic
|
||||
if data.len() < 4 || &data[0..4] != MAGIC {
|
||||
return Err(CheckpointError::InvalidMagic);
|
||||
}
|
||||
cursor += 4;
|
||||
|
||||
// Version
|
||||
let version = read_u32(data, &mut cursor).ok_or(CheckpointError::InvalidMagic)?;
|
||||
if version != VERSION {
|
||||
return Err(CheckpointError::VersionMismatch { found: version, expected: VERSION });
|
||||
}
|
||||
|
||||
// Field count
|
||||
let field_count = read_u32(data, &mut cursor)
|
||||
.ok_or_else(|| CheckpointError::Io("truncated at field_count".to_string()))? as usize;
|
||||
|
||||
// Read fields
|
||||
let mut fields = Vec::with_capacity(field_count);
|
||||
let mut checksum_computed = 0u32;
|
||||
|
||||
for _ in 0..field_count {
|
||||
// name_len
|
||||
let name_len = read_u32(data, &mut cursor)
|
||||
.ok_or_else(|| CheckpointError::Io("truncated at name_len".to_string()))? as usize;
|
||||
|
||||
// name bytes
|
||||
if cursor + name_len > data.len() {
|
||||
return Err(CheckpointError::Io("truncated at name bytes".to_string()));
|
||||
}
|
||||
let name = String::from_utf8(data[cursor..cursor + name_len].to_vec())
|
||||
.map_err(|e| CheckpointError::Io(format!("invalid UTF-8 in field name: {e}")))?;
|
||||
cursor += name_len;
|
||||
|
||||
// value_count
|
||||
let value_count = read_u32(data, &mut cursor)
|
||||
.ok_or_else(|| CheckpointError::Io("truncated at value_count".to_string()))? as usize;
|
||||
|
||||
// f32 values
|
||||
let byte_len = value_count * 4;
|
||||
if cursor + byte_len > data.len() {
|
||||
return Err(CheckpointError::Io(format!(
|
||||
"truncated reading values for field '{name}': need {byte_len} bytes at offset {cursor}, have {}",
|
||||
data.len()
|
||||
)));
|
||||
}
|
||||
let mut values = Vec::with_capacity(value_count);
|
||||
for i in 0..value_count {
|
||||
let off = cursor + i * 4;
|
||||
let bytes = [data[off], data[off + 1], data[off + 2], data[off + 3]];
|
||||
for b in bytes {
|
||||
checksum_computed = checksum_computed.wrapping_add(b as u32);
|
||||
}
|
||||
values.push(f32::from_le_bytes(bytes));
|
||||
}
|
||||
cursor += byte_len;
|
||||
|
||||
fields.push(WeightField { name, values });
|
||||
}
|
||||
|
||||
// Step
|
||||
let step = read_u64(data, &mut cursor)
|
||||
.ok_or_else(|| CheckpointError::Io("truncated at step".to_string()))? as usize;
|
||||
|
||||
// Checksum
|
||||
let checksum_stored = read_u32(data, &mut cursor)
|
||||
.ok_or_else(|| CheckpointError::Io("truncated at checksum".to_string()))?;
|
||||
|
||||
if checksum_stored != checksum_computed {
|
||||
return Err(CheckpointError::ChecksumMismatch);
|
||||
}
|
||||
|
||||
Ok(JepaModelCheckpoint {
|
||||
step,
|
||||
fields,
|
||||
config_summary: String::new(),
|
||||
mean_loss: 0.0,
|
||||
ema_tau: 0.0,
|
||||
})
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// File I/O
|
||||
// ============================================================================
|
||||
|
||||
/// Write a checkpoint to disk at `path`.
|
||||
pub fn save_checkpoint(ckpt: &JepaModelCheckpoint, path: &str) -> Result<(), CheckpointError> {
|
||||
let bytes = serialize_checkpoint(ckpt);
|
||||
std::fs::write(path, &bytes)
|
||||
.map_err(|e| CheckpointError::Io(format!("failed to write '{path}': {e}")))?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Load a checkpoint from disk.
|
||||
pub fn load_checkpoint(path: &str) -> Result<JepaModelCheckpoint, CheckpointError> {
|
||||
let data = std::fs::read(path)
|
||||
.map_err(|e| CheckpointError::Io(format!("failed to read '{path}': {e}")))?;
|
||||
deserialize_checkpoint(&data)
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// Encoder <-> WeightField conversion
|
||||
// ============================================================================
|
||||
|
||||
/// Extract weights from a [`CpuViTEncoder`] into a list of [`WeightField`]s.
|
||||
///
|
||||
/// Field names:
|
||||
/// - `"patch_embed"` — `[num_patches * embed_dim]` sinusoidal PE + projection
|
||||
/// - `"proj_w"` — `[embed_dim * embed_dim]` patch projection weights
|
||||
/// - `"proj_b"` — `[embed_dim]` patch projection bias
|
||||
/// - `"block_{i}_qkv_w"` — `[embed_dim * 3 * embed_dim]`
|
||||
/// - `"block_{i}_qkv_b"` — `[3 * embed_dim]`
|
||||
/// - `"block_{i}_out_w"` — `[embed_dim * embed_dim]`
|
||||
/// - `"block_{i}_out_b"` — `[embed_dim]`
|
||||
/// - `"block_{i}_ffn1_w"` — `[embed_dim * ffn_dim]`
|
||||
/// - `"block_{i}_ffn1_b"` — `[ffn_dim]`
|
||||
/// - `"block_{i}_ffn2_w"` — `[ffn_dim * embed_dim]`
|
||||
/// - `"block_{i}_ffn2_b"` — `[embed_dim]`
|
||||
pub fn encoder_to_fields(encoder: &CpuViTEncoder) -> Vec<WeightField> {
|
||||
let mut fields = Vec::new();
|
||||
|
||||
fields.push(WeightField {
|
||||
name: "patch_embed".to_string(),
|
||||
values: encoder.patch_embed.clone(),
|
||||
});
|
||||
fields.push(WeightField {
|
||||
name: "proj_w".to_string(),
|
||||
values: encoder.proj_w.clone(),
|
||||
});
|
||||
fields.push(WeightField {
|
||||
name: "proj_b".to_string(),
|
||||
values: encoder.proj_b.clone(),
|
||||
});
|
||||
|
||||
for (i, block) in encoder.blocks.iter().enumerate() {
|
||||
fields.push(WeightField {
|
||||
name: format!("block_{i}_qkv_w"),
|
||||
values: block.qkv_w.clone(),
|
||||
});
|
||||
fields.push(WeightField {
|
||||
name: format!("block_{i}_qkv_b"),
|
||||
values: block.qkv_b.clone(),
|
||||
});
|
||||
fields.push(WeightField {
|
||||
name: format!("block_{i}_out_w"),
|
||||
values: block.out_w.clone(),
|
||||
});
|
||||
fields.push(WeightField {
|
||||
name: format!("block_{i}_out_b"),
|
||||
values: block.out_b.clone(),
|
||||
});
|
||||
fields.push(WeightField {
|
||||
name: format!("block_{i}_ffn1_w"),
|
||||
values: block.ffn1_w.clone(),
|
||||
});
|
||||
fields.push(WeightField {
|
||||
name: format!("block_{i}_ffn1_b"),
|
||||
values: block.ffn1_b.clone(),
|
||||
});
|
||||
fields.push(WeightField {
|
||||
name: format!("block_{i}_ffn2_w"),
|
||||
values: block.ffn2_w.clone(),
|
||||
});
|
||||
fields.push(WeightField {
|
||||
name: format!("block_{i}_ffn2_b"),
|
||||
values: block.ffn2_b.clone(),
|
||||
});
|
||||
}
|
||||
|
||||
fields
|
||||
}
|
||||
|
||||
/// Apply [`WeightField`]s back into a [`CpuViTEncoder`] in-place.
|
||||
///
|
||||
/// Returns `Err(MissingField)` if a required field is absent,
|
||||
/// or `Err(WrongSize)` if a field has incorrect length.
|
||||
pub fn apply_fields_to_encoder(
|
||||
encoder: &mut CpuViTEncoder,
|
||||
fields: &[WeightField],
|
||||
) -> Result<(), CheckpointError> {
|
||||
// Build a lookup map from field name to values slice
|
||||
let map: std::collections::HashMap<&str, &[f32]> = fields
|
||||
.iter()
|
||||
.map(|f| (f.name.as_str(), f.values.as_slice()))
|
||||
.collect();
|
||||
|
||||
macro_rules! apply_field {
|
||||
($name:expr, $target:expr) => {{
|
||||
let name: &str = $name;
|
||||
let vals = map
|
||||
.get(name)
|
||||
.ok_or_else(|| CheckpointError::MissingField(name.to_string()))?;
|
||||
if vals.len() != $target.len() {
|
||||
return Err(CheckpointError::WrongSize {
|
||||
field: name.to_string(),
|
||||
expected: $target.len(),
|
||||
found: vals.len(),
|
||||
});
|
||||
}
|
||||
$target.copy_from_slice(vals);
|
||||
}};
|
||||
}
|
||||
|
||||
apply_field!("patch_embed", encoder.patch_embed);
|
||||
apply_field!("proj_w", encoder.proj_w);
|
||||
apply_field!("proj_b", encoder.proj_b);
|
||||
|
||||
for (i, block) in encoder.blocks.iter_mut().enumerate() {
|
||||
apply_field!(&format!("block_{i}_qkv_w"), block.qkv_w);
|
||||
apply_field!(&format!("block_{i}_qkv_b"), block.qkv_b);
|
||||
apply_field!(&format!("block_{i}_out_w"), block.out_w);
|
||||
apply_field!(&format!("block_{i}_out_b"), block.out_b);
|
||||
apply_field!(&format!("block_{i}_ffn1_w"), block.ffn1_w);
|
||||
apply_field!(&format!("block_{i}_ffn1_b"), block.ffn1_b);
|
||||
apply_field!(&format!("block_{i}_ffn2_w"), block.ffn2_w);
|
||||
apply_field!(&format!("block_{i}_ffn2_b"), block.ffn2_b);
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// Tests
|
||||
// ============================================================================
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use super::super::jepa_vit::{CpuViTEncoder, JepaViTConfig};
|
||||
|
||||
fn tiny_config() -> JepaViTConfig {
|
||||
JepaViTConfig {
|
||||
embed_dim: 32,
|
||||
depth: 2,
|
||||
num_heads: 4,
|
||||
mlp_ratio: 2.0,
|
||||
patch_size: 16,
|
||||
image_size: 64,
|
||||
}
|
||||
}
|
||||
|
||||
fn empty_ckpt(step: usize) -> JepaModelCheckpoint {
|
||||
JepaModelCheckpoint {
|
||||
step,
|
||||
fields: Vec::new(),
|
||||
config_summary: String::new(),
|
||||
mean_loss: 0.0,
|
||||
ema_tau: 0.0,
|
||||
}
|
||||
}
|
||||
|
||||
fn single_field_ckpt(step: usize) -> JepaModelCheckpoint {
|
||||
JepaModelCheckpoint {
|
||||
step,
|
||||
fields: vec![WeightField {
|
||||
name: "test".to_string(),
|
||||
values: vec![1.0, 2.0, 3.0],
|
||||
}],
|
||||
config_summary: String::new(),
|
||||
mean_loss: 0.0,
|
||||
ema_tau: 0.0,
|
||||
}
|
||||
}
|
||||
|
||||
// 1. Magic bytes
|
||||
#[test]
|
||||
fn test_magic_bytes() {
|
||||
let ckpt = empty_ckpt(0);
|
||||
let bytes = serialize_checkpoint(&ckpt);
|
||||
assert_eq!(&bytes[0..4], b"JEPA", "first 4 bytes must be b\"JEPA\"");
|
||||
}
|
||||
|
||||
// 2. Version
|
||||
#[test]
|
||||
fn test_version() {
|
||||
let ckpt = empty_ckpt(0);
|
||||
let bytes = serialize_checkpoint(&ckpt);
|
||||
let v = u32::from_le_bytes([bytes[4], bytes[5], bytes[6], bytes[7]]);
|
||||
assert_eq!(v, 1, "version field must be 1");
|
||||
}
|
||||
|
||||
// 3. Round-trip: empty fields
|
||||
#[test]
|
||||
fn test_round_trip_empty_fields() {
|
||||
let ckpt = empty_ckpt(42);
|
||||
let bytes = serialize_checkpoint(&ckpt);
|
||||
let loaded = deserialize_checkpoint(&bytes).expect("deserialize should succeed");
|
||||
assert_eq!(loaded.step, 42);
|
||||
assert_eq!(loaded.fields.len(), 0);
|
||||
}
|
||||
|
||||
// 4. Round-trip: single field
|
||||
#[test]
|
||||
fn test_round_trip_single_field() {
|
||||
let ckpt = single_field_ckpt(7);
|
||||
let bytes = serialize_checkpoint(&ckpt);
|
||||
let loaded = deserialize_checkpoint(&bytes).expect("deserialize should succeed");
|
||||
assert_eq!(loaded.step, 7);
|
||||
assert_eq!(loaded.fields.len(), 1);
|
||||
assert_eq!(loaded.fields[0].name, "test");
|
||||
assert_eq!(loaded.fields[0].values, vec![1.0f32, 2.0, 3.0]);
|
||||
}
|
||||
|
||||
// 5. Round-trip: multiple fields
|
||||
#[test]
|
||||
fn test_round_trip_multiple_fields() {
|
||||
let fields: Vec<WeightField> = (0..5)
|
||||
.map(|i| WeightField {
|
||||
name: format!("field_{i}"),
|
||||
values: (0..i * 10 + 1).map(|j| j as f32 * 0.1).collect(),
|
||||
})
|
||||
.collect();
|
||||
let ckpt = JepaModelCheckpoint {
|
||||
step: 100,
|
||||
fields,
|
||||
config_summary: "test".to_string(),
|
||||
mean_loss: 1.23,
|
||||
ema_tau: 0.996,
|
||||
};
|
||||
let bytes = serialize_checkpoint(&ckpt);
|
||||
let loaded = deserialize_checkpoint(&bytes).expect("deserialize should succeed");
|
||||
assert_eq!(loaded.step, 100);
|
||||
assert_eq!(loaded.fields.len(), 5);
|
||||
for i in 0..5 {
|
||||
assert_eq!(loaded.fields[i].name, format!("field_{i}"));
|
||||
assert_eq!(loaded.fields[i].values.len(), i * 10 + 1);
|
||||
}
|
||||
}
|
||||
|
||||
// 6. Round-trip: large field
|
||||
#[test]
|
||||
fn test_round_trip_large_field() {
|
||||
let n = 50_000;
|
||||
let ckpt = JepaModelCheckpoint {
|
||||
step: 999,
|
||||
fields: vec![WeightField {
|
||||
name: "large".to_string(),
|
||||
values: (0..n).map(|i| i as f32 * 1e-5).collect(),
|
||||
}],
|
||||
config_summary: String::new(),
|
||||
mean_loss: 0.0,
|
||||
ema_tau: 0.0,
|
||||
};
|
||||
let bytes = serialize_checkpoint(&ckpt);
|
||||
let loaded = deserialize_checkpoint(&bytes).expect("large field should round-trip");
|
||||
assert_eq!(loaded.fields[0].values.len(), n);
|
||||
// Spot-check a few values
|
||||
assert!((loaded.fields[0].values[0] - 0.0).abs() < 1e-10);
|
||||
assert!((loaded.fields[0].values[1] - 1e-5).abs() < 1e-12);
|
||||
}
|
||||
|
||||
// 7. Invalid magic
|
||||
#[test]
|
||||
fn test_invalid_magic() {
|
||||
let ckpt = empty_ckpt(0);
|
||||
let mut bytes = serialize_checkpoint(&ckpt);
|
||||
bytes[0] = 0xFF; // corrupt first magic byte
|
||||
let result = deserialize_checkpoint(&bytes);
|
||||
assert!(
|
||||
matches!(result, Err(CheckpointError::InvalidMagic)),
|
||||
"expected InvalidMagic, got {result:?}"
|
||||
);
|
||||
}
|
||||
|
||||
// 8. Version mismatch
|
||||
#[test]
|
||||
fn test_version_mismatch() {
|
||||
let ckpt = empty_ckpt(0);
|
||||
let mut bytes = serialize_checkpoint(&ckpt);
|
||||
// Overwrite version bytes (offset 4..8) with version=2
|
||||
let v2 = 2u32.to_le_bytes();
|
||||
bytes[4] = v2[0];
|
||||
bytes[5] = v2[1];
|
||||
bytes[6] = v2[2];
|
||||
bytes[7] = v2[3];
|
||||
let result = deserialize_checkpoint(&bytes);
|
||||
assert!(
|
||||
matches!(result, Err(CheckpointError::VersionMismatch { found: 2, expected: 1 })),
|
||||
"expected VersionMismatch{{found:2, expected:1}}, got {result:?}"
|
||||
);
|
||||
}
|
||||
|
||||
// 9. Checksum corruption
|
||||
#[test]
|
||||
fn test_checksum_corruption() {
|
||||
let ckpt = single_field_ckpt(1);
|
||||
let mut bytes = serialize_checkpoint(&ckpt);
|
||||
// Flip the very last byte (checksum)
|
||||
let last = bytes.len() - 1;
|
||||
bytes[last] ^= 0xFF;
|
||||
let result = deserialize_checkpoint(&bytes);
|
||||
assert!(
|
||||
matches!(result, Err(CheckpointError::ChecksumMismatch)),
|
||||
"expected ChecksumMismatch, got {result:?}"
|
||||
);
|
||||
}
|
||||
|
||||
// 10. Save / load round-trip via temp file
|
||||
#[test]
|
||||
fn test_save_load_roundtrip() {
|
||||
let ckpt = JepaModelCheckpoint {
|
||||
step: 12345,
|
||||
fields: vec![
|
||||
WeightField { name: "w1".to_string(), values: vec![1.0, 2.0, 3.0] },
|
||||
WeightField { name: "w2".to_string(), values: vec![4.0, 5.0] },
|
||||
],
|
||||
config_summary: "{}".to_string(),
|
||||
mean_loss: 0.42,
|
||||
ema_tau: 0.999,
|
||||
};
|
||||
|
||||
let tmp = std::env::temp_dir().join("jepa_test_save_load.jepa");
|
||||
let path = tmp.to_str().expect("temp path is valid UTF-8");
|
||||
|
||||
save_checkpoint(&ckpt, path).expect("save_checkpoint should succeed");
|
||||
let loaded = load_checkpoint(path).expect("load_checkpoint should succeed");
|
||||
|
||||
assert_eq!(loaded.step, 12345);
|
||||
assert_eq!(loaded.fields.len(), 2);
|
||||
assert_eq!(loaded.fields[0].name, "w1");
|
||||
assert_eq!(loaded.fields[0].values, vec![1.0f32, 2.0, 3.0]);
|
||||
assert_eq!(loaded.fields[1].name, "w2");
|
||||
assert_eq!(loaded.fields[1].values, vec![4.0f32, 5.0]);
|
||||
|
||||
// Clean up
|
||||
let _ = std::fs::remove_file(path);
|
||||
}
|
||||
|
||||
// 11. Load nonexistent file → Err
|
||||
#[test]
|
||||
fn test_load_nonexistent_file() {
|
||||
let result = load_checkpoint("nonexistent_jepa_file_xyz_abc.jepa");
|
||||
assert!(result.is_err(), "loading a nonexistent file must return Err");
|
||||
}
|
||||
|
||||
// 12. encoder_to_fields produces "patch_embed" field
|
||||
#[test]
|
||||
fn test_encoder_to_fields_names() {
|
||||
let enc = CpuViTEncoder::new(tiny_config());
|
||||
let fields = encoder_to_fields(&enc);
|
||||
let names: Vec<&str> = fields.iter().map(|f| f.name.as_str()).collect();
|
||||
assert!(
|
||||
names.contains(&"patch_embed"),
|
||||
"fields must contain 'patch_embed'; got: {names:?}"
|
||||
);
|
||||
}
|
||||
|
||||
// 13. encoder_to_fields count (tiny has depth=2 → 3 + 8*2 = 19 fields)
|
||||
#[test]
|
||||
fn test_encoder_to_fields_count() {
|
||||
let enc = CpuViTEncoder::new(tiny_config());
|
||||
let fields = encoder_to_fields(&enc);
|
||||
// 3 base fields (patch_embed, proj_w, proj_b) + 8 per block × depth
|
||||
let expected = 3 + 8 * tiny_config().depth;
|
||||
assert_eq!(
|
||||
fields.len(),
|
||||
expected,
|
||||
"expected {expected} fields for tiny config (depth=2), got {}",
|
||||
fields.len()
|
||||
);
|
||||
}
|
||||
|
||||
// 14. apply_fields_to_encoder round-trip: extract from A, apply to B → same weights
|
||||
#[test]
|
||||
fn test_apply_fields_roundtrip() {
|
||||
let enc_a = CpuViTEncoder::new(tiny_config());
|
||||
let mut enc_b = CpuViTEncoder::new(tiny_config());
|
||||
|
||||
// Modify enc_b's patch_embed so they start different
|
||||
for v in enc_b.patch_embed.iter_mut() {
|
||||
*v = 0.0;
|
||||
}
|
||||
assert_ne!(enc_a.patch_embed, enc_b.patch_embed, "encoders should start different");
|
||||
|
||||
// Extract from A, apply to B
|
||||
let fields = encoder_to_fields(&enc_a);
|
||||
apply_fields_to_encoder(&mut enc_b, &fields).expect("apply_fields should succeed");
|
||||
|
||||
// Now B's patch_embed should match A's
|
||||
assert_eq!(
|
||||
enc_a.patch_embed, enc_b.patch_embed,
|
||||
"patch_embed should match after apply_fields"
|
||||
);
|
||||
}
|
||||
|
||||
// 15. apply_fields_to_encoder: missing required field → MissingField error
|
||||
#[test]
|
||||
fn test_apply_fields_missing() {
|
||||
let mut enc = CpuViTEncoder::new(tiny_config());
|
||||
// Provide only a subset of fields (missing patch_embed)
|
||||
let fields: Vec<WeightField> = Vec::new();
|
||||
let result = apply_fields_to_encoder(&mut enc, &fields);
|
||||
assert!(
|
||||
matches!(result, Err(CheckpointError::MissingField(ref name)) if name == "patch_embed"),
|
||||
"expected MissingField(\"patch_embed\"), got {result:?}"
|
||||
);
|
||||
}
|
||||
|
||||
// 16. apply_fields_to_encoder: wrong size → WrongSize error
|
||||
#[test]
|
||||
fn test_apply_fields_wrong_size() {
|
||||
let mut enc = CpuViTEncoder::new(tiny_config());
|
||||
// Build a mostly-correct field list but wrong size for patch_embed
|
||||
let mut fields = encoder_to_fields(&enc);
|
||||
// Corrupt patch_embed size
|
||||
for f in fields.iter_mut() {
|
||||
if f.name == "patch_embed" {
|
||||
f.values = vec![1.0, 2.0]; // wrong length
|
||||
break;
|
||||
}
|
||||
}
|
||||
let result = apply_fields_to_encoder(&mut enc, &fields);
|
||||
assert!(
|
||||
matches!(result, Err(CheckpointError::WrongSize { ref field, .. }) if field == "patch_embed"),
|
||||
"expected WrongSize for patch_embed, got {result:?}"
|
||||
);
|
||||
}
|
||||
|
||||
// 17. CheckpointError::Display is non-empty
|
||||
#[test]
|
||||
fn test_checkpoint_error_display() {
|
||||
let s = CheckpointError::InvalidMagic.to_string();
|
||||
assert!(!s.is_empty(), "Display for InvalidMagic should be non-empty");
|
||||
|
||||
let s2 = CheckpointError::VersionMismatch { found: 2, expected: 1 }.to_string();
|
||||
assert!(!s2.is_empty(), "Display for VersionMismatch should be non-empty");
|
||||
|
||||
let s3 = CheckpointError::ChecksumMismatch.to_string();
|
||||
assert!(!s3.is_empty(), "Display for ChecksumMismatch should be non-empty");
|
||||
|
||||
let s4 = CheckpointError::MissingField("foo".to_string()).to_string();
|
||||
assert!(!s4.is_empty(), "Display for MissingField should be non-empty");
|
||||
|
||||
let s5 = CheckpointError::Io("disk full".to_string()).to_string();
|
||||
assert!(!s5.is_empty(), "Display for Io should be non-empty");
|
||||
}
|
||||
|
||||
// 18. Step is preserved exactly
|
||||
#[test]
|
||||
fn test_step_preserved() {
|
||||
let ckpt = empty_ckpt(99999);
|
||||
let bytes = serialize_checkpoint(&ckpt);
|
||||
let loaded = deserialize_checkpoint(&bytes).expect("deserialize should succeed");
|
||||
assert_eq!(loaded.step, 99999, "step must be preserved exactly");
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user