Skip to content

Commit ba78fad

Browse files
committed
fix: avoid unused patch state globals
1 parent 851bbda commit ba78fad

2 files changed

Lines changed: 5 additions & 11 deletions

File tree

vllm_fl/dispatch/backends/vendor/gcu/impl/flash_attn_backend.py

Lines changed: 2 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -38,15 +38,12 @@
3838

3939
logger = logging.getLogger(__name__)
4040

41-
_patched = False
42-
4341
_FA_UTILS = "vllm.v1.attention.backends.fa_utils"
4442
_FLASH_ATTN = "vllm.v1.attention.backends.flash_attn"
4543

4644

4745
def apply_flash_attn_backend_gcu_patch() -> None:
48-
global _patched
49-
if _patched:
46+
if getattr(sys.modules.get(_FA_UTILS), "_gcu_flash_attn_patched", False):
5047
return
5148

5249
try:
@@ -72,6 +69,7 @@ def apply_flash_attn_backend_gcu_patch() -> None:
7269
fa_utils.get_scheduler_metadata = get_scheduler_metadata
7370
fa_utils.reshape_and_cache_flash = reshape_and_cache_flash
7471
fa_utils._GCU_FLASH_ATTN_AVAILABLE = True
72+
fa_utils._gcu_flash_attn_patched = True
7573
fa_utils.is_flash_attn_varlen_func_available = lambda: True
7674

7775
# If flash_attn.py already imported (its gate ran False and it skipped the
@@ -84,7 +82,6 @@ def apply_flash_attn_backend_gcu_patch() -> None:
8482
flash_attn_mod.flash_attn_supports_sinks = fa_utils.flash_attn_supports_sinks
8583
flash_attn_mod.is_flash_attn_varlen_func_available = lambda: True
8684

87-
_patched = True
8885
logger.info(
8986
"GCU: enabled native FLASH_ATTN backend "
9087
"(vendor flash_attn_varlen_func + flag_gems reshape_and_cache_flash)"

vllm_fl/dispatch/backends/vendor/gcu/impl/slot_mapping.py

Lines changed: 3 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -40,8 +40,6 @@
4040

4141
logger = logging.getLogger(__name__)
4242

43-
_patched = False
44-
4543

4644
def _compute_slot_mapping_int32(self, num_reqs, query_start_loc, positions):
4745
from vllm.v1.attention.backends.utils import PAD_SLOT_ID
@@ -105,11 +103,10 @@ def _compute_slot_mapping_int32(self, num_reqs, query_start_loc, positions):
105103

106104
def apply_slot_mapping_gcu_patch() -> None:
107105
"""Replace BlockTable.compute_slot_mapping with the on-device int32 version."""
108-
global _patched
109-
if _patched:
110-
return
111106
from vllm.v1.worker.block_table import BlockTable
112107

108+
if BlockTable.compute_slot_mapping is _compute_slot_mapping_int32:
109+
return
110+
113111
BlockTable.compute_slot_mapping = _compute_slot_mapping_int32
114-
_patched = True
115112
logger.info("GCU: patched BlockTable.compute_slot_mapping (on-device int32)")

0 commit comments

Comments
 (0)