Skip to content

Commit 2b046fa

Browse files
committed
KUNLUNXIN][BUF-FIX]fix that transformer engine disable benchmark logger
the transformer_engine may disable flagGems benchMark logging, which lead to "pytest --record log" failure
1 parent a1d4c62 commit 2b046fa

3 files changed

Lines changed: 25 additions & 10 deletions

File tree

benchmark/conftest.py

Lines changed: 17 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -22,6 +22,8 @@
2222

2323
device = flag_gems.device
2424
vendor_name = flag_gems.vendor_name
25+
recordLogger = logging.getLogger("flag_gems.benchmark.record")
26+
recordLogger.propagate = False
2527

2628

2729
class BenchConfig:
@@ -165,12 +167,21 @@ def pytest_configure(config):
165167
for arg in config.invocation_params.args
166168
]
167169

168-
logging.basicConfig(
169-
filename="result_{}.log".format("_".join(cmd_args)).replace("_-", "-"),
170-
filemode="w",
171-
level=logging.INFO,
172-
format="[%(levelname)s] %(message)s",
173-
)
170+
log_file = "result_{}.log".format("_".join(cmd_args)).replace("_-", "-")
171+
172+
for h in list(recordLogger.handlers):
173+
recordLogger.removeHandler(h)
174+
try:
175+
h.close()
176+
except Exception:
177+
pass
178+
179+
handler = logging.FileHandler(log_file, mode="w", encoding="utf-8")
180+
handler.setLevel(logging.INFO)
181+
handler.setFormatter(logging.Formatter("[%(levelname)s] %(message)s"))
182+
recordLogger.addHandler(handler)
183+
recordLogger.setLevel(logging.INFO)
184+
recordLogger.info("Benchmark record logger enabled")
174185

175186

176187
BUILTIN_MARKS = {

benchmark/performance_utils.py

Lines changed: 3 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,5 @@
11
import gc
22
import importlib
3-
import logging
43
import os
54
import time
65
from typing import Any, Generator, List, Optional, Tuple
@@ -26,7 +25,7 @@
2625
OperationAttribute,
2726
check_metric_dependencies,
2827
)
29-
from .conftest import Config
28+
from .conftest import Config, recordLogger
3029

3130
torch_backend_device = flag_gems.runtime.torch_backend_device
3231
torch_device_fn = flag_gems.runtime.torch_device_fn
@@ -372,7 +371,7 @@ def run(self):
372371
shape_desc=self.shape_desc,
373372
)
374373
print(attri)
375-
logging.info(attri.to_dict())
374+
recordLogger.info(attri.to_dict())
376375
return
377376
self.init_user_config()
378377
for dtype in self.to_bench_dtypes:
@@ -425,7 +424,7 @@ def run(self):
425424
result=metrics,
426425
)
427426
print(result)
428-
logging.info(result.to_json())
427+
recordLogger.info(result.to_json())
429428

430429

431430
class GenericBenchmark(Benchmark):

benchmark/test_transformer_engine_perf.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,11 @@
88
try:
99
from transformer_engine.pytorch import cpp_extensions as tex
1010

11+
# Note: Importing transformer_engine (especially in some versions like on python 3.10) may automatically
12+
# configure the Root Logger (adding handlers). This can cause subsequent `logging.basicConfig` calls
13+
# (used by FlagGems benchmark) to be ignored/no-op, leading to missing result log files.
14+
# See: https://github.com/NVIDIA/TransformerEngine/issues/1065
15+
1116
TE_AVAILABLE = True
1217
except ImportError:
1318
TE_AVAILABLE = False

0 commit comments

Comments
 (0)