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:
+3
-27
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user