forked from Karylab-cklius/vllm
Fix Qwen3-VL and Qwen3-omni-thinker accuracy degradation from deepstack inputs under torch.compile (#43617)
Signed-off-by: Dakai An <dakaian108@gmail.com>
This commit is contained in:
@@ -1778,8 +1778,8 @@ class Qwen3OmniMoeThinkerForConditionalGeneration(
|
||||
) -> IntermediateTensors | None:
|
||||
if not getattr(self, "deepstack_input_embeds", None):
|
||||
return None # If vision tower is skipped
|
||||
if getattr(self, "deepstack_input_embeds_num_tokens", 0) == 0:
|
||||
return None
|
||||
if num_tokens > self.deepstack_input_embeds[0].size(0):
|
||||
self._resize_deepstack_input_embeds(num_tokens)
|
||||
|
||||
# get deepstack_input_embeds from buffer, and clear the buffer
|
||||
return IntermediateTensors(
|
||||
@@ -1791,6 +1791,17 @@ class Qwen3OmniMoeThinkerForConditionalGeneration(
|
||||
}
|
||||
)
|
||||
|
||||
def _resize_deepstack_input_embeds(self, num_tokens: int) -> None:
|
||||
self.deepstack_input_embeds = [
|
||||
torch.zeros(
|
||||
num_tokens,
|
||||
self.config.text_config.hidden_size,
|
||||
device=self.deepstack_input_embeds[0].device,
|
||||
dtype=self.deepstack_input_embeds[0].dtype,
|
||||
)
|
||||
for _ in range(self.deepstack_num_level)
|
||||
]
|
||||
|
||||
def _set_deepstack_input_embeds(self, deepstack_input_embeds: torch.Tensor) -> None:
|
||||
if not getattr(self, "deepstack_input_embeds", None):
|
||||
return
|
||||
@@ -1798,15 +1809,7 @@ class Qwen3OmniMoeThinkerForConditionalGeneration(
|
||||
# set deepstack_input_embeds to buffer
|
||||
num_tokens = deepstack_input_embeds.size(1)
|
||||
if num_tokens > self.deepstack_input_embeds[0].size(0):
|
||||
self.deepstack_input_embeds = [
|
||||
torch.zeros(
|
||||
num_tokens,
|
||||
self.config.text_config.hidden_size,
|
||||
device=self.deepstack_input_embeds[0].device,
|
||||
dtype=self.deepstack_input_embeds[0].dtype,
|
||||
)
|
||||
for _ in range(self.deepstack_num_level)
|
||||
]
|
||||
self._resize_deepstack_input_embeds(num_tokens)
|
||||
for idx in range(self.deepstack_num_level):
|
||||
self.deepstack_input_embeds[idx][:num_tokens].copy_(
|
||||
deepstack_input_embeds[idx]
|
||||
|
||||
@@ -1715,8 +1715,8 @@ class Qwen3VLForConditionalGeneration(
|
||||
) -> IntermediateTensors | None:
|
||||
if not getattr(self, "deepstack_input_embeds", None):
|
||||
return None # If vision tower is skipped
|
||||
if getattr(self, "deepstack_input_embeds_num_tokens", 0) == 0:
|
||||
return None
|
||||
if num_tokens > self.deepstack_input_embeds[0].size(0):
|
||||
self._resize_deepstack_input_embeds(num_tokens)
|
||||
|
||||
# get deepstack_input_embeds from buffer, and clear the buffer
|
||||
return IntermediateTensors(
|
||||
@@ -1728,6 +1728,17 @@ class Qwen3VLForConditionalGeneration(
|
||||
}
|
||||
)
|
||||
|
||||
def _resize_deepstack_input_embeds(self, num_tokens: int) -> None:
|
||||
self.deepstack_input_embeds = [
|
||||
torch.zeros(
|
||||
num_tokens,
|
||||
self.config.text_config.hidden_size,
|
||||
device=self.deepstack_input_embeds[0].device,
|
||||
dtype=self.deepstack_input_embeds[0].dtype,
|
||||
)
|
||||
for _ in range(self.deepstack_num_level)
|
||||
]
|
||||
|
||||
def _set_deepstack_input_embeds(self, deepstack_input_embeds: torch.Tensor) -> None:
|
||||
if not getattr(self, "deepstack_input_embeds", None):
|
||||
return
|
||||
@@ -1735,15 +1746,7 @@ class Qwen3VLForConditionalGeneration(
|
||||
# set deepstack_input_embeds to buffer
|
||||
num_tokens = deepstack_input_embeds.size(1)
|
||||
if num_tokens > self.deepstack_input_embeds[0].size(0):
|
||||
self.deepstack_input_embeds = [
|
||||
torch.zeros(
|
||||
num_tokens,
|
||||
self.config.text_config.hidden_size,
|
||||
device=self.deepstack_input_embeds[0].device,
|
||||
dtype=self.deepstack_input_embeds[0].dtype,
|
||||
)
|
||||
for _ in range(self.deepstack_num_level)
|
||||
]
|
||||
self._resize_deepstack_input_embeds(num_tokens)
|
||||
for idx in range(self.deepstack_num_level):
|
||||
self.deepstack_input_embeds[idx][:num_tokens].copy_(
|
||||
deepstack_input_embeds[idx]
|
||||
|
||||
Reference in New Issue
Block a user