Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
7127e49d1c |
@@ -34,6 +34,7 @@ class BatchExecutionDescriptor:
|
||||
|
||||
cg_mode: CUDAGraphMode
|
||||
num_tokens: int
|
||||
num_tokens_for_attn: int | None
|
||||
num_reqs: int | None # None means no request padding is needed (PIECEWISE graphs)
|
||||
uniform_token_count: int | None = None
|
||||
|
||||
@@ -120,27 +121,31 @@ class CudaGraphManager:
|
||||
and decode_mode
|
||||
and self.decode_query_len <= num_tokens <= max_decode_tokens
|
||||
):
|
||||
num_reqs = num_tokens // self.decode_query_len
|
||||
desc = BatchExecutionDescriptor(
|
||||
cg_mode=decode_mode,
|
||||
num_tokens=num_tokens,
|
||||
num_reqs=num_tokens // self.decode_query_len,
|
||||
num_tokens_for_attn=num_reqs * self.decode_query_len,
|
||||
num_reqs=num_reqs,
|
||||
uniform_token_count=self.decode_query_len,
|
||||
)
|
||||
descs_by_mode[decode_mode].append(desc)
|
||||
descs_by_token_count[num_tokens].append(desc)
|
||||
|
||||
if mixed_mode:
|
||||
# for PIECEWISE graphs there is no limit on requests when replaying
|
||||
# i.e. no request padding is needed
|
||||
# so we leave it as None
|
||||
num_reqs = (
|
||||
min(num_tokens, self.max_num_reqs)
|
||||
if mixed_mode == CUDAGraphMode.FULL
|
||||
else None
|
||||
)
|
||||
if mixed_mode == CUDAGraphMode.FULL:
|
||||
num_reqs = min(num_tokens, self.max_num_reqs)
|
||||
num_tokens_for_attn = num_tokens
|
||||
else:
|
||||
# for PIECEWISE graphs there is no limit on requests when replaying
|
||||
# i.e. no request padding is needed
|
||||
# so we leave it as None
|
||||
num_reqs = None
|
||||
num_tokens_for_attn = None
|
||||
desc = BatchExecutionDescriptor(
|
||||
cg_mode=mixed_mode,
|
||||
num_tokens=num_tokens,
|
||||
num_tokens_for_attn=num_tokens_for_attn,
|
||||
num_reqs=num_reqs,
|
||||
)
|
||||
descs_by_mode[mixed_mode].append(desc)
|
||||
@@ -233,7 +238,10 @@ class CudaGraphManager:
|
||||
if _is_compatible(desc, num_reqs, num_tokens, uniform_token_count):
|
||||
return desc
|
||||
return BatchExecutionDescriptor(
|
||||
cg_mode=CUDAGraphMode.NONE, num_tokens=num_tokens, num_reqs=num_reqs
|
||||
cg_mode=CUDAGraphMode.NONE,
|
||||
num_tokens=num_tokens,
|
||||
num_tokens_for_attn=None,
|
||||
num_reqs=num_reqs,
|
||||
)
|
||||
|
||||
def run_fullgraph(self, desc: BatchExecutionDescriptor):
|
||||
|
||||
@@ -47,7 +47,7 @@ def sync_cudagraph_and_dp_padding(
|
||||
|
||||
if torch.all(num_tokens_across_dp == 0).item():
|
||||
synced_desc = BatchExecutionDescriptor(
|
||||
cg_mode=CUDAGraphMode.NONE, num_tokens=0, num_reqs=0
|
||||
cg_mode=CUDAGraphMode.NONE, num_tokens=0, num_tokens_for_attn=0, num_reqs=0
|
||||
)
|
||||
return synced_desc, None
|
||||
|
||||
@@ -58,6 +58,7 @@ def sync_cudagraph_and_dp_padding(
|
||||
return BatchExecutionDescriptor(
|
||||
cg_mode=CUDAGraphMode.NONE,
|
||||
num_tokens=num_tokens,
|
||||
num_tokens_for_attn=None,
|
||||
num_reqs=num_reqs,
|
||||
), num_tokens_across_dp
|
||||
|
||||
|
||||
@@ -53,6 +53,7 @@ class InputBatch:
|
||||
# sum(num_scheduled_tokens)
|
||||
num_tokens: int
|
||||
num_tokens_after_padding: int
|
||||
num_tokens_for_attn: int
|
||||
num_draft_tokens: int
|
||||
|
||||
# [num_reqs + 1]
|
||||
@@ -132,6 +133,7 @@ class InputBatch:
|
||||
num_scheduled_tokens=num_scheduled_tokens,
|
||||
num_tokens=num_tokens,
|
||||
num_tokens_after_padding=num_tokens,
|
||||
num_tokens_for_attn=num_tokens,
|
||||
num_draft_tokens=0,
|
||||
query_start_loc=query_start_loc,
|
||||
query_start_loc_np=query_start_loc_np,
|
||||
|
||||
@@ -599,8 +599,12 @@ class GPUModelRunner(LoRAModelRunnerMixin):
|
||||
self, scheduler_output: SchedulerOutput, batch_desc: BatchExecutionDescriptor
|
||||
) -> InputBatch:
|
||||
num_tokens = scheduler_output.total_num_scheduled_tokens
|
||||
num_tokens_after_padding = batch_desc.num_tokens
|
||||
assert num_tokens > 0
|
||||
num_tokens_after_padding = batch_desc.num_tokens
|
||||
if batch_desc.num_tokens_for_attn is not None:
|
||||
num_tokens_for_attn = batch_desc.num_tokens_for_attn
|
||||
else:
|
||||
num_tokens_for_attn = num_tokens
|
||||
num_tokens_per_req = scheduler_output.num_scheduled_tokens
|
||||
num_reqs = len(num_tokens_per_req)
|
||||
|
||||
@@ -721,6 +725,7 @@ class GPUModelRunner(LoRAModelRunnerMixin):
|
||||
num_scheduled_tokens=num_scheduled_tokens,
|
||||
num_tokens=num_tokens,
|
||||
num_tokens_after_padding=num_tokens_after_padding,
|
||||
num_tokens_for_attn=num_tokens_for_attn,
|
||||
num_draft_tokens=total_num_draft_tokens,
|
||||
query_start_loc=query_start_loc,
|
||||
query_start_loc_np=query_start_loc_np,
|
||||
|
||||
@@ -150,11 +150,10 @@ class DefaultModelState(ModelState):
|
||||
if cudagraph_mode == CUDAGraphMode.FULL:
|
||||
# Use padded sizes - padding is handled by model_runner.prepare_attn.
|
||||
num_reqs = input_batch.num_reqs_after_padding
|
||||
num_tokens = input_batch.num_tokens_after_padding
|
||||
else:
|
||||
# For piecewise cudagraphs and eager, use unpadded sizes.
|
||||
num_reqs = input_batch.num_reqs
|
||||
num_tokens = input_batch.num_tokens
|
||||
num_tokens = input_batch.num_tokens_for_attn
|
||||
query_start_loc_cpu = torch.from_numpy(input_batch.query_start_loc_np)
|
||||
max_query_len = input_batch.num_scheduled_tokens.max().item()
|
||||
attn_metadata = build_attn_metadata(
|
||||
|
||||
Reference in New Issue
Block a user