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