Skip to content

Commit 330d206

Browse files
committed
fix: align dict keys with framework naming and remove invalid lru_cache
- Rename build_vsa_metadata() keys to tile_partition_indices / reverse_tile_partition_indices to match VideoSparseAttentionMetadata - Remove @lru_cache from get_non_pad_index (torch.Tensor is unhashable) - Update tests accordingly
1 parent e5e3485 commit 330d206

2 files changed

Lines changed: 8 additions & 9 deletions

File tree

fastvideo-kernel/python/fastvideo_kernel/vsa_utils.py

Lines changed: 3 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -85,7 +85,6 @@ def _sizes(dim_len: int, tile: int, n: int) -> torch.LongTensor:
8585
).reshape(-1)
8686

8787

88-
@functools.lru_cache(maxsize=10)
8988
def get_non_pad_index(
9089
variable_block_sizes: torch.LongTensor,
9190
max_block_size: int,
@@ -116,7 +115,7 @@ def build_vsa_metadata(
116115
device: Target device for index tensors.
117116
118117
Returns:
119-
Dict with keys: tile_indices, reverse_tile_indices,
118+
Dict with keys: tile_partition_indices, reverse_tile_partition_indices,
120119
variable_block_sizes, non_pad_index, num_tiles, max_block_size.
121120
"""
122121
if isinstance(device, str):
@@ -137,8 +136,8 @@ def build_vsa_metadata(
137136
npi = get_non_pad_index(vbs, max_block_size)
138137

139138
return {
140-
"tile_indices": tile_indices,
141-
"reverse_tile_indices": reverse_tile_indices,
139+
"tile_partition_indices": tile_indices,
140+
"reverse_tile_partition_indices": reverse_tile_indices,
142141
"variable_block_sizes": vbs,
143142
"non_pad_index": npi,
144143
"num_tiles": num_tiles,

fastvideo-kernel/tests/test_vsa_utils.py

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -160,7 +160,7 @@ def test_all_keys_present(self):
160160
"""build_vsa_metadata returns all expected keys."""
161161
meta = build_vsa_metadata((8, 16, 16), device="cpu")
162162
expected_keys = {
163-
"tile_indices", "reverse_tile_indices",
163+
"tile_partition_indices", "reverse_tile_partition_indices",
164164
"variable_block_sizes", "non_pad_index",
165165
"num_tiles", "max_block_size",
166166
}
@@ -169,8 +169,8 @@ def test_all_keys_present(self):
169169
def test_types(self):
170170
"""Return types are correct."""
171171
meta = build_vsa_metadata((8, 16, 16), device="cpu")
172-
assert isinstance(meta["tile_indices"], torch.Tensor)
173-
assert isinstance(meta["reverse_tile_indices"], torch.Tensor)
172+
assert isinstance(meta["tile_partition_indices"], torch.Tensor)
173+
assert isinstance(meta["reverse_tile_partition_indices"], torch.Tensor)
174174
assert isinstance(meta["variable_block_sizes"], torch.Tensor)
175175
assert isinstance(meta["non_pad_index"], torch.Tensor)
176176
assert isinstance(meta["num_tiles"], tuple)
@@ -191,8 +191,8 @@ def test_consistency(self):
191191
shape = (8, 16, 16)
192192
meta = build_vsa_metadata(shape, device="cpu")
193193
n = math.prod(shape)
194-
assert meta["tile_indices"].shape == (n,)
195-
assert meta["reverse_tile_indices"].shape == (n,)
194+
assert meta["tile_partition_indices"].shape == (n,)
195+
assert meta["reverse_tile_partition_indices"].shape == (n,)
196196
assert meta["variable_block_sizes"].sum().item() == n
197197
assert meta["non_pad_index"].shape[0] == n
198198

0 commit comments

Comments
 (0)