Mike0021 commited on
Commit
7dfb205
·
verified ·
1 Parent(s): 492f937

functional index_copy fix

Browse files
Files changed (1) hide show
  1. 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.index_copy(
92
- torch.zeros((batch, heads, padded_len, dim), dtype=query.dtype, device=query.device),
93
- 2, untile_index, query.transpose(1, 2),
94
- )
95
- key_tiled = torch.index_copy(
96
- torch.zeros_like(query_tiled), 2, untile_index, key.transpose(1, 2),
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