use torsh_cli::commands::model::pytorch_reader::{read_state_dict, TensorDType, TensorData};
use torsh_cli::commands::real_training::{
build_mlp, evaluate_loss, synthetic_regression, train_regression, MlpConfig,
};
#[test]
fn f005_real_training_decreases_loss_on_synthetic_regression() {
let cfg = MlpConfig {
input_dim: 8,
hidden_dim: 16,
output_dim: 1,
};
let model = build_mlp(&cfg).expect("build MLP");
let data =
synthetic_regression(64, cfg.input_dim, cfg.output_dim, 0xC0FF_EE01).expect("dataset");
let initial = evaluate_loss(&model, &data).expect("initial loss");
let history = train_regression(&model, &data, 0.05, 20).expect("training");
assert_eq!(history.len(), 20, "one loss reading per step");
let final_loss = *history.last().expect("final loss");
assert!(
final_loss < initial,
"real training must reduce the loss: initial={initial}, final={final_loss}"
);
assert!(
final_loss.is_finite(),
"loss must be finite, got {final_loss}"
);
assert!(
final_loss >= 0.0,
"MSE loss cannot be negative, got {final_loss}"
);
}
#[test]
fn f005_training_history_is_deterministic_not_random() {
let cfg = MlpConfig {
input_dim: 4,
hidden_dim: 8,
output_dim: 2,
};
let data = synthetic_regression(32, cfg.input_dim, cfg.output_dim, 42).expect("dataset");
let model_a = build_mlp(&cfg).expect("model a");
let model_b = build_mlp(&cfg).expect("model b");
let hist_a = train_regression(&model_a, &data, 0.05, 30).expect("train a");
let hist_b = train_regression(&model_b, &data, 0.05, 30).expect("train b");
let init = evaluate_loss(&build_mlp(&cfg).expect("m"), &data).expect("init");
assert!(hist_a.last().unwrap() < &init);
assert!(hist_b.last().unwrap() < &init);
}
fn push_stored_entry(out: &mut Vec<u8>, name: &str, body: &[u8]) {
out.extend_from_slice(b"PK\x03\x04");
out.extend_from_slice(&20u16.to_le_bytes()); out.extend_from_slice(&0u16.to_le_bytes()); out.extend_from_slice(&0u16.to_le_bytes()); out.extend_from_slice(&0u16.to_le_bytes()); out.extend_from_slice(&0u16.to_le_bytes()); out.extend_from_slice(&0u32.to_le_bytes()); out.extend_from_slice(&(body.len() as u32).to_le_bytes()); out.extend_from_slice(&(body.len() as u32).to_le_bytes()); out.extend_from_slice(&(name.len() as u16).to_le_bytes());
out.extend_from_slice(&0u16.to_le_bytes()); out.extend_from_slice(name.as_bytes());
out.extend_from_slice(body);
}
fn crafted_pickle() -> Vec<u8> {
let mut p = Vec::new();
p.extend_from_slice(&[0x80, 0x02]); p.push(b'}'); p.extend_from_slice(&[b'q', 0x00]); p.push(b'('); p.extend_from_slice(&[0x8c, 0x01]); p.push(b'w'); p.extend_from_slice(b"ctorch._utils\n_rebuild_tensor_v2\n"); p.push(b'('); p.push(b'('); p.extend_from_slice(&[0x8c, 0x07]);
p.extend_from_slice(b"storage");
p.extend_from_slice(b"ctorch\nFloatStorage\n"); p.extend_from_slice(&[0x8c, 0x01]);
p.push(b'0'); p.extend_from_slice(&[0x8c, 0x03]);
p.extend_from_slice(b"cpu");
p.extend_from_slice(&[b'K', 0x06]); p.push(b't'); p.push(b'Q'); p.extend_from_slice(&[b'K', 0x00]); p.push(b'('); p.extend_from_slice(&[b'K', 0x02, b'K', 0x03]);
p.push(b't');
p.push(b'('); p.extend_from_slice(&[b'K', 0x03, b'K', 0x01]);
p.push(b't');
p.push(0x89); p.push(b'N'); p.push(b't'); p.push(b'R'); p.push(b'u'); p.push(b'.'); p
}
#[test]
fn pytorch_reader_reconstructs_real_tensor_values() {
let mut storage = Vec::new();
for v in [1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0] {
storage.extend_from_slice(&v.to_le_bytes());
}
let mut zip = Vec::new();
push_stored_entry(&mut zip, "archive/data.pkl", &crafted_pickle());
push_stored_entry(&mut zip, "archive/data/0", &storage);
let tensors = read_state_dict(&zip).expect("reader must reconstruct the checkpoint");
assert_eq!(tensors.len(), 1, "one tensor in the state_dict");
let t = &tensors[0];
assert_eq!(t.name, "w");
assert_eq!(t.dtype, TensorDType::F32);
assert_eq!(t.shape, vec![2, 3]);
match &t.data {
TensorData::F32(values) => {
assert_eq!(values, &vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]);
}
other => panic!("expected F32 data, got {other:?}"),
}
}
#[test]
fn pytorch_reader_rejects_non_zip_input() {
let err = read_state_dict(b"\x80\x02}q\x00.").unwrap_err();
assert!(
err.to_string().contains("zip-based"),
"unexpected error: {err}"
);
}