This commit is contained in:
2
ppocr/utils/loggers/__init__.py
Normal file
2
ppocr/utils/loggers/__init__.py
Normal file
@@ -0,0 +1,2 @@
|
||||
from .wandb_logger import WandbLogger
|
||||
from .loggers import Loggers
|
||||
16
ppocr/utils/loggers/base_logger.py
Normal file
16
ppocr/utils/loggers/base_logger.py
Normal file
@@ -0,0 +1,16 @@
|
||||
import os
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
|
||||
class BaseLogger(ABC):
|
||||
def __init__(self, save_dir):
|
||||
self.save_dir = save_dir
|
||||
os.makedirs(self.save_dir, exist_ok=True)
|
||||
|
||||
@abstractmethod
|
||||
def log_metrics(self, metrics, prefix=None):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def close(self):
|
||||
pass
|
||||
19
ppocr/utils/loggers/loggers.py
Normal file
19
ppocr/utils/loggers/loggers.py
Normal file
@@ -0,0 +1,19 @@
|
||||
from .wandb_logger import WandbLogger
|
||||
|
||||
|
||||
class Loggers(object):
|
||||
def __init__(self, loggers):
|
||||
super().__init__()
|
||||
self.loggers = loggers
|
||||
|
||||
def log_metrics(self, metrics, prefix=None, step=None):
|
||||
for logger in self.loggers:
|
||||
logger.log_metrics(metrics, prefix=prefix, step=step)
|
||||
|
||||
def log_model(self, is_best, prefix, metadata=None):
|
||||
for logger in self.loggers:
|
||||
logger.log_model(is_best=is_best, prefix=prefix, metadata=metadata)
|
||||
|
||||
def close(self):
|
||||
for logger in self.loggers:
|
||||
logger.close()
|
||||
84
ppocr/utils/loggers/wandb_logger.py
Normal file
84
ppocr/utils/loggers/wandb_logger.py
Normal file
@@ -0,0 +1,84 @@
|
||||
import os
|
||||
from .base_logger import BaseLogger
|
||||
from ppocr.utils.logging import get_logger
|
||||
|
||||
|
||||
class WandbLogger(BaseLogger):
|
||||
def __init__(
|
||||
self,
|
||||
project=None,
|
||||
name=None,
|
||||
id=None,
|
||||
entity=None,
|
||||
save_dir=None,
|
||||
config=None,
|
||||
**kwargs,
|
||||
):
|
||||
try:
|
||||
import wandb
|
||||
|
||||
self.wandb = wandb
|
||||
except ModuleNotFoundError:
|
||||
raise ModuleNotFoundError("Please install wandb using `pip install wandb`")
|
||||
|
||||
self.project = project
|
||||
self.name = name
|
||||
self.id = id
|
||||
self.save_dir = save_dir
|
||||
self.config = config
|
||||
self.kwargs = kwargs
|
||||
self.entity = entity
|
||||
self._run = None
|
||||
self._wandb_init = dict(
|
||||
project=self.project,
|
||||
name=self.name,
|
||||
id=self.id,
|
||||
entity=self.entity,
|
||||
dir=self.save_dir,
|
||||
resume="allow",
|
||||
)
|
||||
self._wandb_init.update(**kwargs)
|
||||
self.logger = get_logger()
|
||||
|
||||
_ = self.run
|
||||
|
||||
if self.config:
|
||||
self.run.config.update(self.config)
|
||||
|
||||
@property
|
||||
def run(self):
|
||||
if self._run is None:
|
||||
if self.wandb.run is not None:
|
||||
self.logger.info(
|
||||
"There is a wandb run already in progress "
|
||||
"and newly created instances of `WandbLogger` will reuse"
|
||||
" this run. If this is not desired, call `wandb.finish()`"
|
||||
"before instantiating `WandbLogger`."
|
||||
)
|
||||
self._run = self.wandb.run
|
||||
else:
|
||||
self._run = self.wandb.init(**self._wandb_init)
|
||||
return self._run
|
||||
|
||||
def log_metrics(self, metrics, prefix=None, step=None):
|
||||
if not prefix:
|
||||
prefix = ""
|
||||
updated_metrics = {prefix.lower() + "/" + k: v for k, v in metrics.items()}
|
||||
|
||||
self.run.log(updated_metrics, step=step)
|
||||
|
||||
def log_model(self, is_best, prefix, metadata=None):
|
||||
model_path = os.path.join(self.save_dir, prefix + ".pdparams")
|
||||
artifact = self.wandb.Artifact(
|
||||
"model-{}".format(self.run.id), type="model", metadata=metadata
|
||||
)
|
||||
artifact.add_file(model_path, name="model_ckpt.pdparams")
|
||||
|
||||
aliases = [prefix]
|
||||
if is_best:
|
||||
aliases.append("best")
|
||||
|
||||
self.run.log_artifact(artifact, aliases=aliases)
|
||||
|
||||
def close(self):
|
||||
self.run.finish()
|
||||
Reference in New Issue
Block a user