From cbc584470df3faed5976731aae857258891006ad Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Sat, 8 Aug 2026 12:32:51 +0800 Subject: [PATCH] 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 --- astrai/__init__.py | 30 +++--------------------------- astrai/inference/core/executor.py | 9 ++++----- astrai/logging.py | 27 +++++++++++++++++++++++++++ scripts/demo/stream_chat.py | 2 +- scripts/tools/benchmark.py | 2 -- scripts/tools/generate.py | 2 -- scripts/tools/preprocess.py | 2 -- scripts/tools/server.py | 2 -- scripts/tools/train.py | 2 -- 9 files changed, 35 insertions(+), 43 deletions(-) create mode 100644 astrai/logging.py diff --git a/astrai/__init__.py b/astrai/__init__.py index 0e559f7..de331b8 100644 --- a/astrai/__init__.py +++ b/astrai/__init__.py @@ -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() diff --git a/astrai/inference/core/executor.py b/astrai/inference/core/executor.py index 6705172..e0aba1d 100644 --- a/astrai/inference/core/executor.py +++ b/astrai/inference/core/executor.py @@ -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 diff --git a/astrai/logging.py b/astrai/logging.py new file mode 100644 index 0000000..d1fc8f1 --- /dev/null +++ b/astrai/logging.py @@ -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) diff --git a/scripts/demo/stream_chat.py b/scripts/demo/stream_chat.py index 3fdc563..e09d4bc 100644 --- a/scripts/demo/stream_chat.py +++ b/scripts/demo/stream_chat.py @@ -3,7 +3,7 @@ from pathlib import Path import torch -from astrai.inference import InferenceEngine +from astrai import InferenceEngine from astrai.model import AutoModel from astrai.tokenize import AutoTokenizer diff --git a/scripts/tools/benchmark.py b/scripts/tools/benchmark.py index 7623754..9a7dbc2 100644 --- a/scripts/tools/benchmark.py +++ b/scripts/tools/benchmark.py @@ -5,7 +5,6 @@ from typing import Optional, Union import click import torch -from astrai import setup_logging from astrai.config import BaseModelConfig, ConfigFactory from astrai.extension import ATTN_BACKEND, AttentionBackendFactory, attn_backend from astrai.inference.core.cache import PagePool @@ -478,5 +477,4 @@ def benchmark_command( if __name__ == "__main__": - setup_logging() benchmark_command() diff --git a/scripts/tools/generate.py b/scripts/tools/generate.py index 9263878..cd425ca 100644 --- a/scripts/tools/generate.py +++ b/scripts/tools/generate.py @@ -6,7 +6,6 @@ import click import torch from tqdm import tqdm -from astrai import setup_logging from astrai.inference import InferenceEngine from astrai.model import AutoModel from astrai.tokenize import AutoTokenizer @@ -157,5 +156,4 @@ def generate_command(**kwargs): if __name__ == "__main__": - setup_logging() generate_command() diff --git a/scripts/tools/preprocess.py b/scripts/tools/preprocess.py index 1b3b364..c1c18b2 100644 --- a/scripts/tools/preprocess.py +++ b/scripts/tools/preprocess.py @@ -2,7 +2,6 @@ import click -from astrai import setup_logging from astrai.config.preprocess_config import PipelineConfig from astrai.preprocessing.pipeline import Pipeline @@ -48,5 +47,4 @@ def preprocess_command(inputs, output_dir, pipeline_config, tokenizer_path, batc if __name__ == "__main__": - setup_logging() preprocess_command() diff --git a/scripts/tools/server.py b/scripts/tools/server.py index 31d43bf..66c62cb 100644 --- a/scripts/tools/server.py +++ b/scripts/tools/server.py @@ -3,7 +3,6 @@ from pathlib import Path import click import torch -from astrai import setup_logging from astrai.inference import run_server _DTYPES = ["bfloat16", "float16", "float32"] @@ -65,5 +64,4 @@ def server_command( if __name__ == "__main__": - setup_logging() server_command() diff --git a/scripts/tools/train.py b/scripts/tools/train.py index 452ca92..fd0402b 100644 --- a/scripts/tools/train.py +++ b/scripts/tools/train.py @@ -8,7 +8,6 @@ import torch from click.core import ParameterSource from torch import optim -from astrai import setup_logging from astrai.config import AutoRegressiveLMConfig, TrainConfig from astrai.dataset import DatasetFactory, dpo_collate_fn, grpo_collate_fn from astrai.model import AutoRegressiveLM @@ -828,5 +827,4 @@ def train( if __name__ == "__main__": - setup_logging() train_command()