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:
Dakai An
2026-05-27 15:34:08 -07:00
committed by GitHub
parent 7fb9c0197a
commit 5963c19478
2 changed files with 28 additions and 22 deletions
@@ -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]
+14 -11
View File
@@ -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]