[DSA] Make MQA logits free memory ratio configurable (#25859)
This commit is contained in:
@@ -474,6 +474,7 @@ class Envs:
|
|||||||
SGLANG_DSA_HIP_DISABLE_PRESHUFFLE = EnvBoolWithAlias(
|
SGLANG_DSA_HIP_DISABLE_PRESHUFFLE = EnvBoolWithAlias(
|
||||||
False, deprecated_name="SGLANG_NSA_HIP_DISABLE_PRESHUFFLE"
|
False, deprecated_name="SGLANG_NSA_HIP_DISABLE_PRESHUFFLE"
|
||||||
)
|
)
|
||||||
|
SGLANG_DSA_MQA_LOGITS_FREE_MEM_FRACTION = EnvFloat(0.2)
|
||||||
SGLANG_USE_FUSED_METADATA_COPY = EnvBool(True)
|
SGLANG_USE_FUSED_METADATA_COPY = EnvBool(True)
|
||||||
|
|
||||||
# sgl-kernel
|
# sgl-kernel
|
||||||
|
|||||||
@@ -178,10 +178,13 @@ def rotate_activation(x: torch.Tensor) -> torch.Tensor:
|
|||||||
class Indexer(MultiPlatformOp):
|
class Indexer(MultiPlatformOp):
|
||||||
_MQA_LOGITS_BYTES_PER_ELEM = 4
|
_MQA_LOGITS_BYTES_PER_ELEM = 4
|
||||||
_MQA_LOGITS_STATIC_SKIP_ELEMS = 8_000_000
|
_MQA_LOGITS_STATIC_SKIP_ELEMS = 8_000_000
|
||||||
_MQA_LOGITS_FREE_MEM_FRACTION = 0.5
|
|
||||||
_MQA_LOGITS_TOTAL_MEM_FRACTION = 0.3
|
_MQA_LOGITS_TOTAL_MEM_FRACTION = 0.3
|
||||||
_mqa_logits_budget_bytes: Dict[int, int] = {}
|
_mqa_logits_budget_bytes: Dict[int, int] = {}
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _mqa_logits_free_mem_fraction() -> float:
|
||||||
|
return envs.SGLANG_DSA_MQA_LOGITS_FREE_MEM_FRACTION.get()
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
hidden_size: int,
|
hidden_size: int,
|
||||||
@@ -559,6 +562,7 @@ class Indexer(MultiPlatformOp):
|
|||||||
return topk_result
|
return topk_result
|
||||||
|
|
||||||
def _get_mqa_logits_budget_bytes(self, device_index: int) -> int:
|
def _get_mqa_logits_budget_bytes(self, device_index: int) -> int:
|
||||||
|
free_mem_fraction = self._mqa_logits_free_mem_fraction()
|
||||||
cached_budget = self._mqa_logits_budget_bytes.get(device_index)
|
cached_budget = self._mqa_logits_budget_bytes.get(device_index)
|
||||||
if cached_budget is not None:
|
if cached_budget is not None:
|
||||||
return cached_budget
|
return cached_budget
|
||||||
@@ -572,7 +576,7 @@ class Indexer(MultiPlatformOp):
|
|||||||
else:
|
else:
|
||||||
static_free_mem = int(total_mem * max(0.0, 1.0 - mem_fraction_static))
|
static_free_mem = int(total_mem * max(0.0, 1.0 - mem_fraction_static))
|
||||||
static_budget = min(
|
static_budget = min(
|
||||||
int(static_free_mem * self._MQA_LOGITS_FREE_MEM_FRACTION),
|
int(static_free_mem * free_mem_fraction),
|
||||||
total_mem_budget,
|
total_mem_budget,
|
||||||
)
|
)
|
||||||
static_budget = max(1, static_budget)
|
static_budget = max(1, static_budget)
|
||||||
@@ -587,9 +591,7 @@ class Indexer(MultiPlatformOp):
|
|||||||
# torch.cuda.mem_get_info synchronizes the host, so cache the result,
|
# torch.cuda.mem_get_info synchronizes the host, so cache the result,
|
||||||
# capped by the workload-independent serving-memory headroom.
|
# capped by the workload-independent serving-memory headroom.
|
||||||
free_mem, _ = torch.cuda.mem_get_info(device_index)
|
free_mem, _ = torch.cuda.mem_get_info(device_index)
|
||||||
budget_bytes = min(
|
budget_bytes = min(int(free_mem * free_mem_fraction), static_budget)
|
||||||
int(free_mem * self._MQA_LOGITS_FREE_MEM_FRACTION), static_budget
|
|
||||||
)
|
|
||||||
|
|
||||||
budget_bytes = max(1, budget_bytes)
|
budget_bytes = max(1, budget_bytes)
|
||||||
self._mqa_logits_budget_bytes[device_index] = budget_bytes
|
self._mqa_logits_budget_bytes[device_index] = budget_bytes
|
||||||
|
|||||||
Reference in New Issue
Block a user