[SGL] sync patch: Remove sync points, prefill cudagraph for DP, disable cache reset in mem check (#19190)
Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com> Co-authored-by: ispobock <ispobaoke@gmail.com>
This commit is contained in:
co-authored by
Claude Opus 4.6
ispobock
parent
8c0f2d40bd
commit
b5a8e4179e
@@ -685,6 +685,7 @@ class TboForwardBatchPreparer:
|
|||||||
for key in [
|
for key in [
|
||||||
"forward_mode",
|
"forward_mode",
|
||||||
"is_extend_in_batch",
|
"is_extend_in_batch",
|
||||||
|
"all_extend_in_batch",
|
||||||
"return_logprob",
|
"return_logprob",
|
||||||
"req_to_token_pool",
|
"req_to_token_pool",
|
||||||
"token_to_kv_pool",
|
"token_to_kv_pool",
|
||||||
|
|||||||
@@ -22,7 +22,7 @@ class ConnectorType(str, enum.Enum):
|
|||||||
INSTANCE = "instance"
|
INSTANCE = "instance"
|
||||||
|
|
||||||
|
|
||||||
def create_remote_connector(url, device, **kwargs) -> BaseConnector:
|
def create_remote_connector(url, device=None, **kwargs) -> BaseConnector:
|
||||||
connector_type = parse_connector_type(url)
|
connector_type = parse_connector_type(url)
|
||||||
if connector_type == "redis":
|
if connector_type == "redis":
|
||||||
return RedisConnector(url)
|
return RedisConnector(url)
|
||||||
|
|||||||
@@ -519,11 +519,11 @@ class LogitsProcessor(nn.Module):
|
|||||||
if hidden_states_before_norm is not None:
|
if hidden_states_before_norm is not None:
|
||||||
pruned_states_before_norm = torch.cat(pruned_states_before_norm_list)
|
pruned_states_before_norm = torch.cat(pruned_states_before_norm_list)
|
||||||
sample_indices = torch.tensor(
|
sample_indices = torch.tensor(
|
||||||
sample_indices, device=pruned_states.device, dtype=torch.int64
|
sample_indices, dtype=torch.int64, pin_memory=True
|
||||||
)
|
).to(pruned_states.device, non_blocking=True)
|
||||||
input_logprob_indices = torch.tensor(
|
input_logprob_indices = torch.tensor(
|
||||||
input_logprob_indices, device=pruned_states.device, dtype=torch.int64
|
input_logprob_indices, dtype=torch.int64, pin_memory=True
|
||||||
)
|
).to(pruned_states.device, non_blocking=True)
|
||||||
|
|
||||||
return (
|
return (
|
||||||
pruned_states,
|
pruned_states,
|
||||||
@@ -590,19 +590,24 @@ class LogitsProcessor(nn.Module):
|
|||||||
def _expand_metadata_for_logprobs(
|
def _expand_metadata_for_logprobs(
|
||||||
self, logits_metadata: LogitsMetadata, device: torch.device
|
self, logits_metadata: LogitsMetadata, device: torch.device
|
||||||
):
|
):
|
||||||
|
# Avoid implicit device sync inside repeat_interleave by providing output_size,
|
||||||
|
# which we can compute from CPU metadata.
|
||||||
|
total_pruned_len = sum(logits_metadata.extend_logprob_pruned_lens_cpu)
|
||||||
pruned_lens = torch.tensor(
|
pruned_lens = torch.tensor(
|
||||||
logits_metadata.extend_logprob_pruned_lens_cpu,
|
logits_metadata.extend_logprob_pruned_lens_cpu,
|
||||||
device=device,
|
pin_memory=True,
|
||||||
)
|
).to(device, non_blocking=True)
|
||||||
if logits_metadata.temp_scaled_logprobs:
|
if logits_metadata.temp_scaled_logprobs:
|
||||||
logits_metadata.temperature = torch.repeat_interleave(
|
logits_metadata.temperature = torch.repeat_interleave(
|
||||||
logits_metadata.temperature.view(-1),
|
logits_metadata.temperature.view(-1),
|
||||||
pruned_lens,
|
pruned_lens,
|
||||||
|
output_size=total_pruned_len,
|
||||||
).view(-1, 1)
|
).view(-1, 1)
|
||||||
if logits_metadata.top_p_normalized_logprobs:
|
if logits_metadata.top_p_normalized_logprobs:
|
||||||
logits_metadata.top_p = torch.repeat_interleave(
|
logits_metadata.top_p = torch.repeat_interleave(
|
||||||
logits_metadata.top_p,
|
logits_metadata.top_p,
|
||||||
pruned_lens,
|
pruned_lens,
|
||||||
|
output_size=total_pruned_len,
|
||||||
)
|
)
|
||||||
|
|
||||||
def process_input_logprobs(self, input_logits, logits_metadata: LogitsMetadata):
|
def process_input_logprobs(self, input_logits, logits_metadata: LogitsMetadata):
|
||||||
|
|||||||
@@ -1226,6 +1226,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
global_num_tokens: Optional[List[int]] = None
|
global_num_tokens: Optional[List[int]] = None
|
||||||
global_num_tokens_for_logprob: Optional[List[int]] = None
|
global_num_tokens_for_logprob: Optional[List[int]] = None
|
||||||
is_extend_in_batch: bool = False
|
is_extend_in_batch: bool = False
|
||||||
|
all_extend_in_batch: bool = False
|
||||||
can_run_dp_cuda_graph: bool = False
|
can_run_dp_cuda_graph: bool = False
|
||||||
tbo_split_seq_index: Optional[int] = None
|
tbo_split_seq_index: Optional[int] = None
|
||||||
global_forward_mode: Optional[ForwardMode] = None
|
global_forward_mode: Optional[ForwardMode] = None
|
||||||
@@ -1985,22 +1986,34 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
self.seq_lens_sum += bs
|
self.seq_lens_sum += bs
|
||||||
|
|
||||||
if get_global_server_args().enable_mamba_extra_buffer():
|
if get_global_server_args().enable_mamba_extra_buffer():
|
||||||
self.mamba_track_indices = torch.tensor(
|
# Build indices fully on GPU without scalar extraction.
|
||||||
[
|
# Each slice is shape [1]; cat -> [bs].
|
||||||
req.mamba_ping_pong_track_buffer[req.mamba_next_track_idx]
|
if len(self.reqs) == 0:
|
||||||
for req in self.reqs
|
self.mamba_track_indices = torch.empty(
|
||||||
],
|
(0,), dtype=torch.int64, device=self.device
|
||||||
dtype=torch.int64,
|
)
|
||||||
device=self.device,
|
else:
|
||||||
)
|
self.mamba_track_indices = torch.cat(
|
||||||
|
[
|
||||||
|
(
|
||||||
|
req.mamba_ping_pong_track_buffer[1:]
|
||||||
|
if req.mamba_next_track_idx == 1
|
||||||
|
else req.mamba_ping_pong_track_buffer[:1]
|
||||||
|
)
|
||||||
|
for req in self.reqs
|
||||||
|
],
|
||||||
|
dim=0,
|
||||||
|
).to(torch.int64)
|
||||||
|
|
||||||
|
# Keep mask construction in the pinned-tensor form.
|
||||||
self.mamba_track_mask = torch.tensor(
|
self.mamba_track_mask = torch.tensor(
|
||||||
[
|
[
|
||||||
sl % get_global_server_args().mamba_track_interval == 0
|
sl % get_global_server_args().mamba_track_interval == 0
|
||||||
for sl in self.seq_lens_cpu
|
for sl in self.seq_lens_cpu
|
||||||
],
|
],
|
||||||
dtype=torch.bool,
|
dtype=torch.bool,
|
||||||
device=self.device,
|
pin_memory=True,
|
||||||
)
|
).to(device=self.device, non_blocking=True)
|
||||||
|
|
||||||
def maybe_wait_verify_done(self):
|
def maybe_wait_verify_done(self):
|
||||||
if self.is_spec_v2:
|
if self.is_spec_v2:
|
||||||
@@ -2170,6 +2183,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
global_num_tokens=self.global_num_tokens,
|
global_num_tokens=self.global_num_tokens,
|
||||||
global_num_tokens_for_logprob=self.global_num_tokens_for_logprob,
|
global_num_tokens_for_logprob=self.global_num_tokens_for_logprob,
|
||||||
is_extend_in_batch=self.is_extend_in_batch,
|
is_extend_in_batch=self.is_extend_in_batch,
|
||||||
|
all_extend_in_batch=self.all_extend_in_batch,
|
||||||
can_run_dp_cuda_graph=self.can_run_dp_cuda_graph,
|
can_run_dp_cuda_graph=self.can_run_dp_cuda_graph,
|
||||||
tbo_split_seq_index=self.tbo_split_seq_index,
|
tbo_split_seq_index=self.tbo_split_seq_index,
|
||||||
global_forward_mode=self.global_forward_mode,
|
global_forward_mode=self.global_forward_mode,
|
||||||
@@ -2227,6 +2241,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
global_num_tokens=self.global_num_tokens,
|
global_num_tokens=self.global_num_tokens,
|
||||||
global_num_tokens_for_logprob=self.global_num_tokens_for_logprob,
|
global_num_tokens_for_logprob=self.global_num_tokens_for_logprob,
|
||||||
can_run_dp_cuda_graph=self.can_run_dp_cuda_graph,
|
can_run_dp_cuda_graph=self.can_run_dp_cuda_graph,
|
||||||
|
all_extend_in_batch=self.all_extend_in_batch,
|
||||||
is_extend_in_batch=self.is_extend_in_batch,
|
is_extend_in_batch=self.is_extend_in_batch,
|
||||||
is_prefill_only=self.is_prefill_only,
|
is_prefill_only=self.is_prefill_only,
|
||||||
seq_lens_cpu=self.seq_lens_cpu,
|
seq_lens_cpu=self.seq_lens_cpu,
|
||||||
@@ -2331,6 +2346,7 @@ class ModelWorkerBatch:
|
|||||||
global_num_tokens: Optional[List[int]]
|
global_num_tokens: Optional[List[int]]
|
||||||
global_num_tokens_for_logprob: Optional[List[int]]
|
global_num_tokens_for_logprob: Optional[List[int]]
|
||||||
is_extend_in_batch: bool
|
is_extend_in_batch: bool
|
||||||
|
all_extend_in_batch: bool
|
||||||
can_run_dp_cuda_graph: bool
|
can_run_dp_cuda_graph: bool
|
||||||
tbo_split_seq_index: Optional[int]
|
tbo_split_seq_index: Optional[int]
|
||||||
global_forward_mode: Optional[ForwardMode]
|
global_forward_mode: Optional[ForwardMode]
|
||||||
|
|||||||
@@ -343,10 +343,18 @@ class MambaPool:
|
|||||||
|
|
||||||
select_index = self.free_slots[:need_size]
|
select_index = self.free_slots[:need_size]
|
||||||
self.free_slots = self.free_slots[need_size:]
|
self.free_slots = self.free_slots[need_size:]
|
||||||
# clear at alloc time, fill allocated slots with zeros
|
# clear at alloc time — expand a scalar GPU zero to the right shape, no CPU-GPU sync
|
||||||
for i in range(len(self.mamba_cache.conv)):
|
for i in range(len(self.mamba_cache.conv)):
|
||||||
self.mamba_cache.conv[i][:, select_index] = 0
|
t = self.mamba_cache.conv[i]
|
||||||
self.mamba_cache.temporal[:, select_index] = 0
|
z = torch.zeros(1, dtype=t.dtype, device=t.device).expand(
|
||||||
|
t.shape[0], need_size, *t.shape[2:]
|
||||||
|
)
|
||||||
|
t[:, select_index] = z
|
||||||
|
t = self.mamba_cache.temporal
|
||||||
|
z = torch.zeros(1, dtype=t.dtype, device=t.device).expand(
|
||||||
|
t.shape[0], need_size, *t.shape[2:]
|
||||||
|
)
|
||||||
|
t[:, select_index] = z
|
||||||
|
|
||||||
return select_index
|
return select_index
|
||||||
|
|
||||||
@@ -514,8 +522,8 @@ class HybridReqToTokenPool(ReqToTokenPool):
|
|||||||
if select_index is None:
|
if select_index is None:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
mamba_index = []
|
mamba_indices: list[torch.Tensor] = []
|
||||||
mamba_ping_pong_track_buffer_list = []
|
mamba_ping_pong_track_buffers: list[torch.Tensor] = []
|
||||||
for req in reqs:
|
for req in reqs:
|
||||||
mid = None
|
mid = None
|
||||||
if req.mamba_pool_idx is not None: # for radix cache
|
if req.mamba_pool_idx is not None: # for radix cache
|
||||||
@@ -527,7 +535,7 @@ class HybridReqToTokenPool(ReqToTokenPool):
|
|||||||
), f"Not enough space for mamba cache, try to increase --mamba-full-memory-ratio or --max-mamba-cache-size. {mid=}, {self.mamba_pool.size=}, {self.mamba_pool.available_size()=}, {len(reqs)=}"
|
), f"Not enough space for mamba cache, try to increase --mamba-full-memory-ratio or --max-mamba-cache-size. {mid=}, {self.mamba_pool.size=}, {self.mamba_pool.available_size()=}, {len(reqs)=}"
|
||||||
mid = mid[0]
|
mid = mid[0]
|
||||||
req.mamba_pool_idx = mid
|
req.mamba_pool_idx = mid
|
||||||
mamba_index.append(mid)
|
mamba_indices.append(mid)
|
||||||
if self.enable_mamba_extra_buffer:
|
if self.enable_mamba_extra_buffer:
|
||||||
if req.mamba_ping_pong_track_buffer is None:
|
if req.mamba_ping_pong_track_buffer is None:
|
||||||
req.mamba_ping_pong_track_buffer = self.mamba_pool.alloc(
|
req.mamba_ping_pong_track_buffer = self.mamba_pool.alloc(
|
||||||
@@ -537,26 +545,22 @@ class HybridReqToTokenPool(ReqToTokenPool):
|
|||||||
req.mamba_ping_pong_track_buffer is not None
|
req.mamba_ping_pong_track_buffer is not None
|
||||||
), "Not enough space for mamba ping pong idx, try to increase --mamba-full-memory-ratio."
|
), "Not enough space for mamba ping pong idx, try to increase --mamba-full-memory-ratio."
|
||||||
req.mamba_next_track_idx = 0
|
req.mamba_next_track_idx = 0
|
||||||
mamba_ping_pong_track_buffer_list.append(
|
mamba_ping_pong_track_buffers.append(req.mamba_ping_pong_track_buffer)
|
||||||
req.mamba_ping_pong_track_buffer.tolist()
|
|
||||||
)
|
|
||||||
assert len(select_index) == len(
|
assert len(select_index) == len(
|
||||||
mamba_index
|
mamba_indices
|
||||||
), f"Not enough space for mamba cache, try to increase --mamba-full-memory-ratio or --max-mamba-cache-size."
|
), f"Not enough space for mamba cache, try to increase --mamba-full-memory-ratio or --max-mamba-cache-size."
|
||||||
if self.enable_mamba_extra_buffer:
|
if self.enable_mamba_extra_buffer:
|
||||||
assert len(select_index) == len(
|
assert len(select_index) == len(
|
||||||
mamba_ping_pong_track_buffer_list
|
mamba_ping_pong_track_buffers
|
||||||
), f"Not enough space for mamba ping pong idx, try to increase --mamba-full-memory-ratio."
|
), f"Not enough space for mamba ping pong idx, try to increase --mamba-full-memory-ratio."
|
||||||
self.req_index_to_mamba_index_mapping[select_index] = torch.tensor(
|
mamba_index_tensor = torch.stack(mamba_indices).to(dtype=torch.int32)
|
||||||
mamba_index, dtype=torch.int32, device=self.device
|
self.req_index_to_mamba_index_mapping[select_index] = mamba_index_tensor
|
||||||
)
|
|
||||||
if self.enable_mamba_extra_buffer:
|
if self.enable_mamba_extra_buffer:
|
||||||
|
ping_pong_tensor = torch.stack(mamba_ping_pong_track_buffers).to(
|
||||||
|
dtype=torch.int32
|
||||||
|
)
|
||||||
self.req_index_to_mamba_ping_pong_track_buffer_mapping[select_index] = (
|
self.req_index_to_mamba_ping_pong_track_buffer_mapping[select_index] = (
|
||||||
torch.tensor(
|
ping_pong_tensor
|
||||||
mamba_ping_pong_track_buffer_list,
|
|
||||||
dtype=torch.int32,
|
|
||||||
device=self.device,
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
return select_index
|
return select_index
|
||||||
|
|
||||||
@@ -593,11 +597,28 @@ class HybridReqToTokenPool(ReqToTokenPool):
|
|||||||
0,
|
0,
|
||||||
1,
|
1,
|
||||||
], f"mamba_ping_pong_track_buffer_to_keep must be 0 or 1, {mamba_ping_pong_track_buffer_to_keep=}"
|
], f"mamba_ping_pong_track_buffer_to_keep must be 0 or 1, {mamba_ping_pong_track_buffer_to_keep=}"
|
||||||
idx_to_free = list(range(self.mamba_ping_pong_track_buffer_size))
|
# Avoid Python-list advanced indexing on a device tensor.
|
||||||
idx_to_free.remove(mamba_ping_pong_track_buffer_to_keep)
|
# The ping-pong buffer size is either 2 (normal) or 1 (spec decode).
|
||||||
mamba_ping_pong_track_buffer_to_free = (
|
if self.mamba_ping_pong_track_buffer_size == 2:
|
||||||
mamba_ping_pong_track_buffer_to_free[idx_to_free]
|
idx_to_free = 1 - mamba_ping_pong_track_buffer_to_keep
|
||||||
)
|
mamba_ping_pong_track_buffer_to_free = (
|
||||||
|
mamba_ping_pong_track_buffer_to_free[
|
||||||
|
idx_to_free : idx_to_free + 1
|
||||||
|
]
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
assert self.mamba_ping_pong_track_buffer_size == 1, (
|
||||||
|
f"Unexpected mamba_ping_pong_track_buffer_size="
|
||||||
|
f"{self.mamba_ping_pong_track_buffer_size}"
|
||||||
|
)
|
||||||
|
assert mamba_ping_pong_track_buffer_to_keep == 0, (
|
||||||
|
"mamba_ping_pong_track_buffer_to_keep must be 0 when "
|
||||||
|
"mamba_ping_pong_track_buffer_size is 1"
|
||||||
|
)
|
||||||
|
# Keep the only slot, so free nothing.
|
||||||
|
mamba_ping_pong_track_buffer_to_free = (
|
||||||
|
mamba_ping_pong_track_buffer_to_free[0:0]
|
||||||
|
)
|
||||||
self.mamba_pool.free(mamba_ping_pong_track_buffer_to_free)
|
self.mamba_pool.free(mamba_ping_pong_track_buffer_to_free)
|
||||||
|
|
||||||
def clear(self):
|
def clear(self):
|
||||||
|
|||||||
@@ -338,6 +338,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
dp_local_num_tokens: Optional[torch.Tensor] = None # cached info at runtime
|
dp_local_num_tokens: Optional[torch.Tensor] = None # cached info at runtime
|
||||||
global_dp_buffer_len: Optional[int] = None
|
global_dp_buffer_len: Optional[int] = None
|
||||||
is_extend_in_batch: bool = False
|
is_extend_in_batch: bool = False
|
||||||
|
all_extend_in_batch: bool = False
|
||||||
can_run_dp_cuda_graph: bool = False
|
can_run_dp_cuda_graph: bool = False
|
||||||
global_forward_mode: Optional[ForwardMode] = None
|
global_forward_mode: Optional[ForwardMode] = None
|
||||||
|
|
||||||
@@ -404,6 +405,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
top_logprobs_nums=batch.top_logprobs_nums,
|
top_logprobs_nums=batch.top_logprobs_nums,
|
||||||
token_ids_logprobs=batch.token_ids_logprobs,
|
token_ids_logprobs=batch.token_ids_logprobs,
|
||||||
is_extend_in_batch=batch.is_extend_in_batch,
|
is_extend_in_batch=batch.is_extend_in_batch,
|
||||||
|
all_extend_in_batch=batch.all_extend_in_batch,
|
||||||
can_run_dp_cuda_graph=batch.can_run_dp_cuda_graph,
|
can_run_dp_cuda_graph=batch.can_run_dp_cuda_graph,
|
||||||
global_forward_mode=batch.global_forward_mode,
|
global_forward_mode=batch.global_forward_mode,
|
||||||
is_prefill_only=batch.is_prefill_only,
|
is_prefill_only=batch.is_prefill_only,
|
||||||
|
|||||||
@@ -1133,7 +1133,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
"""Update engine weights in-place from the disk."""
|
"""Update engine weights in-place from the disk."""
|
||||||
logger.info(
|
logger.info(
|
||||||
f"Update engine weights online from disk begin. "
|
f"Update engine weights online from disk begin. "
|
||||||
f"avail mem={get_available_gpu_memory(self.device, self.gpu_id):.2f} GB"
|
f"avail mem={get_available_gpu_memory(self.device, self.gpu_id, empty_cache=False):.2f} GB"
|
||||||
)
|
)
|
||||||
|
|
||||||
target_device = torch.device(self.device)
|
target_device = torch.device(self.device)
|
||||||
|
|||||||
Reference in New Issue
Block a user