functional index_copy fix
Browse files- aoti_attention.py +6 -10
aoti_attention.py
CHANGED
|
@@ -88,16 +88,12 @@ def sparse_attention_functional(
|
|
| 88 |
num_q_video = n_tiles - num_prefix_tiles
|
| 89 |
|
| 90 |
# Tile: scatter the packed rows into the kernels' `[B, H, S_pad, D]` layout, pad slots zero.
|
| 91 |
-
query_tiled = torch.
|
| 92 |
-
|
| 93 |
-
|
| 94 |
-
)
|
| 95 |
-
|
| 96 |
-
|
| 97 |
-
)
|
| 98 |
-
value_tiled = torch.index_copy(
|
| 99 |
-
torch.zeros_like(query_tiled), 2, untile_index, value.transpose(1, 2),
|
| 100 |
-
)
|
| 101 |
|
| 102 |
# Per-head fp32 pooled tile scores.
|
| 103 |
q_pool = query_tiled.view(batch, heads, n_tiles, 64, dim).sum(dim=3, dtype=torch.float32) / tile_divisor
|
|
|
|
| 88 |
num_q_video = n_tiles - num_prefix_tiles
|
| 89 |
|
| 90 |
# Tile: scatter the packed rows into the kernels' `[B, H, S_pad, D]` layout, pad slots zero.
|
| 91 |
+
query_tiled = torch.zeros((batch, heads, padded_len, dim), dtype=query.dtype, device=query.device)
|
| 92 |
+
query_tiled.index_copy_(2, untile_index, query.transpose(1, 2))
|
| 93 |
+
key_tiled = torch.zeros_like(query_tiled)
|
| 94 |
+
key_tiled.index_copy_(2, untile_index, key.transpose(1, 2))
|
| 95 |
+
value_tiled = torch.zeros_like(query_tiled)
|
| 96 |
+
value_tiled.index_copy_(2, untile_index, value.transpose(1, 2))
|
|
|
|
|
|
|
|
|
|
|
|
|
| 97 |
|
| 98 |
# Per-head fp32 pooled tile scores.
|
| 99 |
q_pool = query_tiled.view(batch, heads, n_tiles, 64, dim).sum(dim=3, dtype=torch.float32) / tile_divisor
|