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"
|
__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()
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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)
|
||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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,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()
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
Reference in New Issue
Block a user