import os import tempfile import unittest from pathlib import Path import torch from torch import nn from experiments.exp_base import load_custom_checkpoint class _LightningLoopMetadata: def __init__(self, phase): self.phase = phase class DummyFrameMemoryBlock(nn.Module): def __init__(self, hidden_size=2): super().__init__() self.r_adaLN_modulation = nn.Sequential(nn.SiLU(), nn.Linear(hidden_size, 6 * hidden_size)) self.r_mlp = nn.Linear(hidden_size, hidden_size) self.r_attn_anchor = nn.Linear(hidden_size, hidden_size, bias=False) self.r_attn_dynamic = nn.Linear(hidden_size, hidden_size, bias=False) self.r_attn_revisit = nn.Linear(hidden_size, hidden_size, bias=False) class DummyDeMemWM(nn.Module): def __init__(self): super().__init__() self.backbone = nn.Linear(2, 2) self.blocks = nn.ModuleList([DummyFrameMemoryBlock()]) def _ones_state(module, skip_frame_memory=False): state = {} memory_markers = ( "r_attn_anchor", "r_attn_dynamic", "r_attn_revisit", "r_adaLN_modulation", "r_mlp", ) for key, value in module.state_dict().items(): if skip_frame_memory and any(marker in key for marker in memory_markers): continue state[key] = torch.ones_like(value) return state class CheckpointLoadingTests(unittest.TestCase): def test_prefix_matching_loads_full_checkpoint_into_submodule(self): target = nn.Linear(2, 2) state = {f"diffusion_model.model.{key}": torch.full_like(value, 2.0) for key, value in target.state_dict().items()} with tempfile.TemporaryDirectory() as tmpdir: ckpt_path = Path(tmpdir) / "prefixed.ckpt" torch.save({"state_dict": state}, ckpt_path) load_custom_checkpoint(target, ckpt_path) self.assertTrue(torch.equal(target.weight, torch.full_like(target.weight, 2.0))) self.assertTrue(torch.equal(target.bias, torch.full_like(target.bias, 2.0))) def test_lightning_ckpt_pickle_metadata_still_loads_state_dict(self): target = nn.Linear(2, 2) state = {key: torch.full_like(value, 5.0) for key, value in target.state_dict().items()} with tempfile.TemporaryDirectory() as tmpdir: ckpt_path = Path(tmpdir) / "lightning_metadata.ckpt" torch.save( { "state_dict": state, "loops": {"fit": _LightningLoopMetadata("fit")}, }, ckpt_path, ) load_custom_checkpoint(target, ckpt_path) self.assertTrue(torch.equal(target.weight, torch.full_like(target.weight, 5.0))) self.assertTrue(torch.equal(target.bias, torch.full_like(target.bias, 5.0))) def test_shape_mismatch_is_filtered_while_compatible_keys_load(self): target = nn.Linear(2, 2) original_weight = target.weight.detach().clone() state = { "weight": torch.ones((3, 2)), "bias": torch.full_like(target.bias, 4.0), } with tempfile.TemporaryDirectory() as tmpdir: ckpt_path = Path(tmpdir) / "shape_mismatch.pt" torch.save(state, ckpt_path) load_custom_checkpoint(target, ckpt_path) self.assertTrue(torch.equal(target.weight, original_weight)) self.assertTrue(torch.equal(target.bias, torch.full_like(target.bias, 4.0))) def test_base_checkpoint_missing_frame_memory_keys_loads_and_zeroes_gates(self): target = DummyDeMemWM() state = _ones_state(target, skip_frame_memory=True) with tempfile.TemporaryDirectory() as tmpdir: ckpt_path = Path(tmpdir) / "base.ckpt" torch.save({"state_dict": state}, ckpt_path) load_custom_checkpoint(target, ckpt_path) gate = target.blocks[0].r_adaLN_modulation[1] chunk = gate.weight.shape[0] // 6 self.assertTrue(torch.equal(target.backbone.weight, torch.ones_like(target.backbone.weight))) self.assertTrue(torch.equal(gate.weight[2 * chunk:3 * chunk], torch.zeros_like(gate.weight[2 * chunk:3 * chunk]))) self.assertTrue(torch.equal(gate.weight[5 * chunk:6 * chunk], torch.zeros_like(gate.weight[5 * chunk:6 * chunk]))) self.assertTrue(torch.equal(gate.bias[2 * chunk:3 * chunk], torch.zeros_like(gate.bias[2 * chunk:3 * chunk]))) self.assertTrue(torch.equal(gate.bias[5 * chunk:6 * chunk], torch.zeros_like(gate.bias[5 * chunk:6 * chunk]))) def test_incomplete_dememwm_checkpoint_missing_stream_keys_fails(self): target = DummyDeMemWM() state = _ones_state(target, skip_frame_memory=True) for key, value in target.state_dict().items(): if "r_attn_anchor" in key or "r_adaLN_modulation" in key or "r_mlp" in key: state[key] = torch.ones_like(value) with tempfile.TemporaryDirectory() as tmpdir: ckpt_path = Path(tmpdir) / "incomplete_dememwm.ckpt" torch.save({"state_dict": state}, ckpt_path) with self.assertRaisesRegex(RuntimeError, "Incomplete DeMemWM checkpoint.*r_attn_dynamic"): load_custom_checkpoint(target, ckpt_path) def test_incomplete_dememwm_checkpoint_missing_adaln_mlp_fails(self): target = DummyDeMemWM() state = _ones_state(target, skip_frame_memory=True) for key, value in target.state_dict().items(): if "r_attn_" in key: state[key] = torch.ones_like(value) with tempfile.TemporaryDirectory() as tmpdir: ckpt_path = Path(tmpdir) / "incomplete_dememwm_missing_adaln_mlp.ckpt" torch.save({"state_dict": state}, ckpt_path) with self.assertRaisesRegex(RuntimeError, "Incomplete DeMemWM checkpoint.*r_adaLN_modulation"): load_custom_checkpoint(target, ckpt_path) def test_incomplete_dememwm_checkpoint_missing_r_mlp_fails(self): target = DummyDeMemWM() state = _ones_state(target, skip_frame_memory=True) for key, value in target.state_dict().items(): if "r_attn_" in key or "r_adaLN_modulation" in key: state[key] = torch.ones_like(value) with tempfile.TemporaryDirectory() as tmpdir: ckpt_path = Path(tmpdir) / "incomplete_dememwm_missing_r_mlp.ckpt" torch.save({"state_dict": state}, ckpt_path) with self.assertRaisesRegex(RuntimeError, "Incomplete DeMemWM checkpoint.*r_mlp"): load_custom_checkpoint(target, ckpt_path) def test_directory_load_skips_unreadable_latest_checkpoint(self): target = nn.Linear(2, 2) state = {key: torch.full_like(value, 3.0) for key, value in target.state_dict().items()} with tempfile.TemporaryDirectory() as tmpdir: tmpdir = Path(tmpdir) valid_ckpt = tmpdir / "epoch0_step1.ckpt" broken_ckpt = tmpdir / "epoch0_step2.ckpt" torch.save({"state_dict": state}, valid_ckpt) broken_ckpt.write_bytes(b"not a checkpoint") os.utime(broken_ckpt, (valid_ckpt.stat().st_mtime + 10, valid_ckpt.stat().st_mtime + 10)) load_custom_checkpoint(target, tmpdir) self.assertTrue(torch.equal(target.weight, torch.full_like(target.weight, 3.0))) self.assertTrue(torch.equal(target.bias, torch.full_like(target.bias, 3.0))) if __name__ == "__main__": unittest.main()