Skip to content

Commit 511aab9

Browse files
authored
fix backend bugs (#2781)
1 parent 2e38d26 commit 511aab9

11 files changed

Lines changed: 25 additions & 26 deletions

File tree

src/flag_gems/runtime/backend/__init__.py

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -40,7 +40,6 @@ def __init__(self):
4040

4141
# Global singleton instance
4242
_state = BackendState()
43-
vendor_module = _state.vendor_module
4443

4544

4645
class BackendArchEvent:

src/flag_gems/runtime/backend/_cambricon/ops/cat.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -85,10 +85,10 @@ def gen_imports(self):
8585
self.writeline("import triton.language as tl")
8686
self.newline()
8787
self.writeline("from flag_gems.runtime import torch_device_fn")
88-
self.writeline("from flag_gems.runtime.backend import vendor_module")
88+
self.writeline("from flag_gems.runtime.backend import _state")
8989
self.writeline("from flag_gems.utils import libentry, libtuner")
9090
self.newline()
91-
self.writeline("TOTAL_CORE_NUM = vendor_module.TOTAL_CORE_NUM")
91+
self.writeline("TOTAL_CORE_NUM = _state.vendor_module.TOTAL_CORE_NUM")
9292
self.newline()
9393
self.newline()
9494

src/flag_gems/runtime/backend/_cambricon/ops/flip.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -81,9 +81,9 @@ def __imports(self):
8181
from triton import language as tl
8282
8383
from flag_gems.utils import libentry
84-
from flag_gems.runtime.backend import vendor_module
85-
TOTAL_CORE_NUM = vendor_module.utils.TOTAL_CORE_NUM
86-
MAX_NRAM_SIZE = vendor_module.utils.MAX_NRAM_SIZE
84+
from flag_gems.runtime.backend import _state
85+
TOTAL_CORE_NUM = _state.vendor_module.utils.TOTAL_CORE_NUM
86+
MAX_NRAM_SIZE = _state.vendor_module.utils.MAX_NRAM_SIZE
8787
8888
8989
"""

src/flag_gems/runtime/backend/_cambricon/ops/repeat.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -65,8 +65,8 @@ def generate_imports(code: IndentedBuffer) -> IndentedBuffer:
6565
code.writeline("from flag_gems.runtime import torch_device_fn")
6666
code.writeline("from flag_gems.utils.shape_utils import volume")
6767
code.writeline("from flag_gems.utils import libentry")
68-
code.writeline("from flag_gems.runtime.backend import vendor_module")
69-
code.writeline("MAX_GRID_SIZE_X = vendor_module.MAX_GRID_SIZE_X")
68+
code.writeline("from flag_gems.runtime.backend import _state")
69+
code.writeline("MAX_GRID_SIZE_X = _state.vendor_module.MAX_GRID_SIZE_X")
7070
code.writeline("from flag_gems.utils.type_utils import type_promotion")
7171
code.writeline("from flag_gems.utils import triton_lang_extension as tle")
7272
code.newline()

src/flag_gems/runtime/backend/_cambricon/ops/stack.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -52,9 +52,9 @@ def __imports(self):
5252
from triton import language as tl
5353
from typing import List, Tuple, Union
5454
from flag_gems.utils import libentry
55-
from flag_gems.runtime.backend import vendor_module
56-
TOTAL_CORE_NUM = vendor_module.TOTAL_CORE_NUM
57-
MAX_NRAM_SIZE = vendor_module.MAX_NRAM_SIZE
55+
from flag_gems.runtime.backend import _state
56+
TOTAL_CORE_NUM = _state.vendor_module.TOTAL_CORE_NUM
57+
MAX_NRAM_SIZE = _state.vendor_module.MAX_NRAM_SIZE
5858
5959
"""
6060
self.tpl(textwrap.dedent(tpl))

src/flag_gems/runtime/backend/_cambricon/ops/tile.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -65,8 +65,8 @@ def generate_imports(code: IndentedBuffer) -> IndentedBuffer:
6565
code.writeline("from flag_gems.runtime import torch_device_fn")
6666
code.writeline("from flag_gems.utils.shape_utils import volume")
6767
code.writeline("from flag_gems.utils import libentry, libtuner")
68-
code.writeline("from flag_gems.runtime.backend import vendor_module")
69-
code.writeline("MAX_GRID_SIZE_X = vendor_module.MAX_GRID_SIZE_X")
68+
code.writeline("from flag_gems.runtime.backend import _state")
69+
code.writeline("MAX_GRID_SIZE_X = _state.vendor_module.MAX_GRID_SIZE_X")
7070
code.writeline("from flag_gems.utils.type_utils import type_promotion")
7171
code.writeline("from flag_gems.utils import triton_lang_extension as tle")
7272
code.newline()

src/flag_gems/runtime/backend/_cambricon/ops/vstack.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -77,7 +77,7 @@ def __imports(self):
7777
from triton import language as tl
7878
from flag_gems.runtime import torch_device_fn
7979
from flag_gems.utils import libentry, libtuner
80-
from flag_gems.runtime.backend import vendor_module
80+
from flag_gems.runtime.backend import _state
8181
TOTAL_CORE_NUM = vendor_module.TOTAL_CORE_NUM
8282
MAX_NRAM_SIZE = vendor_module.MAX_NRAM_SIZE
8383
"""

src/flag_gems/runtime/backend/_enflame/gcu300/utils/codegen_config_utils.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,7 @@
55
import triton
66

77
from flag_gems.runtime import device
8-
from flag_gems.runtime.backend import vendor_module
8+
from flag_gems.runtime.backend import _state
99
from flag_gems.runtime.common import vendors
1010

1111
ENFLAME_GCU300_4SIPS = int(os.getenv("ENFLAME_GCU300_4SIPS", "0"))
@@ -67,12 +67,12 @@ def __post_init__(self):
6767
vendors.CAMBRICON: (
6868
CodeGenConfig(
6969
8192,
70-
tuple([vendor_module.TOTAL_CORE_NUM, 1, 1]),
70+
tuple([_state.vendor_module.TOTAL_CORE_NUM, 1, 1]),
7171
32,
7272
False,
7373
prefer_1d_tile=int(triton.__version__[0]) < 3,
7474
)
75-
if vendor_module.vendor_info.vendor_name == "cambricon"
75+
if _state.vendor_module.vendor_info.vendor_name == "cambricon"
7676
else None
7777
),
7878
vendors.METAX: CodeGenConfig(

src/flag_gems/runtime/backend/_enflame/gcu400/utils/codegen_config_utils.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,7 @@
44
import triton
55

66
from flag_gems.runtime import device
7-
from flag_gems.runtime.backend import vendor_module
7+
from flag_gems.runtime.backend import _state
88
from flag_gems.runtime.common import vendors
99

1010

@@ -61,12 +61,12 @@ def __post_init__(self):
6161
vendors.CAMBRICON: (
6262
CodeGenConfig(
6363
8192,
64-
tuple([vendor_module.TOTAL_CORE_NUM, 1, 1]),
64+
tuple([_state.vendor_module.TOTAL_CORE_NUM, 1, 1]),
6565
32,
6666
False,
6767
prefer_1d_tile=int(triton.__version__[0]) < 3,
6868
)
69-
if vendor_module.vendor_info.vendor_name == "cambricon"
69+
if _state.vendor_module.vendor_info.vendor_name == "cambricon"
7070
else None
7171
),
7272
vendors.METAX: CodeGenConfig(

src/flag_gems/runtime/backend/_kunlunxin/utils/codegen_config_utils.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,7 @@
44
import triton
55

66
from flag_gems.runtime import device
7-
from flag_gems.runtime.backend import vendor_module
7+
from flag_gems.runtime.backend import _state
88
from flag_gems.runtime.common import vendors
99

1010

@@ -65,12 +65,12 @@ def __post_init__(self):
6565
),
6666
vendors.CAMBRICON: CodeGenConfig(
6767
8192,
68-
tuple([vendor_module.TOTAL_CORE_NUM, 1, 1]),
68+
tuple([_state.vendor_module.TOTAL_CORE_NUM, 1, 1]),
6969
32,
7070
False,
7171
prefer_1d_tile=int(triton.__version__[0]) < 3,
7272
)
73-
if vendor_module.vendor_info.vendor_name == "cambricon"
73+
if _state.vendor_module.vendor_info.vendor_name == "cambricon"
7474
else None,
7575
vendors.METAX: CodeGenConfig(
7676
2048,

0 commit comments

Comments
 (0)