Skip to content

Commit 657b04d

Browse files
committed
polish code
1 parent 29fc068 commit 657b04d

11 files changed

Lines changed: 37 additions & 37 deletions

File tree

vllm_fl/dispatch/backends/flaggems/impl/activation.py

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,8 @@
77
from __future__ import annotations
88

99
import torch
10-
from vllm_fl.utils import use_flaggems_vllm
10+
11+
import vllm_fl.envs as fl_envs
1112

1213

1314
def silu_and_mul_flaggems(obj, x: torch.Tensor) -> torch.Tensor:
@@ -21,7 +22,7 @@ def silu_and_mul_flaggems(obj, x: torch.Tensor) -> torch.Tensor:
2122
Returns:
2223
Output tensor of shape [..., d]
2324
"""
24-
if use_flaggems_vllm():
25+
if fl_envs.VLLM_FL_USE_FLAGGEMS_VLLM:
2526
from flaggems_vllm.ops.silu_and_mul import silu_and_mul
2627
else:
2728
from flag_gems import silu_and_mul
@@ -42,7 +43,7 @@ def gelu_and_mul_flaggems(obj, x: torch.Tensor) -> torch.Tensor:
4243
Returns:
4344
Output tensor of shape [..., d]
4445
"""
45-
if use_flaggems_vllm():
46+
if fl_envs.VLLM_FL_USE_FLAGGEMS_VLLM:
4647
from flaggems_vllm.ops.gelu_and_mul import gelu_and_mul
4748
else:
4849
from flag_gems import gelu_and_mul
@@ -69,7 +70,7 @@ def silu_and_mul_with_clamp_flaggems(x: torch.Tensor, swiglu_limit: torch.Tensor
6970
Returns:
7071
Output tensor of shape [..., d]
7172
"""
72-
if use_flaggems_vllm():
73+
if fl_envs.VLLM_FL_USE_FLAGGEMS_VLLM:
7374
from flaggems_vllm.ops.silu_and_mul_with_clamp import silu_and_mul_with_clamp_kernel
7475
else:
7576
from flag_gems.fused.silu_and_mul_with_clamp import silu_and_mul_with_clamp_kernel

vllm_fl/dispatch/backends/flaggems/impl/attention.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -40,9 +40,10 @@
4040
)
4141
from vllm.v1.kv_cache_interface import AttentionSpec
4242
from vllm.platforms.interface import DeviceCapability
43-
from vllm_fl.utils import use_flaggems_vllm
4443

45-
if use_flaggems_vllm():
44+
import vllm_fl.envs as fl_envs
45+
46+
if fl_envs.VLLM_FL_USE_FLAGGEMS_VLLM:
4647
from flaggems_vllm.ops.flash_attn_varlen_func import flash_attn_varlen_func
4748
from flaggems_vllm.ops.reshape_and_cache_flash import reshape_and_cache_flash
4849
else:

vllm_fl/dispatch/backends/flaggems/impl/fused_moe.py

Lines changed: 8 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -9,7 +9,8 @@
99
import torch
1010
from vllm.triton_utils import triton
1111
from vllm.utils.math_utils import round_up
12-
from vllm_fl.utils import use_flaggems_vllm
12+
13+
import vllm_fl.envs as fl_envs
1314

1415

1516
def moe_align_block_size_flaggems(
@@ -20,7 +21,7 @@ def moe_align_block_size_flaggems(
2021
pad_sorted_ids: bool = False,
2122
ignore_invalid_experts: bool = False,
2223
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
23-
if use_flaggems_vllm():
24+
if fl_envs.VLLM_FL_USE_FLAGGEMS_VLLM:
2425
from flaggems_vllm.ops.moe_align_block_size import moe_align_block_size_triton
2526
else:
2627
from flag_gems import moe_align_block_size_triton
@@ -60,7 +61,7 @@ def moe_align_block_size_flaggems(
6061
def topk_softmax_flaggems(
6162
topk_weights, topk_indices, token_expert_indices, gating_output, renormalize=False
6263
):
63-
if use_flaggems_vllm():
64+
if fl_envs.VLLM_FL_USE_FLAGGEMS_VLLM:
6465
from flaggems_vllm.ops.topk_softmax import topk_softmax
6566
topk_softmax(
6667
topk_weights,
@@ -141,7 +142,7 @@ def _alloc_topk_buffers():
141142
return topk_weights, topk_ids
142143

143144
elif scoring_func == "sqrtsoftplus":
144-
if use_flaggems_vllm():
145+
if fl_envs.VLLM_FL_USE_FLAGGEMS_VLLM:
145146
from flaggems_vllm.ops.topk_softplus_sqrt import topk_softplus_sqrt
146147
else:
147148
from flag_gems import topk_softplus_sqrt
@@ -205,7 +206,7 @@ def invoke_fused_moe_triton_kernel_flaggems(
205206
block_shape=None,
206207
B_bias=None,
207208
):
208-
if use_flaggems_vllm():
209+
if fl_envs.VLLM_FL_USE_FLAGGEMS_VLLM:
209210
from flaggems_vllm.ops.invoke_fused_moe_kernel import invoke_fused_moe_triton_kernel
210211
else:
211212
from flag_gems import invoke_fused_moe_triton_kernel
@@ -244,7 +245,7 @@ def grouped_topk_flaggems(
244245
bias,
245246
scoring_func=0,
246247
):
247-
if use_flaggems_vllm():
248+
if fl_envs.VLLM_FL_USE_FLAGGEMS_VLLM:
248249
from flaggems_vllm.ops.grouped_topk import grouped_topk
249250
else:
250251
from flag_gems import grouped_topk
@@ -262,7 +263,7 @@ def grouped_topk_flaggems(
262263

263264

264265
def moe_sum_flaggems(inp, out):
265-
if use_flaggems_vllm():
266+
if fl_envs.VLLM_FL_USE_FLAGGEMS_VLLM:
266267
from flaggems_vllm.ops.moe_sum import moe_sum
267268
else:
268269
from flag_gems import moe_sum

vllm_fl/dispatch/backends/flaggems/impl/mhc.py

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,8 @@
55
"""
66

77
import torch
8-
from vllm_fl.utils import use_flaggems_vllm
8+
9+
import vllm_fl.envs as fl_envs
910

1011

1112
def mhc_pre_flaggems(
@@ -23,7 +24,7 @@ def mhc_pre_flaggems(
2324
"""FlagGems native implementation of mhc_pre."""
2425
import importlib
2526

26-
if use_flaggems_vllm():
27+
if fl_envs.VLLM_FL_USE_FLAGGEMS_VLLM:
2728
from flaggems_vllm.ops.mhc.mhc_pre import mhc_pre
2829
_mhc_pre_mod = importlib.import_module('flaggems_vllm.ops.mhc.mhc_pre')
2930
else:
@@ -85,7 +86,7 @@ def mhc_post_flaggems(
8586
comb: torch.Tensor,
8687
) -> torch.Tensor:
8788
"""FlagGems native implementation of mhc_post."""
88-
if use_flaggems_vllm():
89+
if fl_envs.VLLM_FL_USE_FLAGGEMS_VLLM:
8990
from flaggems_vllm.ops.mhc_post import mhc_post
9091
else:
9192
from flag_gems import mhc_post
@@ -105,7 +106,7 @@ def hc_head_fused_kernel_flaggems(
105106
hc_mult: int,
106107
) -> None:
107108
"""FlagGems native implementation of hc_head_fused_kernel. Mutates `out` in-place."""
108-
if use_flaggems_vllm():
109+
if fl_envs.VLLM_FL_USE_FLAGGEMS_VLLM:
109110
from flaggems_vllm.ops.hc_head_fused_kernel import hc_head_fused_kernel
110111
else:
111112
from flag_gems import hc_head_fused_kernel

vllm_fl/dispatch/backends/flaggems/impl/mla.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -23,9 +23,9 @@
2323
MLACommonMetadata,
2424
)
2525

26-
from vllm_fl.utils import use_flaggems_vllm
26+
import vllm_fl.envs as fl_envs
2727

28-
if use_flaggems_vllm():
28+
if fl_envs.VLLM_FL_USE_FLAGGEMS_VLLM:
2929
from flaggems_vllm.ops.flash_attn_varlen_func import flash_attn_varlen_func
3030
from flaggems_vllm.ops.flash_mla import flash_mla
3131
else:

vllm_fl/dispatch/backends/flaggems/impl/normalization.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -9,7 +9,8 @@
99
from typing import Optional, Union
1010

1111
import torch
12-
from vllm_fl.utils import use_flaggems_vllm
12+
13+
import vllm_fl.envs as fl_envs
1314

1415

1516
def rms_norm_flaggems(
@@ -28,7 +29,7 @@ def rms_norm_flaggems(
2829
Returns:
2930
Normalized tensor, or tuple of (normalized, residual) if residual is provided
3031
"""
31-
if use_flaggems_vllm():
32+
if fl_envs.VLLM_FL_USE_FLAGGEMS_VLLM:
3233
from flaggems_vllm.ops.rms_norm import gems_rms_forward
3334
else:
3435
from flag_gems import rms_norm_forward as gems_rms_forward

vllm_fl/dispatch/backends/flaggems/impl/rotary.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,8 @@
77
from __future__ import annotations
88

99
import torch
10-
from vllm_fl.utils import use_flaggems_vllm
10+
11+
import vllm_fl.envs as fl_envs
1112

1213

1314
def rotary_embedding_flaggems(
@@ -36,7 +37,7 @@ def rotary_embedding_flaggems(
3637
Returns:
3738
Tuple of (embedded_query, embedded_key)
3839
"""
39-
if use_flaggems_vllm():
40+
if fl_envs.VLLM_FL_USE_FLAGGEMS_VLLM:
4041
from flaggems_vllm.ops.rope import gems_rope_forward
4142
else:
4243
from flag_gems import apply_rotary_pos_emb as gems_rope_forward

vllm_fl/dispatch/backends/flaggems/impl/router_gemm.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -3,13 +3,14 @@
33
"""FlagGems MoE router GEMM implementation."""
44

55
import torch
6-
from vllm_fl.utils import use_flaggems_vllm
6+
7+
import vllm_fl.envs as fl_envs
78

89

910
def router_gemm_bf16_fp32_flaggems(
1011
x: torch.Tensor, weight: torch.Tensor
1112
) -> torch.Tensor:
12-
if use_flaggems_vllm():
13+
if fl_envs.VLLM_FL_USE_FLAGGEMS_VLLM:
1314
from flaggems_vllm.ops.router_gemm import router_gemm
1415
else:
1516
from flag_gems import router_gemm

vllm_fl/envs.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,9 @@
1313
"FLAGGEMS_ENABLE_OPLIST_PATH", "/tmp/flaggems_enable_oplist.txt"
1414
),
1515
"USE_FLAGGEMS": use_flaggems,
16+
"VLLM_FL_USE_FLAGGEMS_VLLM": lambda: (
17+
os.environ.get("VLLM_FL_USE_FLAGGEMS_VLLM", "1").lower() in ("1", "true")
18+
),
1619
}
1720

1821

vllm_fl/ops/fused_moe/fused_moe_utils.py

Lines changed: 2 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -44,6 +44,7 @@
4444
from vllm.platforms import current_platform
4545
from vllm.triton_utils import tl
4646

47+
import vllm_fl.envs as fl_envs
4748
from vllm_fl.dispatch import CachedOp
4849
from vllm_fl.ops.fused_moe.activation import apply_moe_activation
4950
from vllm_fl.utils import use_flaggems
@@ -328,9 +329,7 @@ def apply(
328329
# Fast path (no LoRA, NVIDIA only): let FlagGems own both expert GEMMs
329330
# for unquantized and W8A16 inputs.
330331
if self._lora_context is None and current_platform.is_cuda():
331-
from vllm_fl.utils import use_flaggems_vllm
332-
333-
if use_flaggems_vllm():
332+
if fl_envs.VLLM_FL_USE_FLAGGEMS_VLLM:
334333
import flaggems_vllm
335334
fused_experts_impl = flaggems_vllm.fused_experts_impl
336335
else:

0 commit comments

Comments
 (0)