diff --git a/nemo_rl/models/generation/vllm/quantization/fp8.py b/nemo_rl/models/generation/vllm/quantization/fp8.py index d89f6ed194..478ab9e210 100644 --- a/nemo_rl/models/generation/vllm/quantization/fp8.py +++ b/nemo_rl/models/generation/vllm/quantization/fp8.py @@ -903,56 +903,111 @@ def process_weights_after_loading_moe(self, layer) -> None: ) -def process_weights_after_loading_mxfp8_moe(self, layer) -> None: - """Shuffle weights and scales into FlashInfer TRTLLM MXFP8 layout.""" +def _mxfp8_moe_row_permutations( + layer, + w13_weight: torch.Tensor, + w2_weight: torch.Tensor, + is_gated: bool, + epilogue_tile_m: int, +) -> tuple[torch.Tensor, torch.Tensor]: + """Return row permutations equivalent to the FlashInfer shuffle calls.""" + perm_w13 = getattr(layer, "_mxfp8_shuffle_perm_w13", None) + perm_w2 = getattr(layer, "_mxfp8_shuffle_perm_w2", None) + if perm_w13 is None or perm_w2 is None: + from flashinfer.fused_moe.core import ( + get_reorder_rows_for_gated_act_gemm_row_indices, + ) + from flashinfer.utils import get_shuffle_matrix_a_row_indices + + perm_w13 = get_shuffle_matrix_a_row_indices(w13_weight[0], epilogue_tile_m) + if is_gated: + reorder = get_reorder_rows_for_gated_act_gemm_row_indices(w13_weight[0]) + perm_w13 = reorder[perm_w13] + perm_w2 = get_shuffle_matrix_a_row_indices(w2_weight[0], epilogue_tile_m) + layer._mxfp8_shuffle_perm_w13 = perm_w13 + layer._mxfp8_shuffle_perm_w2 = perm_w2 + device = w13_weight.device + return perm_w13.to(device), perm_w2.to(device) + + +def _shuffle_mxfp8_moe_batched( + layer, + w13_weight: torch.Tensor, + w2_weight: torch.Tensor, + w13_scale: torch.Tensor, + w2_scale: torch.Tensor, + is_gated: bool, + epilogue_tile_m: int, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + """Shuffle stacked expert values and scales with four batched gathers.""" + from flashinfer import block_scale_interleave + from vllm.model_executor.layers.quantization.utils.mxfp8_utils import ( + MXFP8_SCALE_DTYPE, + MXFP8_VALUE_DTYPE, + ) + + perm_w13, perm_w2 = _mxfp8_moe_row_permutations( + layer, w13_weight, w2_weight, is_gated, epilogue_tile_m + ) + num_experts = w13_weight.shape[0] + w13_u8 = w13_weight.view(torch.uint8) + w2_u8 = w2_weight.view(torch.uint8) + w13_shuffled = torch.index_select(w13_u8, 1, perm_w13) + w2_shuffled = torch.index_select(w2_u8, 1, perm_w2) + + w13_scale_u8 = pad_flashinfer_scale_k(w13_scale.view(torch.uint8)) + w2_scale_u8 = pad_flashinfer_scale_k(w2_scale.view(torch.uint8)) + assert w13_scale_u8.shape[1] % 128 == 0 + assert w2_scale_u8.shape[1] % 128 == 0 + w13_scale_gathered = torch.index_select(w13_scale_u8, 1, perm_w13) + w2_scale_gathered = torch.index_select(w2_scale_u8, 1, perm_w2) + w13_scale_shuffled = ( + block_scale_interleave(w13_scale_gathered) + .view(MXFP8_SCALE_DTYPE) + .view(num_experts, -1) + ) + w2_scale_shuffled = ( + block_scale_interleave(w2_scale_gathered) + .view(MXFP8_SCALE_DTYPE) + .view(num_experts, -1) + ) + return ( + w13_shuffled.view(MXFP8_VALUE_DTYPE), + w2_shuffled.view(MXFP8_VALUE_DTYPE), + w13_scale_shuffled, + w2_scale_shuffled, + ) + + +def _shuffle_mxfp8_moe_per_expert( + w13_weight: torch.Tensor, + w2_weight: torch.Tensor, + w13_scale: torch.Tensor, + w2_scale: torch.Tensor, + is_gated: bool, + epilogue_tile_m: int, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + """Run the original per-expert FlashInfer shuffle as a reference path.""" from flashinfer import ( reorder_rows_for_gated_act_gemm, shuffle_matrix_a, shuffle_matrix_sf_a, ) - from vllm.model_executor.layers.fused_moe import FusedMoeWeightScaleSupported - from vllm.model_executor.layers.quantization.utils.flashinfer_utils import ( - swap_w13_to_w31, - ) from vllm.model_executor.layers.quantization.utils.mxfp8_utils import ( MXFP8_SCALE_DTYPE, MXFP8_VALUE_DTYPE, ) - from vllm.model_executor.parameter import ModelWeightParameter - from vllm.model_executor.utils import set_weight_attrs - - epilogue_tile_m = 128 - num_experts = layer.w13_weight.shape[0] - is_gated = self.moe.is_act_and_mul - intermediate_size_factor = 2 if is_gated else 1 - - w13_weight = layer.w13_weight.data - if not hasattr(layer, "w13_weight_scale_from_checkpoint"): - w13_scale = layer.w13_weight_scale.data - else: - w13_scale = layer.w13_weight_scale_from_checkpoint.data - if is_gated: - # FI TRTLLM gated kernels use W31 ordering. Model checkpoints store - # gated projection as W13, so convert once before shuffling. - w13_weight = swap_w13_to_w31(w13_weight) - w13_scale = swap_w13_to_w31(w13_scale) - w2_weight = layer.w2_weight.data - if not hasattr(layer, "w2_weight_scale_from_checkpoint"): - w2_scale = layer.w2_weight_scale.data - else: - w2_scale = layer.w2_weight_scale_from_checkpoint.data + num_experts = w13_weight.shape[0] + w13_rows = w13_weight.shape[1] + w2_rows = w2_weight.shape[1] w13_weight_shuffled = [] w2_weight_shuffled = [] w13_scale_shuffled = [] w2_scale_shuffled = [] for i in range(num_experts): - w13_i = w13_weight[i].reshape( - intermediate_size_factor * layer.intermediate_size_per_partition, -1 - ) - w13_sf_i = w13_scale[i].reshape( - intermediate_size_factor * layer.intermediate_size_per_partition, -1 - ) + w13_i = w13_weight[i].reshape(w13_rows, -1) + w13_sf_i = w13_scale[i].reshape(w13_rows, -1) if is_gated: # Reorder rows for gated activation layout expected by TRTLLM. w13_i = reorder_rows_for_gated_act_gemm(w13_i.clone()) @@ -965,18 +1020,11 @@ def process_weights_after_loading_mxfp8_moe(self, layer) -> None: w13_weight_shuffled.append(w13_shuffled_i.contiguous().view(MXFP8_VALUE_DTYPE)) w2_weight_shuffled.append(w2_shuffled_i.contiguous().view(MXFP8_VALUE_DTYPE)) w13_sf_shuffled_i = shuffle_matrix_sf_a( - pad_flashinfer_scale_k( - w13_sf_i.view(torch.uint8).reshape( - intermediate_size_factor * layer.intermediate_size_per_partition, - -1, - ) - ), + pad_flashinfer_scale_k(w13_sf_i.view(torch.uint8).reshape(w13_rows, -1)), epilogue_tile_m, ) w2_sf_shuffled_i = shuffle_matrix_sf_a( - pad_flashinfer_scale_k( - w2_scale[i].view(torch.uint8).reshape(layer.hidden_size, -1) - ), + pad_flashinfer_scale_k(w2_scale[i].view(torch.uint8).reshape(w2_rows, -1)), epilogue_tile_m, ) w13_scale_shuffled.append( @@ -984,6 +1032,64 @@ def process_weights_after_loading_mxfp8_moe(self, layer) -> None: ) w2_scale_shuffled.append(w2_sf_shuffled_i.contiguous().view(MXFP8_SCALE_DTYPE)) + return ( + torch.stack(w13_weight_shuffled).contiguous(), + torch.stack(w2_weight_shuffled).contiguous(), + torch.stack(w13_scale_shuffled).contiguous(), + torch.stack(w2_scale_shuffled).contiguous(), + ) + + +def process_weights_after_loading_mxfp8_moe(self, layer) -> None: + """Shuffle weights and scales into FlashInfer TRTLLM MXFP8 layout.""" + from vllm.model_executor.layers.fused_moe import FusedMoeWeightScaleSupported + from vllm.model_executor.layers.fused_moe.oracle.fp8 import Fp8MoeBackend + from vllm.model_executor.layers.quantization.utils.flashinfer_utils import ( + swap_w13_to_w31, + ) + from vllm.model_executor.parameter import ModelWeightParameter + from vllm.model_executor.utils import set_weight_attrs + + if self.mxfp8_backend != Fp8MoeBackend.FLASHINFER_TRTLLM: + raise NotImplementedError( + "MXFP8 MoE refit layout conversion only supports FLASHINFER_TRTLLM; " + f"got {self.mxfp8_backend}." + ) + + epilogue_tile_m = 128 + is_gated = self.moe.is_act_and_mul + w13_weight = layer.w13_weight.data + if not hasattr(layer, "w13_weight_scale_from_checkpoint"): + w13_scale = layer.w13_weight_scale.data + else: + w13_scale = layer.w13_weight_scale_from_checkpoint.data + if is_gated: + # FI TRTLLM gated kernels use W31 ordering. Model checkpoints store + # gated projection as W13, so convert once before shuffling. + w13_weight = swap_w13_to_w31(w13_weight) + w13_scale = swap_w13_to_w31(w13_scale) + w2_weight = layer.w2_weight.data + if not hasattr(layer, "w2_weight_scale_from_checkpoint"): + w2_scale = layer.w2_weight_scale.data + else: + w2_scale = layer.w2_weight_scale_from_checkpoint.data + + shuffled = _shuffle_mxfp8_moe_batched( + layer, + w13_weight, + w2_weight, + w13_scale, + w2_scale, + is_gated, + epilogue_tile_m, + ) + ( + w13_weight_shuffled, + w2_weight_shuffled, + w13_scale_shuffled, + w2_scale_shuffled, + ) = shuffled + if not hasattr(layer, "w13_weight_scale_from_checkpoint"): layer.w13_weight_scale_from_checkpoint = ModelWeightParameter( data=layer.w13_weight_scale.data, @@ -1018,16 +1124,31 @@ def process_weights_after_loading_mxfp8_moe(self, layer) -> None: {"quant_method": FusedMoeWeightScaleSupported.BLOCK.value}, ) layer.w13_weight_scale = torch.nn.Parameter( - torch.stack(w13_scale_shuffled).contiguous(), requires_grad=False + w13_scale_shuffled, requires_grad=False ) layer.w2_weight_scale = torch.nn.Parameter( - torch.stack(w2_scale_shuffled).contiguous(), requires_grad=False + w2_scale_shuffled, requires_grad=False ) else: - layer.w13_weight_scale.copy_(torch.stack(w13_scale_shuffled).contiguous()) - layer.w2_weight_scale.copy_(torch.stack(w2_scale_shuffled).contiguous()) - layer.w13_weight.copy_(torch.stack(w13_weight_shuffled).contiguous()) - layer.w2_weight.copy_(torch.stack(w2_weight_shuffled).contiguous()) + layer.w13_weight_scale.copy_(w13_scale_shuffled) + layer.w2_weight_scale.copy_(w2_scale_shuffled) + layer.w13_weight.copy_(w13_weight_shuffled) + layer.w2_weight.copy_(w2_weight_shuffled) + + if self.moe_kernel is None: + from vllm.model_executor.layers.quantization.fp8 import make_fp8_moe_kernel + + self.moe_quant_config = self.get_fused_moe_quant_config(layer) + assert self.moe_quant_config is not None + assert self.experts_cls is not None + self.moe_kernel = make_fp8_moe_kernel( + moe_quant_config=self.moe_quant_config, + moe_config=self.moe, + fp8_backend=self.mxfp8_backend, + experts_cls=self.experts_cls, + routing_tables=layer._expert_routing_tables(), + layer=layer, + ) def process_weights_after_loading_kv(self, layer) -> None: diff --git a/tests/unit/models/generation/test_vllm_fp8_quantization.py b/tests/unit/models/generation/test_vllm_fp8_quantization.py index 5a5045c43e..076af3696e 100644 --- a/tests/unit/models/generation/test_vllm_fp8_quantization.py +++ b/tests/unit/models/generation/test_vllm_fp8_quantization.py @@ -15,6 +15,7 @@ import types import pytest +import torch pytestmark = pytest.mark.vllm @@ -80,6 +81,304 @@ def test_init_fp8_uses_mxfp8_quantization_config(fp8_module, monkeypatch): assert "VLLM_USE_DEEP_GEMM_E8M0" not in fp8.os.environ +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +@pytest.mark.parametrize( + ("is_gated", "intermediate_size", "hidden_size"), + [ + (True, 128, 256), + (True, 192, 128), + (False, 128, 256), + ], +) +def test_batched_moe_shuffle_matches_per_expert( + fp8_module, monkeypatch, is_gated, intermediate_size, hidden_size +): + pytest.importorskip("flashinfer") + fp8 = fp8_module + torch.manual_seed(0) + num_experts = 4 + w13_rows = (2 if is_gated else 1) * intermediate_size + + def rand_bytes(*shape): + return torch.randint(0, 256, shape, dtype=torch.uint8, device="cuda") + + w13_weight = rand_bytes(num_experts, w13_rows, hidden_size).view( + torch.float8_e4m3fn + ) + w2_weight = rand_bytes(num_experts, hidden_size, intermediate_size).view( + torch.float8_e4m3fn + ) + w13_scale = rand_bytes(num_experts, w13_rows, hidden_size // 32) + w2_scale = rand_bytes(num_experts, hidden_size, intermediate_size // 32) + + original_index_select = torch.index_select + index_select_out_tensors = [] + + def track_index_select(*args, **kwargs): + index_select_out_tensors.append(kwargs.get("out")) + return original_index_select(*args, **kwargs) + + monkeypatch.setattr(torch, "index_select", track_index_select) + batched = fp8._shuffle_mxfp8_moe_batched( + types.SimpleNamespace(), + w13_weight, + w2_weight, + w13_scale, + w2_scale, + is_gated, + 128, + ) + monkeypatch.setattr(torch, "index_select", original_index_select) + + assert len(index_select_out_tensors) == 4 + assert all(tensor is None for tensor in index_select_out_tensors) + + reference = fp8._shuffle_mxfp8_moe_per_expert( + w13_weight, + w2_weight, + w13_scale, + w2_scale, + is_gated, + 128, + ) + + for actual, expected in zip(batched, reference): + assert actual.shape == expected.shape + assert actual.dtype == expected.dtype + assert torch.equal(actual.view(torch.uint8), expected.view(torch.uint8)) + + +@pytest.mark.parametrize("is_gated", [True, False]) +def test_process_mxfp8_moe_refit_uses_batched_flashinfer_shuffle( + fp8_module, monkeypatch, is_gated +): + from vllm.model_executor.layers.fused_moe.oracle.fp8 import Fp8MoeBackend + + fp8 = fp8_module + fp8.global_fp8_config = fp8.FP8Config( + use_fp8_weights=True, + model_parallel_size=1, + is_mx=True, + ) + + w13_weight = torch.nn.Parameter(torch.zeros(2, 4, 3), requires_grad=False) + w2_weight = torch.nn.Parameter(torch.zeros(2, 3, 2), requires_grad=False) + w13_scale = torch.nn.Parameter(torch.zeros(2, 4, 1), requires_grad=False) + w2_scale = torch.nn.Parameter(torch.zeros(2, 3, 1), requires_grad=False) + w13_scale_from_checkpoint = torch.ones_like(w13_scale) + w2_scale_from_checkpoint = torch.ones_like(w2_scale) + layer = types.SimpleNamespace( + w13_weight=w13_weight, + w2_weight=w2_weight, + w13_weight_scale=w13_scale, + w2_weight_scale=w2_scale, + w13_weight_scale_from_checkpoint=types.SimpleNamespace( + data=w13_scale_from_checkpoint + ), + w2_weight_scale_from_checkpoint=types.SimpleNamespace( + data=w2_scale_from_checkpoint + ), + ) + moe_kernel = object() + moe_quant_config = object() + quant_method = types.SimpleNamespace( + moe=types.SimpleNamespace(is_act_and_mul=is_gated), + moe_kernel=moe_kernel, + moe_quant_config=moe_quant_config, + mxfp8_backend=Fp8MoeBackend.FLASHINFER_TRTLLM, + ) + shuffled = ( + torch.full_like(w13_weight, 1), + torch.full_like(w2_weight, 2), + torch.full_like(w13_scale, 3), + torch.full_like(w2_scale, 4), + ) + calls = [] + + def batched_shuffle(*args): + calls.append(("batched", args)) + return shuffled + + monkeypatch.setattr(fp8, "_shuffle_mxfp8_moe_batched", batched_shuffle) + + from vllm.model_executor.layers.quantization.utils import flashinfer_utils + + swap_calls = [] + + def swap_w13_to_w31(tensor): + swap_calls.append(tensor) + return tensor + + monkeypatch.setattr(flashinfer_utils, "swap_w13_to_w31", swap_w13_to_w31) + + parameter_ids = tuple( + id(parameter) + for parameter in ( + layer.w13_weight, + layer.w2_weight, + layer.w13_weight_scale, + layer.w2_weight_scale, + ) + ) + storage_ptrs = tuple( + parameter.data_ptr() + for parameter in ( + layer.w13_weight, + layer.w2_weight, + layer.w13_weight_scale, + layer.w2_weight_scale, + ) + ) + + fp8.process_weights_after_loading_mxfp8_moe(quant_method, layer) + + assert len(calls) == 1 + selected_path, args = calls[0] + assert selected_path == "batched" + assert args[0] is layer + args = args[1:] + assert args[0].data_ptr() == w13_weight.data_ptr() + assert args[1].data_ptr() == w2_weight.data_ptr() + assert args[2].data_ptr() == w13_scale_from_checkpoint.data_ptr() + assert args[3].data_ptr() == w2_scale_from_checkpoint.data_ptr() + assert args[4:] == (is_gated, 128) + expected_swap_ptrs = ( + [w13_weight.data_ptr(), w13_scale_from_checkpoint.data_ptr()] + if is_gated + else [] + ) + assert [tensor.data_ptr() for tensor in swap_calls] == expected_swap_ptrs + + parameters = ( + layer.w13_weight, + layer.w2_weight, + layer.w13_weight_scale, + layer.w2_weight_scale, + ) + assert tuple(id(parameter) for parameter in parameters) == parameter_ids + assert tuple(parameter.data_ptr() for parameter in parameters) == storage_ptrs + assert torch.equal(layer.w13_weight, shuffled[0]) + assert torch.equal(layer.w2_weight, shuffled[1]) + assert torch.equal(layer.w13_weight_scale, shuffled[2]) + assert torch.equal(layer.w2_weight_scale, shuffled[3]) + assert quant_method.moe_kernel is moe_kernel + assert quant_method.moe_quant_config is moe_quant_config + + +def test_process_mxfp8_moe_refit_rejects_non_flashinfer_backend(fp8_module): + from vllm.model_executor.layers.fused_moe.oracle.fp8 import Fp8MoeBackend + + quant_method = types.SimpleNamespace(mxfp8_backend=Fp8MoeBackend.DEEPGEMM) + + with pytest.raises( + NotImplementedError, + match="MXFP8 MoE refit layout conversion only supports FLASHINFER_TRTLLM", + ): + fp8_module.process_weights_after_loading_mxfp8_moe(quant_method, object()) + + +def test_process_mxfp8_moe_initializes_kernel_once(fp8_module, monkeypatch): + from vllm.model_executor.layers.fused_moe.oracle.fp8 import Fp8MoeBackend + + fp8 = fp8_module + fp8.global_fp8_config = fp8.FP8Config( + use_fp8_weights=True, + model_parallel_size=1, + is_mx=True, + ) + + layer = torch.nn.Module() + layer.w13_weight = torch.nn.Parameter(torch.zeros(2, 4, 3), requires_grad=False) + layer.w2_weight = torch.nn.Parameter(torch.zeros(2, 3, 2), requires_grad=False) + layer.w13_weight_scale = torch.nn.Parameter( + torch.zeros(2, 4, 1), requires_grad=False + ) + layer.w2_weight_scale = torch.nn.Parameter( + torch.zeros(2, 3, 1), requires_grad=False + ) + layer.w13_weight_scale.weight_loader = object() + layer.w2_weight_scale.weight_loader = object() + layer._expert_routing_tables = lambda: (None, None, None) + moe_config = types.SimpleNamespace(is_act_and_mul=False) + quant_config = object() + experts_cls = object() + quant_config_calls = [] + + def get_quant_config(_layer): + quant_config_calls.append(_layer) + return quant_config + + quant_method = types.SimpleNamespace( + moe=moe_config, + moe_kernel=None, + mxfp8_backend=Fp8MoeBackend.FLASHINFER_TRTLLM, + experts_cls=experts_cls, + get_fused_moe_quant_config=get_quant_config, + ) + kernel = object() + kernel_calls = [] + shuffle_calls = [] + + def shuffle(*args): + shuffle_calls.append(args) + fill = len(shuffle_calls) + return tuple(torch.full_like(tensor, fill) for tensor in args[1:5]) + + monkeypatch.setattr(fp8, "_shuffle_mxfp8_moe_batched", shuffle) + + from vllm.model_executor import parameter as vllm_parameter + from vllm.model_executor.layers.quantization import fp8 as vllm_fp8 + + monkeypatch.setattr(vllm_parameter, "get_tensor_model_parallel_rank", lambda: 0) + monkeypatch.setattr( + vllm_parameter, "get_tensor_model_parallel_world_size", lambda: 1 + ) + + def make_kernel(**kwargs): + kernel_calls.append(kwargs) + return kernel + + monkeypatch.setattr(vllm_fp8, "make_fp8_moe_kernel", make_kernel) + + fp8.process_weights_after_loading_mxfp8_moe(quant_method, layer) + + runtime_parameters = ( + layer.w13_weight, + layer.w2_weight, + layer.w13_weight_scale, + layer.w2_weight_scale, + ) + parameter_ids = tuple(id(parameter) for parameter in runtime_parameters) + storage_ptrs = tuple(parameter.data_ptr() for parameter in runtime_parameters) + + layer.w13_weight_scale_from_checkpoint.data.fill_(2) + layer.w2_weight_scale_from_checkpoint.data.fill_(2) + fp8.process_weights_after_loading_mxfp8_moe(quant_method, layer) + + assert quant_method.moe_kernel is kernel + assert quant_method.moe_quant_config is quant_config + assert quant_config_calls == [layer] + assert len(kernel_calls) == 1 + assert len(shuffle_calls) == 2 + refit_parameters = ( + layer.w13_weight, + layer.w2_weight, + layer.w13_weight_scale, + layer.w2_weight_scale, + ) + assert tuple(id(parameter) for parameter in refit_parameters) == parameter_ids + assert tuple(parameter.data_ptr() for parameter in refit_parameters) == storage_ptrs + assert all(torch.all(parameter == 2) for parameter in refit_parameters) + assert kernel_calls[0] == { + "moe_quant_config": quant_config, + "moe_config": moe_config, + "fp8_backend": Fp8MoeBackend.FLASHINFER_TRTLLM, + "experts_cls": experts_cls, + "routing_tables": (None, None, None), + "layer": layer, + } + + @pytest.mark.parametrize( ("field", "error"), [