Fix NSA indexer in the nightly test (#17452)
This commit is contained in:
@@ -79,6 +79,22 @@ class MockIndexerMetadata(BaseIndexerMetadata):
|
|||||||
)
|
)
|
||||||
return page_table
|
return page_table
|
||||||
|
|
||||||
|
def get_page_table_1(self) -> torch.Tensor:
|
||||||
|
"""Return: (batch_size, num_blocks) int32, page table with page size 1."""
|
||||||
|
# Create a simple page table for testing with page size 1
|
||||||
|
max_seq_len = max(self.seq_lens)
|
||||||
|
num_blocks = max_seq_len # Page size 1 means num_blocks == max_seq_len
|
||||||
|
page_table = torch.zeros(
|
||||||
|
(self.batch_size, num_blocks), dtype=torch.int32, device=self.device
|
||||||
|
)
|
||||||
|
for i in range(self.batch_size):
|
||||||
|
# Simple linear mapping: block i maps to page i
|
||||||
|
num_blocks_needed = self.seq_lens[i]
|
||||||
|
page_table[i, :num_blocks_needed] = torch.arange(
|
||||||
|
num_blocks_needed, device=self.device
|
||||||
|
)
|
||||||
|
return page_table
|
||||||
|
|
||||||
def get_seqlens_expanded(self) -> torch.Tensor:
|
def get_seqlens_expanded(self) -> torch.Tensor:
|
||||||
"""Return: (sum_extend_seq_len,) int32 tensor"""
|
"""Return: (sum_extend_seq_len,) int32 tensor"""
|
||||||
# For extend mode, each new token attends to progressively more tokens
|
# For extend mode, each new token attends to progressively more tokens
|
||||||
|
|||||||
Reference in New Issue
Block a user