Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 11 additions & 5 deletions tests/datasets/test_qwen35_vl_tokenize_fn.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@
import json
import torch
import parametrize
from xtuner.v1.utils.test_utils import add_video_root
from xtuner.v1.utils.test_utils import add_video_root, get_qwen3_vl_video_chat_template
from packaging.version import Version
from transformers import __version__ as transformers_version
import unittest
Expand Down Expand Up @@ -287,10 +287,16 @@ def test_qwen3_vl_sft_video(self, add_vision_id):
add_video_root(messages, VIDEO_ROOT)

if i not in [8, 9]:
ret = self.processor.apply_chat_template(messages, add_generation_prompt=False, tokenize=True,
do_sample_frames=do_sample_frames,
return_dict=True, add_vision_id=add_vision_id,
return_tensors="pt")
ret = self.processor.apply_chat_template(
messages,
chat_template=get_qwen3_vl_video_chat_template(self.processor),
add_generation_prompt=False,
tokenize=True,
return_dict=True,
add_vision_id=add_vision_id,
return_tensors="pt",
processor_kwargs={"do_sample_frames": do_sample_frames},
)
input_ids_hf = ret['input_ids'][0]
pixel_values_hf = ret['pixel_values_videos']
image_grid_thw_hf = ret['video_grid_thw']
Expand Down
16 changes: 11 additions & 5 deletions tests/datasets/test_qwen3_vl_tokenize_fn.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@
import json
import torch
import parametrize
from xtuner.v1.utils.test_utils import add_video_root
from xtuner.v1.utils.test_utils import add_video_root, get_qwen3_vl_video_chat_template

QWEN3_VL_PATH = os.environ["QWEN3_VL_MOE_PATH"]
VIDEO_ROOT = os.environ["VIDEO_ROOT"]
Expand Down Expand Up @@ -319,10 +319,16 @@ def test_qwen3_vl_sft_video(self, add_vision_id):
add_video_root(messages, VIDEO_ROOT)

if i not in [8, 9]:
ret = self.processor.apply_chat_template(messages, add_generation_prompt=False, tokenize=True,
do_sample_frames=do_sample_frames,
return_dict=True, add_vision_id=add_vision_id,
return_tensors="pt")
ret = self.processor.apply_chat_template(
messages,
chat_template=get_qwen3_vl_video_chat_template(self.processor),
add_generation_prompt=False,
tokenize=True,
return_dict=True,
add_vision_id=add_vision_id,
return_tensors="pt",
processor_kwargs={"do_sample_frames": do_sample_frames},
)
input_ids_hf = ret['input_ids'][0]
pixel_values_hf = ret['pixel_values_videos']
image_grid_thw_hf = ret['video_grid_thw']
Expand Down
76 changes: 54 additions & 22 deletions tests/model/test_qwen3_5.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,7 @@
import os
import unittest
from unittest.mock import patch

import parametrize
import torch
from packaging.version import Version
Expand Down Expand Up @@ -31,16 +33,39 @@
f"transformers >= 5.2.0 is required, but got {transformers_version}"
)
class TestQwen3_5_VL(DeterministicDDPTestCase):
@staticmethod
def _patch_hf_router_keep_fp32_topk_weights():
from transformers.models.qwen3_5_moe.modeling_qwen3_5_moe import Qwen3_5MoeTopKRouter

def _forward(router, hidden_states):
hidden_states = hidden_states.reshape(-1, router.hidden_dim)
router_logits = torch.nn.functional.linear(hidden_states, router.weight)
router_probs = torch.nn.functional.softmax(router_logits, dtype=torch.float, dim=-1)
router_top_value, router_indices = torch.topk(router_probs, router.top_k, dim=-1)
router_top_value /= router_top_value.sum(dim=-1, keepdim=True)

# Transformers 5.2 overwrote router_logits with the fp32 softmax result,
# so this dtype conversion was effectively a no-op. Keep the 5.14 logits
# semantics while restoring the fp32 top-k weights used by XTuner's
# historical HF-parity baseline.
return router_logits, router_top_value, router_indices

return patch.object(Qwen3_5MoeTopKRouter, "forward", _forward)

def _patch_xtuner_fast_pos_embed_interpolate(self) -> None:
from xtuner.v1.model.compose.qwen3_vl.modeling_vision import Qwen3VLVisionModel
from transformers.models.qwen3_5_moe import Qwen3_5MoeVisionModel
from transformers.vision_utils import get_vision_bilinear_indices_and_weights

hf_fast_pos_embed_interpolate = Qwen3_5MoeVisionModel.fast_pos_embed_interpolate
from xtuner.v1.model.compose.qwen3_vl.modeling_vision import Qwen3VLVisionModel

def _fast_pos_embed_interpolate(self, grid_thw):
# Transformers 5.14 accumulates interpolation in fp32 and casts in its
# forward; XTuner's older forward expects this helper to keep the model dtype.
return hf_fast_pos_embed_interpolate(self, grid_thw).to(self.pos_embed.weight.dtype)
indices, weights = get_vision_bilinear_indices_and_weights(
grid_thw,
num_grid_per_side=self.num_grid_per_side,
spatial_merge_size=self.config.spatial_merge_size,
)
return (self.pos_embed(indices) * weights[:, :, None]).sum(0).to(self.pos_embed.weight.dtype)

Qwen3VLVisionModel.fast_pos_embed_interpolate = _fast_pos_embed_interpolate

Expand Down Expand Up @@ -191,19 +216,20 @@ def test_qwen3_5_vl_run(self, device, sp_size, tol):
from transformers import Qwen3_5MoeForConditionalGeneration
QWEN3_VL_MOE_PATH = os.environ["QWEN3_5_MOE_PATH"]

hf_model = Qwen3_5MoeForConditionalGeneration.from_pretrained(
QWEN3_VL_MOE_PATH,
dtype=torch.bfloat16,
attn_implementation="flash_attention_2",
device_map="cuda",
trust_remote_code=True
).eval()
# Cannot understand, but must accept. Once there is no this code, it will appear cuda access illegal memory error in multi-GPU
torch.distributed.barrier()

loss_hf_text = self._forward(hf_model, type='text', device=device, sp_size=sp_size)
loss_hf_image = self._forward(hf_model, type='image', device=device, sp_size=sp_size)
# loss_hf_video = self._forward(hf_model, type='video', device=device, sp_size=sp_size)
with self._patch_hf_router_keep_fp32_topk_weights():
hf_model = Qwen3_5MoeForConditionalGeneration.from_pretrained(
QWEN3_VL_MOE_PATH,
dtype=torch.bfloat16,
attn_implementation="flash_attention_2",
device_map="cuda",
trust_remote_code=True
).eval()
# Cannot understand, but must accept. Once there is no this code, it will appear cuda access illegal memory error in multi-GPU
torch.distributed.barrier()

loss_hf_text = self._forward(hf_model, type='text', device=device, sp_size=sp_size)
loss_hf_image = self._forward(hf_model, type='image', device=device, sp_size=sp_size)
# loss_hf_video = self._forward(hf_model, type='video', device=device, sp_size=sp_size)

del hf_model
torch.cuda.empty_cache()
Expand All @@ -222,8 +248,14 @@ def test_qwen3_5_vl_run(self, device, sp_size, tol):
loss_xtuner_image = self._forward(qwen3vl_model, type='image',device=device, sp_size=sp_size)
loss_xtuner_video = self._forward(qwen3vl_model, type='video',device=device, sp_size=sp_size)

self.assertTrue(torch.allclose(loss_xtuner_text, loss_hf_text.to(loss_xtuner_text.dtype), atol=tol, rtol=tol))
self.assertTrue(torch.allclose(loss_xtuner_image, loss_hf_image.to(loss_xtuner_image.dtype), atol=tol, rtol=tol))
self.assertTrue(
torch.allclose(loss_xtuner_text, loss_hf_text.to(loss_xtuner_text.dtype), atol=tol, rtol=tol),
f"Text loss mismatch: XTuner={loss_xtuner_text.item()}, HF={loss_hf_text.item()}",
)
self.assertTrue(
torch.allclose(loss_xtuner_image, loss_hf_image.to(loss_xtuner_image.dtype), atol=tol, rtol=tol),
f"Image loss mismatch: XTuner={loss_xtuner_image.item()}, HF={loss_hf_image.item()}",
)
# self.assertTrue(torch.allclose(loss_xtuner_video, loss_hf_video.to(loss_xtuner_video.dtype), atol=tol, rtol=tol))

del qwen3vl_model
Expand Down Expand Up @@ -263,12 +295,12 @@ def test_qwen3_5_vl_run_mtp(self, device, sp_size, tol):
self.create_pg(device)
self._patch_xtuner_fast_pos_embed_interpolate()

# pt29 + transformers 5.2.0 with XTUNER_DETERMINISTIC=true, which pins Triton autotune.
# pt29 + transformers 5.14.1 with XTUNER_DETERMINISTIC=true, which pins Triton autotune.
# The 11.5k-token video has a stable SP-specific LM-loss baseline on this path.
loss_reference = {
"text": 1.4981,
"image": 3.6109,
"video": {1: 9.3212, 4: 8.6532}[sp_size],
"video": {1: 8.5521, 4: 8.1323}[sp_size],
}

QWEN3_VL_MOE_PATH = os.environ["QWEN3_5_MOE_PATH"]
Expand Down Expand Up @@ -309,7 +341,7 @@ def test_qwen3_5_vl_run_mtp(self, device, sp_size, tol):
atol=tol,
rtol=tol
),
f"Expected text loss around {key}, but got {loss.item()}"
f"Expected {key} loss around {loss_reference[key]}, but got {loss.item()}"
)


Expand Down
9 changes: 7 additions & 2 deletions tests/model/test_qwen3_5_dense.py
Original file line number Diff line number Diff line change
Expand Up @@ -444,12 +444,17 @@ def _patch_fast_pos_embed_interpolate(self) -> None:
# HF's fast_pos_embed_interpolate returns fp32; the reused XTuner vision forward adds
# pos_embeds without a cast, so cast the result back to the pos_embed dtype here to
# avoid an fp32/bf16 LayerNorm mismatch.
from transformers.models.qwen3_5.modeling_qwen3_5 import Qwen3_5VisionModel
from transformers.vision_utils import get_vision_bilinear_indices_and_weights

from xtuner.v1.model.compose.qwen3_vl.modeling_vision import Qwen3VLVisionModel

def _interp(self, grid_thw):
return Qwen3_5VisionModel.fast_pos_embed_interpolate(self, grid_thw).to(self.pos_embed.weight.dtype)
indices, weights = get_vision_bilinear_indices_and_weights(
grid_thw,
num_grid_per_side=self.num_grid_per_side,
spatial_merge_size=self.config.spatial_merge_size,
)
return (self.pos_embed(indices) * weights[:, :, None]).sum(0).to(self.pos_embed.weight.dtype)

Qwen3VLVisionModel.fast_pos_embed_interpolate = _interp

Expand Down
19 changes: 19 additions & 0 deletions xtuner/v1/utils/test_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -267,3 +267,22 @@ def add_video_root(messages: list[dict], video_root: Path | str):
content["path"] = new_image_list
else:
content["path"] = str(content_path)


def get_qwen3_vl_video_chat_template(processor) -> str:
"""Return a Qwen3-VL template compatible with Transformers 5.14.1.

Qwen3VLProcessor.replace_video_token() already adds a vision-start/end pair for every video frame. Transformers
5.14.1's generic ProcessorMixin replaces only the video token, leaving the template's outer pair in the prompt. Use
a bare video token in the processor template so the expanded prompt keeps exactly one vision-start/end pair per
frame.
"""
chat_template = processor.chat_template
if not isinstance(chat_template, str):
raise TypeError("Qwen3-VL processor chat_template must be a string")

video_placeholder = processor.vision_start_token + processor.video_token + processor.vision_end_token
if video_placeholder not in chat_template:
raise ValueError("Qwen3-VL video placeholder was not found in processor chat_template")

return chat_template.replace(video_placeholder, processor.video_token)
Loading