File size: 1,488 Bytes
ec0a9aa
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
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)