functional name fix
Browse files- aoti_attention.py +1 -0
aoti_attention.py
CHANGED
|
@@ -102,6 +102,7 @@ def sparse_attention_functional(
|
|
| 102 |
# Prefix-exempt top-k mask, assembled from fresh tensors (functional, export-traceable):
|
| 103 |
# video query rows keep the top-k VIDEO kv tiles, prefix kv tiles are always selected,
|
| 104 |
# prefix query rows are dense over everything.
|
|
|
|
| 105 |
video_kv_scores = scores[..., num_prefix_tiles:, num_prefix_tiles:] # [B, H, nq_video, kv_video]
|
| 106 |
indices = video_kv_scores.topk(topk, dim=-1, sorted=False).indices
|
| 107 |
mask_video_kv = torch.zeros_like(video_kv_scores, dtype=torch.bool).scatter(-1, indices, True)
|
|
|
|
| 102 |
# Prefix-exempt top-k mask, assembled from fresh tensors (functional, export-traceable):
|
| 103 |
# video query rows keep the top-k VIDEO kv tiles, prefix kv tiles are always selected,
|
| 104 |
# prefix query rows are dense over everything.
|
| 105 |
+
num_q_video = n_tiles - num_prefix_tiles
|
| 106 |
video_kv_scores = scores[..., num_prefix_tiles:, num_prefix_tiles:] # [B, H, nq_video, kv_video]
|
| 107 |
indices = video_kv_scores.topk(topk, dim=-1, sorted=False).indices
|
| 108 |
mask_video_kv = torch.zeros_like(video_kv_scores, dtype=torch.bool).scatter(-1, indices, True)
|