From 08a37f51f2becd4978e7c236ac334f0d6122332e Mon Sep 17 00:00:00 2001 From: huanghaian Date: Fri, 14 Aug 2026 08:24:47 +0000 Subject: [PATCH 1/2] support transformers 5.14.1 --- tests/datasets/test_qwen35_vl_tokenize_fn.py | 16 +++-- tests/datasets/test_qwen3_vl_tokenize_fn.py | 16 +++-- tests/model/test_qwen3_5.py | 30 +++++--- tests/model/test_qwen3_5_dense.py | 9 ++- xtuner/v1/model/moe/moe.py | 4 ++ .../module/decoder_layer/moe_decoder_layer.py | 72 +++++++++++++------ xtuner/v1/utils/test_utils.py | 19 +++++ 7 files changed, 122 insertions(+), 44 deletions(-) diff --git a/tests/datasets/test_qwen35_vl_tokenize_fn.py b/tests/datasets/test_qwen35_vl_tokenize_fn.py index 2149989766..76c0337b51 100644 --- a/tests/datasets/test_qwen35_vl_tokenize_fn.py +++ b/tests/datasets/test_qwen35_vl_tokenize_fn.py @@ -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 @@ -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'] diff --git a/tests/datasets/test_qwen3_vl_tokenize_fn.py b/tests/datasets/test_qwen3_vl_tokenize_fn.py index 1e6fa4f87f..0c6ffde3aa 100644 --- a/tests/datasets/test_qwen3_vl_tokenize_fn.py +++ b/tests/datasets/test_qwen3_vl_tokenize_fn.py @@ -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"] @@ -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'] diff --git a/tests/model/test_qwen3_5.py b/tests/model/test_qwen3_5.py index c3833f05bc..ed1fd50ebf 100644 --- a/tests/model/test_qwen3_5.py +++ b/tests/model/test_qwen3_5.py @@ -32,15 +32,19 @@ ) class TestQwen3_5_VL(DeterministicDDPTestCase): 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 @@ -213,6 +217,7 @@ def test_qwen3_5_vl_run(self, device, sp_size, tol): # hf_save_cfg of text_model is ignored to align with transformers's forward result model_cfg.text_config.hf_save_cfg = HFSaveCfg() model_cfg.text_config.router_compute_dtype = "native" + model_cfg.text_config.hf_compatible_moe_combine = True qwen3vl_model = model_cfg.build()._to_device_dtype(dtype=torch.bfloat16, skip_buffers_dtype=True) qwen3vl_model.from_hf(QWEN3_VL_MOE_PATH) @@ -222,8 +227,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 @@ -235,6 +246,7 @@ def test_qwen3_5_vl_run(self, device, sp_size, tol): # hf_save_cfg of text_model is ignored to align with transformers's forward result model_cfg.text_config.hf_save_cfg = HFSaveCfg() model_cfg.text_config.router_compute_dtype = "native" + model_cfg.text_config.hf_compatible_moe_combine = True qwen3vl_model = model_cfg.build()._to_device_dtype(dtype=torch.bfloat16, skip_buffers_dtype=True) fsdp_config = FSDPConfig(cpu_offload=False) @@ -263,12 +275,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"] @@ -309,7 +321,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()}" ) diff --git a/tests/model/test_qwen3_5_dense.py b/tests/model/test_qwen3_5_dense.py index 8ad08d46af..181fadb0ab 100644 --- a/tests/model/test_qwen3_5_dense.py +++ b/tests/model/test_qwen3_5_dense.py @@ -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 diff --git a/xtuner/v1/model/moe/moe.py b/xtuner/v1/model/moe/moe.py index f27e0a2dbc..36d2a28640 100644 --- a/xtuner/v1/model/moe/moe.py +++ b/xtuner/v1/model/moe/moe.py @@ -154,6 +154,8 @@ class MoEConfig(TransformerConfig): return_router_results: bool = False gate_bias: bool = False router_compute_dtype: Literal["float32", "native"] = "float32" + # Reproduce Hugging Face's native-dtype weighted expert reduction. This is slower and only supports EP=1. + hf_compatible_moe_combine: bool = False moe_bias: bool = False moe_act_fn_cfg: MoEActFnConfig = MoEActFnConfig() mtp_config: MTPConfig | None = None @@ -995,6 +997,7 @@ def build_layers(self, config: MoEConfig) -> nn.ModuleDict: generate_config=config.generate_config, router_config=config.router, router_compute_dtype=config.router_compute_dtype, + hf_compatible_moe_combine=config.hf_compatible_moe_combine, moe_act_fn_cfg=config.moe_act_fn_cfg, float8_cfg=config.float8_cfg, layer_idx=layer_idx, @@ -1060,6 +1063,7 @@ def build_mtp_block(self, config: MoEConfig) -> MTPBlock: generate_config=config.generate_config, router_config=config.router, router_compute_dtype=config.router_compute_dtype, + hf_compatible_moe_combine=config.hf_compatible_moe_combine, moe_act_fn_cfg=config.moe_act_fn_cfg, float8_cfg=config.float8_cfg, layer_idx=config.num_hidden_layers + i, diff --git a/xtuner/v1/module/decoder_layer/moe_decoder_layer.py b/xtuner/v1/module/decoder_layer/moe_decoder_layer.py index 00e5d6c27e..919b0364cb 100644 --- a/xtuner/v1/module/decoder_layer/moe_decoder_layer.py +++ b/xtuner/v1/module/decoder_layer/moe_decoder_layer.py @@ -28,6 +28,7 @@ from xtuner.v1.module.dispatcher import ( CombineResult, DispatchResult, + NaiveDispatcher, PostDispatchResult, PreCombineResult, PreDispatchResult, @@ -217,6 +218,7 @@ def __init__( generate_config: GenerateConfig | None = None, router_config: GreedyRouterConfig | NoAuxRouterConfig, router_compute_dtype: Literal["float32", "native"] = "float32", + hf_compatible_moe_combine: bool = False, moe_act_fn_cfg: MoEActFnConfig, float8_cfg: Float8Config | None = None, layer_idx: int = 0, @@ -288,6 +290,11 @@ def __init__( training_dtype="fp8" if float8_cfg is not None else "bf16", generate_dtype=generate_config.dtype if generate_config is not None else "bf16", ) + if hf_compatible_moe_combine and not isinstance(self.dispatcher, NaiveDispatcher): + raise ValueError("hf_compatible_moe_combine only supports expert parallel size 1") + if hf_compatible_moe_combine and float8_cfg is not None: + raise ValueError("hf_compatible_moe_combine does not support float8 experts") + self.hf_compatible_moe_combine = hf_compatible_moe_combine def forward( self, @@ -371,6 +378,22 @@ def _hf_expert_forward_for_debug(self, hidden_states: torch.Tensor, router_resul combined_hidden_states = combined_hidden_states.view(*origin_shape) return combined_hidden_states + @staticmethod + def _hf_compatible_moe_combine( + experts_out: torch.Tensor, router_results: RouterResults, origin_shape: torch.Size + ) -> torch.Tensor: + """Reduce routed expert outputs in the model dtype, matching Hugging + Face.""" + topk_ids = router_results["topk_ids"] + topk_weights = router_results["topk_weights"].to(experts_out.dtype) + # The naive dispatcher groups expert inputs with the same stable ordering. + sorted_indices = torch.argsort(topk_ids.flatten(), stable=True) + unpermuted_experts = torch.empty_like(experts_out) + unpermuted_experts.index_copy_(0, sorted_indices, experts_out) + unpermuted_experts = unpermuted_experts.view(*topk_ids.shape, experts_out.shape[-1]) + combined = (unpermuted_experts * topk_weights.unsqueeze(-1)).sum(dim=1) + return combined.view(*origin_shape) + def _forward( self, hidden_states: torch.Tensor, @@ -421,30 +444,33 @@ def _forward( # post_dispatched.get("row_ids_map"), # type: ignore[arg-type] # dispatched["topk_weights"], # ) - pre_combined = self.dispatcher.combine_preprocess( - hidden_states=experts_out, - pre_dispatched=pre_dispatched, - dispatched=dispatched, - post_dispatched=post_dispatched, - decoding=False, - ) + if self.hf_compatible_moe_combine: + combined_hidden_states = self._hf_compatible_moe_combine(experts_out, router_results, origin_shape) + else: + pre_combined = self.dispatcher.combine_preprocess( + hidden_states=experts_out, + pre_dispatched=pre_dispatched, + dispatched=dispatched, + post_dispatched=post_dispatched, + decoding=False, + ) - combined = self.dispatcher.combine( - pre_dispatched=pre_dispatched, - dispatched=dispatched, - post_dispatched=post_dispatched, - pre_combined=pre_combined, - decoding=False, - ) - post_combined = self.dispatcher.combine_postprocess( - pre_dispatched=pre_dispatched, - dispatched=dispatched, - post_dispatched=post_dispatched, - pre_combined=pre_combined, - combined=combined, - ) - combined_hidden_states = post_combined["hidden_states"] - combined_hidden_states = combined_hidden_states.view(*origin_shape) + combined = self.dispatcher.combine( + pre_dispatched=pre_dispatched, + dispatched=dispatched, + post_dispatched=post_dispatched, + pre_combined=pre_combined, + decoding=False, + ) + post_combined = self.dispatcher.combine_postprocess( + pre_dispatched=pre_dispatched, + dispatched=dispatched, + post_dispatched=post_dispatched, + pre_combined=pre_combined, + combined=combined, + ) + combined_hidden_states = post_combined["hidden_states"] + combined_hidden_states = combined_hidden_states.view(*origin_shape) # debug for aligning with hf implementation. # combined_hidden_states = self._hf_expert_forward_for_debug(hidden_states, router_results, origin_shape) diff --git a/xtuner/v1/utils/test_utils.py b/xtuner/v1/utils/test_utils.py index 9748e5a8f8..73bd0cde55 100644 --- a/xtuner/v1/utils/test_utils.py +++ b/xtuner/v1/utils/test_utils.py @@ -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) From c4c3d8fb7f95ae07f5f4f559329ca03ac33db6ba Mon Sep 17 00:00:00 2001 From: huanghaian Date: Tue, 18 Aug 2026 03:22:20 +0000 Subject: [PATCH 2/2] update --- tests/model/test_qwen3_5.py | 50 +++++++++---- xtuner/v1/model/moe/moe.py | 4 -- .../module/decoder_layer/moe_decoder_layer.py | 72 ++++++------------- 3 files changed, 58 insertions(+), 68 deletions(-) diff --git a/tests/model/test_qwen3_5.py b/tests/model/test_qwen3_5.py index ed1fd50ebf..bdf3c870ab 100644 --- a/tests/model/test_qwen3_5.py +++ b/tests/model/test_qwen3_5.py @@ -1,5 +1,7 @@ import os import unittest +from unittest.mock import patch + import parametrize import torch from packaging.version import Version @@ -31,6 +33,25 @@ 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 transformers.vision_utils import get_vision_bilinear_indices_and_weights @@ -195,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() @@ -217,7 +239,6 @@ def test_qwen3_5_vl_run(self, device, sp_size, tol): # hf_save_cfg of text_model is ignored to align with transformers's forward result model_cfg.text_config.hf_save_cfg = HFSaveCfg() model_cfg.text_config.router_compute_dtype = "native" - model_cfg.text_config.hf_compatible_moe_combine = True qwen3vl_model = model_cfg.build()._to_device_dtype(dtype=torch.bfloat16, skip_buffers_dtype=True) qwen3vl_model.from_hf(QWEN3_VL_MOE_PATH) @@ -246,7 +267,6 @@ def test_qwen3_5_vl_run(self, device, sp_size, tol): # hf_save_cfg of text_model is ignored to align with transformers's forward result model_cfg.text_config.hf_save_cfg = HFSaveCfg() model_cfg.text_config.router_compute_dtype = "native" - model_cfg.text_config.hf_compatible_moe_combine = True qwen3vl_model = model_cfg.build()._to_device_dtype(dtype=torch.bfloat16, skip_buffers_dtype=True) fsdp_config = FSDPConfig(cpu_offload=False) diff --git a/xtuner/v1/model/moe/moe.py b/xtuner/v1/model/moe/moe.py index 36d2a28640..f27e0a2dbc 100644 --- a/xtuner/v1/model/moe/moe.py +++ b/xtuner/v1/model/moe/moe.py @@ -154,8 +154,6 @@ class MoEConfig(TransformerConfig): return_router_results: bool = False gate_bias: bool = False router_compute_dtype: Literal["float32", "native"] = "float32" - # Reproduce Hugging Face's native-dtype weighted expert reduction. This is slower and only supports EP=1. - hf_compatible_moe_combine: bool = False moe_bias: bool = False moe_act_fn_cfg: MoEActFnConfig = MoEActFnConfig() mtp_config: MTPConfig | None = None @@ -997,7 +995,6 @@ def build_layers(self, config: MoEConfig) -> nn.ModuleDict: generate_config=config.generate_config, router_config=config.router, router_compute_dtype=config.router_compute_dtype, - hf_compatible_moe_combine=config.hf_compatible_moe_combine, moe_act_fn_cfg=config.moe_act_fn_cfg, float8_cfg=config.float8_cfg, layer_idx=layer_idx, @@ -1063,7 +1060,6 @@ def build_mtp_block(self, config: MoEConfig) -> MTPBlock: generate_config=config.generate_config, router_config=config.router, router_compute_dtype=config.router_compute_dtype, - hf_compatible_moe_combine=config.hf_compatible_moe_combine, moe_act_fn_cfg=config.moe_act_fn_cfg, float8_cfg=config.float8_cfg, layer_idx=config.num_hidden_layers + i, diff --git a/xtuner/v1/module/decoder_layer/moe_decoder_layer.py b/xtuner/v1/module/decoder_layer/moe_decoder_layer.py index 919b0364cb..00e5d6c27e 100644 --- a/xtuner/v1/module/decoder_layer/moe_decoder_layer.py +++ b/xtuner/v1/module/decoder_layer/moe_decoder_layer.py @@ -28,7 +28,6 @@ from xtuner.v1.module.dispatcher import ( CombineResult, DispatchResult, - NaiveDispatcher, PostDispatchResult, PreCombineResult, PreDispatchResult, @@ -218,7 +217,6 @@ def __init__( generate_config: GenerateConfig | None = None, router_config: GreedyRouterConfig | NoAuxRouterConfig, router_compute_dtype: Literal["float32", "native"] = "float32", - hf_compatible_moe_combine: bool = False, moe_act_fn_cfg: MoEActFnConfig, float8_cfg: Float8Config | None = None, layer_idx: int = 0, @@ -290,11 +288,6 @@ def __init__( training_dtype="fp8" if float8_cfg is not None else "bf16", generate_dtype=generate_config.dtype if generate_config is not None else "bf16", ) - if hf_compatible_moe_combine and not isinstance(self.dispatcher, NaiveDispatcher): - raise ValueError("hf_compatible_moe_combine only supports expert parallel size 1") - if hf_compatible_moe_combine and float8_cfg is not None: - raise ValueError("hf_compatible_moe_combine does not support float8 experts") - self.hf_compatible_moe_combine = hf_compatible_moe_combine def forward( self, @@ -378,22 +371,6 @@ def _hf_expert_forward_for_debug(self, hidden_states: torch.Tensor, router_resul combined_hidden_states = combined_hidden_states.view(*origin_shape) return combined_hidden_states - @staticmethod - def _hf_compatible_moe_combine( - experts_out: torch.Tensor, router_results: RouterResults, origin_shape: torch.Size - ) -> torch.Tensor: - """Reduce routed expert outputs in the model dtype, matching Hugging - Face.""" - topk_ids = router_results["topk_ids"] - topk_weights = router_results["topk_weights"].to(experts_out.dtype) - # The naive dispatcher groups expert inputs with the same stable ordering. - sorted_indices = torch.argsort(topk_ids.flatten(), stable=True) - unpermuted_experts = torch.empty_like(experts_out) - unpermuted_experts.index_copy_(0, sorted_indices, experts_out) - unpermuted_experts = unpermuted_experts.view(*topk_ids.shape, experts_out.shape[-1]) - combined = (unpermuted_experts * topk_weights.unsqueeze(-1)).sum(dim=1) - return combined.view(*origin_shape) - def _forward( self, hidden_states: torch.Tensor, @@ -444,33 +421,30 @@ def _forward( # post_dispatched.get("row_ids_map"), # type: ignore[arg-type] # dispatched["topk_weights"], # ) - if self.hf_compatible_moe_combine: - combined_hidden_states = self._hf_compatible_moe_combine(experts_out, router_results, origin_shape) - else: - pre_combined = self.dispatcher.combine_preprocess( - hidden_states=experts_out, - pre_dispatched=pre_dispatched, - dispatched=dispatched, - post_dispatched=post_dispatched, - decoding=False, - ) + pre_combined = self.dispatcher.combine_preprocess( + hidden_states=experts_out, + pre_dispatched=pre_dispatched, + dispatched=dispatched, + post_dispatched=post_dispatched, + decoding=False, + ) - combined = self.dispatcher.combine( - pre_dispatched=pre_dispatched, - dispatched=dispatched, - post_dispatched=post_dispatched, - pre_combined=pre_combined, - decoding=False, - ) - post_combined = self.dispatcher.combine_postprocess( - pre_dispatched=pre_dispatched, - dispatched=dispatched, - post_dispatched=post_dispatched, - pre_combined=pre_combined, - combined=combined, - ) - combined_hidden_states = post_combined["hidden_states"] - combined_hidden_states = combined_hidden_states.view(*origin_shape) + combined = self.dispatcher.combine( + pre_dispatched=pre_dispatched, + dispatched=dispatched, + post_dispatched=post_dispatched, + pre_combined=pre_combined, + decoding=False, + ) + post_combined = self.dispatcher.combine_postprocess( + pre_dispatched=pre_dispatched, + dispatched=dispatched, + post_dispatched=post_dispatched, + pre_combined=pre_combined, + combined=combined, + ) + combined_hidden_states = post_combined["hidden_states"] + combined_hidden_states = combined_hidden_states.view(*origin_shape) # debug for aligning with hf implementation. # combined_hidden_states = self._hf_expert_forward_for_debug(hidden_states, router_results, origin_shape)