| 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() |
|
|