| from pathlib import Path |
| from types import SimpleNamespace |
| import sys |
|
|
| import pytest |
| import torch |
|
|
| sys.path.insert(0, str(Path(__file__).resolve().parents[3])) |
|
|
| from methods.prunning.SiTo.runtime import SiToTokenPruner |
|
|
|
|
| def test_sito_per_frame_roundtrip_preserves_shape_and_kept_tokens(): |
| pruner = SiToTokenPruner(group_mode="per_frame", prune_ratio=0.20, patch_h=2, patch_w=2) |
| hidden_states = torch.randn(2, 3 * 4 * 4, 8) |
| plan = pruner.prepare(hidden_states, video_size=SimpleNamespace(T=3, H=4, W=4)) |
|
|
| assert plan is not None |
| pruned = pruner.prune(hidden_states, plan) |
| recovered = pruner.recover(pruned, plan) |
|
|
| assert recovered.shape == hidden_states.shape |
| assert torch.allclose(recovered[:, plan.keep_indices], hidden_states[:, plan.keep_indices]) |
|
|
|
|
| def test_sito_full_2d_roundtrip_supports_rectangular_grids(): |
| pruner = SiToTokenPruner(group_mode="full_2d", prune_ratio=0.35, patch_h=2, patch_w=2) |
| hidden_states = torch.randn(2, 18 * 10, 16) |
| plan = pruner.prepare(hidden_states, video_size=SimpleNamespace(T=1, H=18, W=10)) |
|
|
| assert plan is not None |
| pruned = pruner.prune(hidden_states, plan) |
| recovered = pruner.recover(pruned, plan) |
|
|
| assert pruned.shape[1] < hidden_states.shape[1] |
| assert recovered.shape == hidden_states.shape |
|
|
|
|
| def test_sito_invalid_prune_ratio_raises(): |
| with pytest.raises(ValueError, match="prune_ratio"): |
| SiToTokenPruner(group_mode="per_frame", prune_ratio=0.80, patch_h=2, patch_w=2) |
|
|