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" __version__ = "1.3.13"
__author__ = "ViperEkura" __author__ = "ViperEkura"
import logging
import os
from astrai.config import ( from astrai.config import (
AutoRegressiveLMConfig, AutoRegressiveLMConfig,
BaseModelConfig, BaseModelConfig,
@@ -28,6 +25,7 @@ from astrai.inference import (
run_server, run_server,
sample, sample,
) )
from astrai.logging import setup_logging
from astrai.model import ( from astrai.model import (
AutoModel, AutoModel,
AutoRegressiveLM, AutoRegressiveLM,
@@ -55,30 +53,6 @@ from astrai.trainer import (
Trainer, 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__ = [ __all__ = [
"AutoRegressiveLM", "AutoRegressiveLM",
"AutoRegressiveLMConfig", "AutoRegressiveLMConfig",
@@ -122,3 +96,5 @@ __all__ = [
"setup_logging", "setup_logging",
"spawn_parallel_fn", "spawn_parallel_fn",
] ]
setup_logging()
+4 -5
View File
@@ -1,5 +1,4 @@
import logging import logging
import os
import time import time
from contextlib import contextmanager from contextlib import contextmanager
from dataclasses import dataclass from dataclasses import dataclass
@@ -23,19 +22,19 @@ from astrai.model.automodel import AutoModel
from astrai.tokenize.tokenizer import AutoTokenizer from astrai.tokenize.tokenizer import AutoTokenizer
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
_TIMED = os.environ.get("ASTRAI_TIMED", "") == "1"
@contextmanager @contextmanager
def timed(label: str, log: Optional[logging.Logger] = None): def timed(label: str, log: Optional[logging.Logger] = None):
"""Wall-clock debug timer, enabled via ``ASTRAI_TIMED=1``.""" """Wall-clock debug timer, enabled when the logger level is DEBUG or lower."""
if not _TIMED: log = log or logger
if not log.isEnabledFor(logging.DEBUG):
yield yield
return return
tic = time.perf_counter() tic = time.perf_counter()
yield yield
elapsed_ms = (time.perf_counter() - tic) * 1000 elapsed_ms = (time.perf_counter() - tic) * 1000
(log or logger).info("%s %.1fms", label, elapsed_ms) log.debug("%s %.1fms", label, elapsed_ms)
@dataclass @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)
+1 -1
View File
@@ -3,7 +3,7 @@ from pathlib import Path
import torch import torch
from astrai.inference import InferenceEngine from astrai import InferenceEngine
from astrai.model import AutoModel from astrai.model import AutoModel
from astrai.tokenize import AutoTokenizer from astrai.tokenize import AutoTokenizer
-2
View File
@@ -5,7 +5,6 @@ from typing import Optional, Union
import click import click
import torch import torch
from astrai import setup_logging
from astrai.config import BaseModelConfig, ConfigFactory from astrai.config import BaseModelConfig, ConfigFactory
from astrai.extension import ATTN_BACKEND, AttentionBackendFactory, attn_backend from astrai.extension import ATTN_BACKEND, AttentionBackendFactory, attn_backend
from astrai.inference.core.cache import PagePool from astrai.inference.core.cache import PagePool
@@ -478,5 +477,4 @@ def benchmark_command(
if __name__ == "__main__": if __name__ == "__main__":
setup_logging()
benchmark_command() benchmark_command()
-2
View File
@@ -6,7 +6,6 @@ import click
import torch import torch
from tqdm import tqdm from tqdm import tqdm
from astrai import setup_logging
from astrai.inference import InferenceEngine from astrai.inference import InferenceEngine
from astrai.model import AutoModel from astrai.model import AutoModel
from astrai.tokenize import AutoTokenizer from astrai.tokenize import AutoTokenizer
@@ -157,5 +156,4 @@ def generate_command(**kwargs):
if __name__ == "__main__": if __name__ == "__main__":
setup_logging()
generate_command() generate_command()
-2
View File
@@ -2,7 +2,6 @@
import click import click
from astrai import setup_logging
from astrai.config.preprocess_config import PipelineConfig from astrai.config.preprocess_config import PipelineConfig
from astrai.preprocessing.pipeline import Pipeline from astrai.preprocessing.pipeline import Pipeline
@@ -48,5 +47,4 @@ def preprocess_command(inputs, output_dir, pipeline_config, tokenizer_path, batc
if __name__ == "__main__": if __name__ == "__main__":
setup_logging()
preprocess_command() preprocess_command()
-2
View File
@@ -3,7 +3,6 @@ from pathlib import Path
import click import click
import torch import torch
from astrai import setup_logging
from astrai.inference import run_server from astrai.inference import run_server
_DTYPES = ["bfloat16", "float16", "float32"] _DTYPES = ["bfloat16", "float16", "float32"]
@@ -65,5 +64,4 @@ def server_command(
if __name__ == "__main__": if __name__ == "__main__":
setup_logging()
server_command() server_command()
-2
View File
@@ -8,7 +8,6 @@ import torch
from click.core import ParameterSource from click.core import ParameterSource
from torch import optim from torch import optim
from astrai import setup_logging
from astrai.config import AutoRegressiveLMConfig, TrainConfig from astrai.config import AutoRegressiveLMConfig, TrainConfig
from astrai.dataset import DatasetFactory, dpo_collate_fn, grpo_collate_fn from astrai.dataset import DatasetFactory, dpo_collate_fn, grpo_collate_fn
from astrai.model import AutoRegressiveLM from astrai.model import AutoRegressiveLM
@@ -828,5 +827,4 @@ def train(
if __name__ == "__main__": if __name__ == "__main__":
setup_logging()
train_command() train_command()