[Bugfix] Fix broken profile_modular_kernel.py (#43300)

This commit is contained in:
x41lakazam
2026-06-11 12:17:23 +01:00
committed by GitHub
parent 432905d5d6
commit 3508cb78d4
2 changed files with 98 additions and 10 deletions
@@ -9,9 +9,19 @@ from typing import Any
import torch
from vllm.config import VllmConfig
from vllm.forward_context import set_forward_context
from vllm.model_executor.layers.fused_moe.activation import MoEActivation
from vllm.model_executor.layers.fused_moe.config import FusedMoEQuantConfig
from vllm.utils.torch_utils import set_random_seed
from vllm.v1.worker.workspace import init_workspace_manager
from .common import Config, RankTensors, WeightTensors, make_modular_kernel
from .common import (
Config,
RankTensors,
WeightTensors,
_make_gscale,
make_modular_kernel,
)
from .parallel_utils import ProcessGroupInfo, parallel_launch_with_config
@@ -35,7 +45,7 @@ def do_profile(
) as tprof:
fn(**fn_kwargs)
device = torch.accelerator.current_device_index()
torch.accelerator.synchronize(device=device)
torch.accelerator.synchronize(device)
# TODO (varun): Add a descriptive trace file name
tprof.export_chrome_trace(
@@ -56,24 +66,60 @@ def profile_modular_kernel(
# weights for rank
rank_weights = weights.slice_weights(pgi.rank, config.num_local_experts)
if config.quant_dtype == "nvfp4":
gscale = _make_gscale(config.num_local_experts)
else:
gscale = None
quant_config = FusedMoEQuantConfig.make(
config.quant_dtype,
w1_scale=rank_weights.w1_scale,
w2_scale=rank_weights.w2_scale,
a1_scale=rank_tensors.hidden_states_scale,
g1_alphas=(1 / rank_weights.w1_gs) if rank_weights.w1_gs is not None else None,
g2_alphas=(1 / rank_weights.w2_gs) if rank_weights.w2_gs is not None else None,
a1_gscale=gscale,
a2_gscale=gscale,
block_shape=config.quant_block_shape,
per_act_token_quant=config.is_per_act_token_quant,
per_out_ch_quant=config.is_per_out_ch_quant,
)
# make modular kernel
mk = make_modular_kernel(config, vllm_config, weights)
mk = make_modular_kernel(config, vllm_config, quant_config)
topk_ids = rank_tensors.topk_ids.to(
mk.prepare_finalize.topk_indices_dtype() or rank_tensors.topk_ids.dtype
)
# impls might update the tensor in place
hidden_states = rank_tensors.hidden_states.clone()
mk_kwargs = {
"hidden_states": rank_tensors.hidden_states,
"hidden_states": hidden_states,
"w1": rank_weights.w1,
"w2": rank_weights.w2,
"topk_weights": rank_tensors.topk_weights,
"topk_ids": rank_tensors.topk_ids,
"topk_ids": topk_ids,
"activation": MoEActivation.SILU,
"expert_map": rank_tensors.expert_map,
"w1_scale": rank_weights.w1_scale,
"w2_scale": rank_weights.w2_scale,
"a1_scale": rank_tensors.hidden_states_scale,
"global_num_experts": config.E,
"apply_router_weight_on_input": config.topk == 1,
"apply_router_weight_on_input": config.topk == 1
and config.supports_apply_weight_on_input(),
}
do_profile(mk.apply, mk_kwargs, pgi, config)
num_tokens = hidden_states.shape[0]
num_tokens_across_dp = torch.tensor(
[num_tokens] * config.world_size, device="cpu", dtype=torch.int
)
with set_forward_context(
None,
vllm_config,
num_tokens=num_tokens,
num_tokens_across_dp=num_tokens_across_dp,
):
do_profile(mk.apply, mk_kwargs, pgi, config)
def rank_worker(
@@ -85,6 +131,10 @@ def rank_worker(
):
set_random_seed(pgi.rank)
# workspace manager is normally initialized by GPUModelRunner; we initialize
# it here for the standalone benchmark process.
init_workspace_manager(torch.device(f"cuda:{pgi.local_rank}"))
# get weights to this device
weights.to_current_device()
@@ -0,0 +1,38 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import pytest
import torch
from vllm.model_executor.layers.fused_moe import TritonExperts
from vllm.model_executor.layers.fused_moe.prepare_finalize import (
MoEPrepareAndFinalizeNoDPEPModular,
)
from .modular_kernel_tools.common import Config
from .modular_kernel_tools.profile_modular_kernel import run
@pytest.mark.skipif(
not torch.cuda.is_available(),
reason="profile_modular_kernel requires a CUDA device",
)
def test_profile_modular_kernel_smoke(tmp_path):
config = Config(
Ms=[16],
K=128,
N=256,
E=4,
topks=[2],
dtype=torch.bfloat16,
quant_config=None,
prepare_finalize_type=MoEPrepareAndFinalizeNoDPEPModular,
fused_experts_type=TritonExperts,
world_size=1,
torch_trace_dir_path=str(tmp_path),
)
run(config)
traces = list(tmp_path.glob("m*_*_trace.json"))
assert traces, "profile_modular_kernel.run did not emit any chrome traces"