[NSA] Avoid repeated NSA MQA logits memory queries (#25299)
This commit is contained in:
@@ -176,6 +176,12 @@ def rotate_activation(x: torch.Tensor) -> torch.Tensor:
|
||||
|
||||
|
||||
class Indexer(MultiPlatformOp):
|
||||
_MQA_LOGITS_BYTES_PER_ELEM = 4
|
||||
_MQA_LOGITS_STATIC_SKIP_ELEMS = 8_000_000
|
||||
_MQA_LOGITS_FREE_MEM_FRACTION = 0.5
|
||||
_MQA_LOGITS_TOTAL_MEM_FRACTION = 0.3
|
||||
_mqa_logits_budget_bytes: Dict[int, int] = {}
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
@@ -552,24 +558,59 @@ class Indexer(MultiPlatformOp):
|
||||
topk_result = torch.cat([topk_result, padding], dim=0)
|
||||
return topk_result
|
||||
|
||||
def _get_mqa_logits_budget_bytes(self, device_index: int) -> int:
|
||||
cached_budget = self._mqa_logits_budget_bytes.get(device_index)
|
||||
if cached_budget is not None:
|
||||
return cached_budget
|
||||
|
||||
total_mem = torch.cuda.get_device_properties(device_index).total_memory
|
||||
|
||||
total_mem_budget = int(total_mem * self._MQA_LOGITS_TOTAL_MEM_FRACTION)
|
||||
mem_fraction_static = get_global_server_args().mem_fraction_static
|
||||
if mem_fraction_static is None:
|
||||
static_budget = total_mem_budget
|
||||
else:
|
||||
static_free_mem = int(total_mem * max(0.0, 1.0 - mem_fraction_static))
|
||||
static_budget = min(
|
||||
int(static_free_mem * self._MQA_LOGITS_FREE_MEM_FRACTION),
|
||||
total_mem_budget,
|
||||
)
|
||||
static_budget = max(1, static_budget)
|
||||
|
||||
# Keep the static serving-memory guard during CUDA graph capture without
|
||||
# caching it. The first non-capture prefill path will cache the real
|
||||
# free-memory budget below.
|
||||
if get_is_capture_mode():
|
||||
return static_budget
|
||||
|
||||
# Match the original free-memory guard: logits_bytes * 2 > free_mem.
|
||||
# torch.cuda.mem_get_info synchronizes the host, so cache the result,
|
||||
# capped by the workload-independent serving-memory headroom.
|
||||
free_mem, _ = torch.cuda.mem_get_info(device_index)
|
||||
budget_bytes = min(
|
||||
int(free_mem * self._MQA_LOGITS_FREE_MEM_FRACTION), static_budget
|
||||
)
|
||||
|
||||
budget_bytes = max(1, budget_bytes)
|
||||
self._mqa_logits_budget_bytes[device_index] = budget_bytes
|
||||
return budget_bytes
|
||||
|
||||
def _should_chunk_mqa_logits(
|
||||
self, num_q: int, num_k: int, device: torch.device
|
||||
self, num_q: int, num_k: int, device_index: int
|
||||
) -> Tuple[bool, int]:
|
||||
"""
|
||||
Detect whether we need to chunk the MQA logits computation to avoid OOM
|
||||
Return: (need_chunk, free_mem)
|
||||
Return: (need_chunk, logits_budget_bytes)
|
||||
"""
|
||||
# Quick static check for normal batches
|
||||
if num_q * num_k < 8_000_000: # 8M elements ≈ 32MB logits
|
||||
if num_q * num_k < self._MQA_LOGITS_STATIC_SKIP_ELEMS:
|
||||
return False, 0
|
||||
|
||||
free_mem, total_mem = torch.cuda.mem_get_info(device)
|
||||
bytes_per_elem = 4 # float32
|
||||
logits_bytes = num_q * num_k * bytes_per_elem
|
||||
logits_bytes = num_q * num_k * self._MQA_LOGITS_BYTES_PER_ELEM
|
||||
logits_budget_bytes = self._get_mqa_logits_budget_bytes(device_index)
|
||||
|
||||
# Logits should not exceed 50% of free memory or 30% of total memory
|
||||
need_chunk = (logits_bytes * 2 > free_mem) or (logits_bytes > total_mem * 0.3)
|
||||
return need_chunk, free_mem
|
||||
need_chunk = logits_bytes > logits_budget_bytes
|
||||
return need_chunk, logits_budget_bytes
|
||||
|
||||
def _get_topk_ragged(
|
||||
self,
|
||||
@@ -618,6 +659,8 @@ class Indexer(MultiPlatformOp):
|
||||
batch_size = len(block_tables)
|
||||
token_nums, _, _ = q_fp8.shape
|
||||
device = q_fp8.device
|
||||
device_index = device.index
|
||||
assert device_index is not None, "q_fp8 must be on an indexed CUDA device"
|
||||
|
||||
topk_result = torch.full(
|
||||
(token_nums, self.index_topk), -1, device=device, dtype=torch.int32
|
||||
@@ -650,7 +693,9 @@ class Indexer(MultiPlatformOp):
|
||||
token_to_batch_idx = metadata.get_token_to_batch_idx()
|
||||
q_offset = ks.shape[0]
|
||||
k_offset = k_fp8.shape[0]
|
||||
need_chunk, free_mem = self._should_chunk_mqa_logits(q_offset, k_offset, device)
|
||||
need_chunk, logits_budget_bytes = self._should_chunk_mqa_logits(
|
||||
q_offset, k_offset, device_index
|
||||
)
|
||||
|
||||
if not need_chunk:
|
||||
assert q_fp8[:q_offset].shape[0] != 0
|
||||
@@ -678,14 +723,14 @@ class Indexer(MultiPlatformOp):
|
||||
topk_result[:q_offset] = raw_topk_result
|
||||
return topk_result
|
||||
|
||||
# Chunk path
|
||||
bytes_per_elem = 4 # float32
|
||||
bytes_per_row = k_offset * bytes_per_elem
|
||||
# Reserve 50% of free memory for logits
|
||||
max_rows = max(1, int((free_mem * 0.5) // max(bytes_per_row, 1)))
|
||||
bytes_per_row = k_offset * self._MQA_LOGITS_BYTES_PER_ELEM
|
||||
max_rows = max(1, int(logits_budget_bytes // max(bytes_per_row, 1)))
|
||||
max_rows = min(max_rows, q_offset)
|
||||
|
||||
global_topk_offset = metadata.attn_metadata.topk_indices_offset
|
||||
cu_seqlens_q_full = None
|
||||
if global_topk_offset is None:
|
||||
cu_seqlens_q_full = torch.ones(q_offset, dtype=torch.int32, device=device)
|
||||
|
||||
assert (
|
||||
seq_lens_expanded.shape[0] == q_offset
|
||||
@@ -733,10 +778,7 @@ class Indexer(MultiPlatformOp):
|
||||
else:
|
||||
# PAGED path: treat each token as a length-1 sequence
|
||||
topk_offset_chunk = None
|
||||
B_chunk = logits_chunk.shape[0]
|
||||
cu_seqlens_q_chunk = torch.ones(
|
||||
B_chunk, dtype=torch.int32, device=device
|
||||
)
|
||||
cu_seqlens_q_chunk = cu_seqlens_q_full[start:end]
|
||||
batch_idx_chunk = token_to_batch_idx[start:end]
|
||||
|
||||
raw_topk_chunk = metadata.topk_transform(
|
||||
|
||||
Reference in New Issue
Block a user