feat : add setup_logging with hierarchical astrai logger

- setup_logging(): attach handler only to astrai logger, not root
- all astrai.* sub-module loggers inherit automatically
- controlled by ASTR_LOG_LEVEL env var (default INFO)
- called in if __name__ == '__main__' of each CLI script
This commit is contained in:
2026-07-27 08:13:48 +08:00
parent 53c804e233
commit 07625057f2
6 changed files with 38 additions and 0 deletions
+28
View File
@@ -1,6 +1,33 @@
__version__ = "1.3.11" __version__ = "1.3.11"
__author__ = "ViperEkura" __author__ = "ViperEkura"
import logging
import os
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)
from astrai.config import ( from astrai.config import (
AutoRegressiveLMConfig, AutoRegressiveLMConfig,
BaseModelConfig, BaseModelConfig,
@@ -94,5 +121,6 @@ __all__ = [
"only_on_rank", "only_on_rank",
"run_server", "run_server",
"sample", "sample",
"setup_logging",
"spawn_parallel_fn", "spawn_parallel_fn",
] ]
+2
View File
@@ -1,6 +1,7 @@
import click import click
import torch import torch
from astrai import setup_logging
from astrai.config import AutoRegressiveLMConfig from astrai.config import AutoRegressiveLMConfig
_DTYPES = ["bfloat16", "float16", "float32"] _DTYPES = ["bfloat16", "float16", "float32"]
@@ -204,4 +205,5 @@ def benchmark_command(
if __name__ == "__main__": if __name__ == "__main__":
setup_logging()
benchmark_command() benchmark_command()
+2
View File
@@ -5,6 +5,7 @@ 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
@@ -162,4 +163,5 @@ def generate_command(**kwargs):
if __name__ == "__main__": if __name__ == "__main__":
setup_logging()
generate_command() generate_command()
+2
View File
@@ -2,6 +2,7 @@
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
@@ -47,4 +48,5 @@ 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,6 +3,7 @@ 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"]
@@ -64,4 +65,5 @@ def server_command(
if __name__ == "__main__": if __name__ == "__main__":
setup_logging()
server_command() server_command()
+2
View File
@@ -7,6 +7,7 @@ import click
import torch import torch
from torch import Tensor, nn, optim from torch import Tensor, nn, 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
@@ -556,4 +557,5 @@ def train(
if __name__ == "__main__": if __name__ == "__main__":
setup_logging()
train_command() train_command()