[Misc] Fix metrics, weight update lock, request logging (#2543)
This commit is contained in:
+33
-44
@@ -14,6 +14,7 @@
|
||||
"""Common utilities."""
|
||||
|
||||
import base64
|
||||
import dataclasses
|
||||
import ipaddress
|
||||
import itertools
|
||||
import json
|
||||
@@ -1238,49 +1239,37 @@ def cuda_device_count_stateless() -> int:
|
||||
return _cuda_device_count_stateless(os.environ.get("CUDA_VISIBLE_DEVICES", None))
|
||||
|
||||
|
||||
def should_use_tensor_core(
|
||||
kv_cache_dtype: torch.dtype,
|
||||
num_attention_heads: int,
|
||||
num_kv_heads: int,
|
||||
) -> bool:
|
||||
"""
|
||||
Determine whether to use tensor cores for attention computation.
|
||||
|
||||
Args:
|
||||
kv_cache_dtype: Data type of the KV cache
|
||||
num_attention_heads: Number of attention heads
|
||||
num_kv_heads: Number of key/value heads
|
||||
|
||||
Returns:
|
||||
bool: Whether to use tensor cores
|
||||
"""
|
||||
# Try to use environment variable first
|
||||
env_override = os.environ.get("SGLANG_FLASHINFER_USE_TENSOR_CORE")
|
||||
if env_override is not None:
|
||||
return env_override.lower() == "true"
|
||||
|
||||
# Try to use _grouped_size_compiled_for_decode_kernels if available
|
||||
# This is for flashinfer <=0.1.6. Otherwise, there is an accuracy bug
|
||||
try:
|
||||
from flashinfer.decode import _grouped_size_compiled_for_decode_kernels
|
||||
|
||||
if not _grouped_size_compiled_for_decode_kernels(
|
||||
num_attention_heads,
|
||||
num_kv_heads,
|
||||
):
|
||||
return True
|
||||
def dataclass_to_string_truncated(data, max_length=2048):
|
||||
if isinstance(data, str):
|
||||
if len(data) > max_length:
|
||||
half_length = max_length // 2
|
||||
return f'"{data[:half_length]} ... {data[-half_length:]}"'
|
||||
else:
|
||||
return False
|
||||
except (ImportError, AttributeError):
|
||||
pass
|
||||
|
||||
# Calculate GQA group size
|
||||
gqa_group_size = num_attention_heads // num_kv_heads
|
||||
|
||||
# Determine based on dtype and GQA group size
|
||||
if kv_cache_dtype in (torch.float8_e4m3fn, torch.float8_e5m2):
|
||||
return True
|
||||
elif kv_cache_dtype in (torch.float16, torch.half, torch.bfloat16):
|
||||
return gqa_group_size > 4
|
||||
return f'"{data}"'
|
||||
elif isinstance(data, (list, tuple)):
|
||||
if len(data) > max_length:
|
||||
half_length = max_length // 2
|
||||
return str(data[:half_length]) + " ... " + str(data[-half_length:])
|
||||
else:
|
||||
return str(data)
|
||||
elif isinstance(data, dict):
|
||||
return (
|
||||
"{"
|
||||
+ ", ".join(
|
||||
f"{k}: {dataclass_to_string_truncated(v, max_length)}"
|
||||
for k, v in data.items()
|
||||
)
|
||||
+ "}"
|
||||
)
|
||||
elif dataclasses.is_dataclass(data):
|
||||
fields = dataclasses.fields(data)
|
||||
return (
|
||||
f"{data.__class__.__name__}("
|
||||
+ ", ".join(
|
||||
f"{f.name}={dataclass_to_string_truncated(getattr(data, f.name), max_length)}"
|
||||
for f in fields
|
||||
)
|
||||
+ ")"
|
||||
)
|
||||
else:
|
||||
return False
|
||||
return str(data)
|
||||
|
||||
Reference in New Issue
Block a user