[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):
|
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__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
hidden_size: int,
|
hidden_size: int,
|
||||||
@@ -552,24 +558,59 @@ class Indexer(MultiPlatformOp):
|
|||||||
topk_result = torch.cat([topk_result, padding], dim=0)
|
topk_result = torch.cat([topk_result, padding], dim=0)
|
||||||
return topk_result
|
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(
|
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]:
|
) -> Tuple[bool, int]:
|
||||||
"""
|
"""
|
||||||
Detect whether we need to chunk the MQA logits computation to avoid OOM
|
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
|
# 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
|
return False, 0
|
||||||
|
|
||||||
free_mem, total_mem = torch.cuda.mem_get_info(device)
|
logits_bytes = num_q * num_k * self._MQA_LOGITS_BYTES_PER_ELEM
|
||||||
bytes_per_elem = 4 # float32
|
logits_budget_bytes = self._get_mqa_logits_budget_bytes(device_index)
|
||||||
logits_bytes = num_q * num_k * bytes_per_elem
|
|
||||||
|
|
||||||
# Logits should not exceed 50% of free memory or 30% of total memory
|
need_chunk = logits_bytes > logits_budget_bytes
|
||||||
need_chunk = (logits_bytes * 2 > free_mem) or (logits_bytes > total_mem * 0.3)
|
return need_chunk, logits_budget_bytes
|
||||||
return need_chunk, free_mem
|
|
||||||
|
|
||||||
def _get_topk_ragged(
|
def _get_topk_ragged(
|
||||||
self,
|
self,
|
||||||
@@ -618,6 +659,8 @@ class Indexer(MultiPlatformOp):
|
|||||||
batch_size = len(block_tables)
|
batch_size = len(block_tables)
|
||||||
token_nums, _, _ = q_fp8.shape
|
token_nums, _, _ = q_fp8.shape
|
||||||
device = q_fp8.device
|
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(
|
topk_result = torch.full(
|
||||||
(token_nums, self.index_topk), -1, device=device, dtype=torch.int32
|
(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()
|
token_to_batch_idx = metadata.get_token_to_batch_idx()
|
||||||
q_offset = ks.shape[0]
|
q_offset = ks.shape[0]
|
||||||
k_offset = k_fp8.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:
|
if not need_chunk:
|
||||||
assert q_fp8[:q_offset].shape[0] != 0
|
assert q_fp8[:q_offset].shape[0] != 0
|
||||||
@@ -678,14 +723,14 @@ class Indexer(MultiPlatformOp):
|
|||||||
topk_result[:q_offset] = raw_topk_result
|
topk_result[:q_offset] = raw_topk_result
|
||||||
return topk_result
|
return topk_result
|
||||||
|
|
||||||
# Chunk path
|
bytes_per_row = k_offset * self._MQA_LOGITS_BYTES_PER_ELEM
|
||||||
bytes_per_elem = 4 # float32
|
max_rows = max(1, int(logits_budget_bytes // max(bytes_per_row, 1)))
|
||||||
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)))
|
|
||||||
max_rows = min(max_rows, q_offset)
|
max_rows = min(max_rows, q_offset)
|
||||||
|
|
||||||
global_topk_offset = metadata.attn_metadata.topk_indices_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 (
|
assert (
|
||||||
seq_lens_expanded.shape[0] == q_offset
|
seq_lens_expanded.shape[0] == q_offset
|
||||||
@@ -733,10 +778,7 @@ class Indexer(MultiPlatformOp):
|
|||||||
else:
|
else:
|
||||||
# PAGED path: treat each token as a length-1 sequence
|
# PAGED path: treat each token as a length-1 sequence
|
||||||
topk_offset_chunk = None
|
topk_offset_chunk = None
|
||||||
B_chunk = logits_chunk.shape[0]
|
cu_seqlens_q_chunk = cu_seqlens_q_full[start:end]
|
||||||
cu_seqlens_q_chunk = torch.ones(
|
|
||||||
B_chunk, dtype=torch.int32, device=device
|
|
||||||
)
|
|
||||||
batch_idx_chunk = token_to_batch_idx[start:end]
|
batch_idx_chunk = token_to_batch_idx[start:end]
|
||||||
|
|
||||||
raw_topk_chunk = metadata.topk_transform(
|
raw_topk_chunk = metadata.topk_transform(
|
||||||
|
|||||||
Reference in New Issue
Block a user