Mike0021 commited on
Commit
883dc36
·
verified ·
1 Parent(s): cf145d5

functional name fix

Browse files
Files changed (1) hide show
  1. 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)