Compare commits

...
Author SHA1 Message Date
yewentao256 9b134c72d3 fix v2 is_prefilling
Signed-off-by: yewentao256 <zhyanwentao@126.com>
2026-04-08 15:24:06 -04:00
4 changed files with 13 additions and 0 deletions
+4
View File
@@ -195,10 +195,13 @@ def build_attn_metadata(
kv_cache_config: KVCacheConfig,
dcp_local_seq_lens: torch.Tensor | None = None,
encoder_seq_lens: dict[int, tuple[torch.Tensor, np.ndarray]] | None = None,
is_prefilling: torch.Tensor | None = None,
) -> dict[str, Any]:
seq_lens = seq_lens[:num_reqs]
if dcp_local_seq_lens is not None:
dcp_local_seq_lens = dcp_local_seq_lens[:num_reqs]
if is_prefilling is not None:
is_prefilling = is_prefilling[:num_reqs]
attn_metadata: dict[str, Any] = {}
num_kv_cache_groups = len(kv_cache_config.kv_cache_groups)
@@ -218,6 +221,7 @@ def build_attn_metadata(
slot_mapping=slot_mapping,
causal=True,
dcp_local_seq_lens=dcp_local_seq_lens,
is_prefilling=is_prefilling,
)
if encoder_seq_lens and i in encoder_seq_lens:
encoder_seq_lens_gpu, encoder_seq_lens_cpu = encoder_seq_lens[i]
+2
View File
@@ -76,6 +76,7 @@ class InputBatch:
# Whether any requests in batch use structured output.
has_structured_output_reqs: bool
is_prefilling: torch.Tensor | None = None
@classmethod
def make_dummy(
@@ -143,6 +144,7 @@ class InputBatch:
cu_num_logits=cu_num_logits,
cu_num_logits_np=cu_num_logits_np,
has_structured_output_reqs=False,
is_prefilling=torch.zeros(num_reqs, dtype=torch.bool),
)
+6
View File
@@ -751,6 +751,11 @@ class GPUModelRunner(LoRAModelRunnerMixin):
self.input_buffers.seq_lens,
)
seq_lens = self.input_buffers.seq_lens[:num_reqs_padded]
is_prefilling = torch.zeros(num_reqs_padded, dtype=torch.bool)
is_prefilling[:num_reqs] = torch.as_tensor(
self.req_states.num_computed_prefill_tokens[idx_mapping_np]
< self.req_states.prefill_len.np[idx_mapping_np]
)
dcp_local_seq_lens = None
if self.use_dcp:
@@ -801,6 +806,7 @@ class GPUModelRunner(LoRAModelRunnerMixin):
cu_num_logits=cu_num_logits,
cu_num_logits_np=cu_num_logits_np,
has_structured_output_reqs=scheduler_output.has_structured_output_requests,
is_prefilling=is_prefilling,
)
def prepare_attn(
@@ -186,5 +186,6 @@ class DefaultModelState(ModelState):
slot_mappings=slot_mappings,
kv_cache_config=kv_cache_config,
dcp_local_seq_lens=input_batch.dcp_local_seq_lens,
is_prefilling=input_batch.is_prefilling,
)
return attn_metadata