refactor: centralize logging in astrai.logging, replace ASTRAI_TIMED with log level

- move setup_logging to astrai/logging.py
- timed() now uses logger.isEnabledFor(DEBUG) instead of separate env var
- enable ASTR_LOG_LEVEL=DEBUG to see per-step timing logs
- call setup_logging() in stream_chat.py
This commit is contained in:
2026-08-08 12:39:27 +08:00
parent cb60713a72
commit cbc584470d
9 changed files with 35 additions and 43 deletions
+3 -27
View File
@@ -1,9 +1,6 @@
__version__ = "1.3.13"
__author__ = "ViperEkura"
import logging
import os
from astrai.config import (
AutoRegressiveLMConfig,
BaseModelConfig,
@@ -28,6 +25,7 @@ from astrai.inference import (
run_server,
sample,
)
from astrai.logging import setup_logging
from astrai.model import (
AutoModel,
AutoRegressiveLM,
@@ -55,30 +53,6 @@ from astrai.trainer import (
Trainer,
)
def setup_logging(level: str = "INFO"):
"""Attach a handler to the ``astrai`` logger (only, not root).
Call once per process, e.g. at the top of CLI scripts.
Set ``ASTR_LOG_LEVEL`` to override the default ``INFO``.
"""
_logger = logging.getLogger("astrai")
if _logger.handlers:
return
_level = getattr(
logging, os.environ.get("ASTR_LOG_LEVEL", level).upper(), logging.INFO
)
_logger.setLevel(_level)
_handler = logging.StreamHandler()
_handler.setFormatter(
logging.Formatter(
"%(asctime)s | %(levelname)-7s | %(name)s | %(message)s",
datefmt="%Y-%m-%d %H:%M:%S",
)
)
_logger.addHandler(_handler)
__all__ = [
"AutoRegressiveLM",
"AutoRegressiveLMConfig",
@@ -122,3 +96,5 @@ __all__ = [
"setup_logging",
"spawn_parallel_fn",
]
setup_logging()
+4 -5
View File
@@ -1,5 +1,4 @@
import logging
import os
import time
from contextlib import contextmanager
from dataclasses import dataclass
@@ -23,19 +22,19 @@ from astrai.model.automodel import AutoModel
from astrai.tokenize.tokenizer import AutoTokenizer
logger = logging.getLogger(__name__)
_TIMED = os.environ.get("ASTRAI_TIMED", "") == "1"
@contextmanager
def timed(label: str, log: Optional[logging.Logger] = None):
"""Wall-clock debug timer, enabled via ``ASTRAI_TIMED=1``."""
if not _TIMED:
"""Wall-clock debug timer, enabled when the logger level is DEBUG or lower."""
log = log or logger
if not log.isEnabledFor(logging.DEBUG):
yield
return
tic = time.perf_counter()
yield
elapsed_ms = (time.perf_counter() - tic) * 1000
(log or logger).info("%s %.1fms", label, elapsed_ms)
log.debug("%s %.1fms", label, elapsed_ms)
@dataclass
+27
View File
@@ -0,0 +1,27 @@
import logging
import os
def setup_logging(level: str = "INFO"):
"""Attach a StreamHandler to the ``astrai`` logger (idempotent).
Call once per process at the top of CLI scripts.
Set ``ASTR_LOG_LEVEL`` env var to override the default level.
Level names: ``DEBUG``, ``INFO``, ``WARNING``, ``ERROR``, ``CRITICAL``.
``DEBUG`` enables per-step prefill/decode timing logs
(:func:`astrai.inference.core.executor.timed`).
"""
logger = logging.getLogger("astrai")
if logger.handlers:
return
level_name = os.environ.get("ASTR_LOG_LEVEL", level).upper()
logger.setLevel(getattr(logging, level_name, logging.INFO))
handler = logging.StreamHandler()
handler.setFormatter(
logging.Formatter(
"%(asctime)s | %(levelname)-7s | %(name)s | %(message)s",
datefmt="%Y-%m-%d %H:%M:%S",
)
)
logger.addHandler(handler)