Compare commits

...
Author SHA1 Message Date
Woosuk Kwon 7127e49d1c [Model Runner V2] Introduce num_tokens_for_attn
Signed-off-by: Woosuk Kwon <woosuk@inferact.ai>
2026-03-11 19:47:58 +00:00
5 changed files with 29 additions and 14 deletions
+18 -10
View File
@@ -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):
+2 -1
View File
@@ -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
+2
View File
@@ -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,
+6 -1
View File
@@ -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,
+1 -2
View File
@@ -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(