[Fix]Fix capture fail bug for DeepSeek (#6275)

This commit is contained in:
Baizhou Zhang
2025-05-21 11:11:20 -07:00
committed by GitHub
parent 55f6005f53
commit d4c038daed
4 changed files with 20 additions and 13 deletions
+3 -2
View File
@@ -836,7 +836,6 @@ class MllamaForConditionalGeneration(nn.Module):
prefix="multi_modal_projector",
)
self.logits_processor = LogitsProcessor(config.text_config)
self.capture_mode = False
def pad_input_ids(self, input_ids: List[int], mm_inputs: MultimodalInputs):
pixel_values = torch.cat(
@@ -969,6 +968,8 @@ class MllamaForConditionalGeneration(nn.Module):
positions: torch.Tensor,
forward_batch: ForwardBatch,
) -> Union[Tuple, CausalLMOutputWithPast]:
from sglang.srt.model_executor.cuda_graph_runner import get_is_capture_mode
batched_images, batched_ar_ids, batched_ar_mask, encoder_lens_need = (
self._batch_image_inputs(forward_batch)
)
@@ -977,7 +978,7 @@ class MllamaForConditionalGeneration(nn.Module):
cross_attention_mask = None
cross_attention_states = None
if self.capture_mode:
if get_is_capture_mode():
# NOTE: when doing cuda graph capture, we do not want to skip cross attention
# Make is a constant value to avoid cuda graph capture issue
skip_cross_attention = False