Skip to content

fix: bound QSA keys by the key length under context parallelism - #222

Open
shiaho777 wants to merge 1 commit into
modelscope:mainfrom
shiaho777:fix/qsa-cp-query-length
Open

shiaho777 wants to merge 1 commit into
modelscope:mainfrom
shiaho777:fix/qsa-cp-query-length

Conversation

@shiaho777

Copy link
Copy Markdown
Contributor

The sparse kernel used the query count as the key bound (offs_k < T) and built the block bitmap with ceil(T / block_size). Context parallelism holds a 1/cp_size query shard and the full key sequence. The CP path therefore copied the local queries into a buffer as long as the keys and filled the rest with -1. Those rows launch tiles and allocate a bitmap for the whole sequence, then the kernel skips them.

The key bound and the bitmap width now follow k.shape[0]. The query grid still follows the query count. The CP path passes the local queries and the gathered keys directly. When the two lengths are equal, the bounds match the previous kernel.

Checked with python3 -m pytest tests/test_qsa_key_length.py. Two queries and eight keys, block size 4, flag key index 6 in block 1. flake8 is clean. The Triton kernels were not executed here; triton is not installed.

The sparse kernel used the query count as the key bound (offs_k < T) and built the block bitmap with ceil(T / block_size). Context parallelism holds a 1/cp_size query shard and the full key sequence. The CP path therefore copied the local queries into a buffer as long as the keys and filled the rest with -1. Those rows launch tiles and allocate a bitmap for the whole sequence, then the kernel skips them.

The key bound and the bitmap width now follow k.shape[0]. The query grid still follows the query count. The CP path passes the local queries and the gathered keys directly. When the two lengths are equal, the bounds match the previous kernel.

Checked with python3 -m pytest tests/test_qsa_key_length.py. Two queries and eight keys, block size 4, flag key index 6 in block 1. flake8 is clean. The Triton kernels were not executed here; triton is not installed.
big = torch.iinfo(torch.int32).max
lo = torch.where(valid, indices, torch.full_like(indices, big)).min(dim=1).values.to(torch.int32)
hi = torch.where(valid, indices, torch.full_like(indices, -1)).max(dim=1).values.to(torch.int32)
zeros = torch.zeros(T, dtype=torch.int32, device=q.device)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

NameError: name 'T' is not defined

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants