add support for eval configuration and fix logger problem

1. add support for configuration for gpu_collect and cache_dir which is used for cpu result gathering, configuration example

```json
"evaluation": {
    "gpu_collect":  false,
    "cache_dir": "path/to/your/local/cache"
}
```

2. fix logger file missing  when log_file is passed to get_logger and add log_file for trainer
3.  automatically create work_dir in rank0 worker
        Link: https://code.alibaba-inc.com/Ali-MaaS/MaaS-lib/codereview/11342068

    * add support for configuration for tmpdir and gpu_collect
This commit is contained in:
wenmeng.zwm
2023-01-09 02:51:35 +08:00
committed by yingda.chen
parent aa541468d1
commit 8f6a0f64e2
4 changed files with 42 additions and 7 deletions

View File

@@ -129,7 +129,6 @@ class EpochBasedTrainer(BaseTrainer):
# add default config
merge_cfg(self.cfg)
self.cfg = self.rebuild_config(self.cfg)
self.logger = get_logger(log_level=self.cfg.get('log_level', 'INFO'))
if 'cfg_options' in kwargs:
self.cfg.merge_from_dict(kwargs['cfg_options'])
@@ -147,8 +146,17 @@ class EpochBasedTrainer(BaseTrainer):
preprocessor)
self._dist = self.init_dist(kwargs.get('launcher'))
if is_master() and not os.path.exists(self.work_dir):
os.makedirs(self.work_dir)
self.device = self.get_device(kwargs.get('device'))
# init logger after distribution init
log_file = os.path.join(self.work_dir, '{}.log'.format(self.timestamp))
self.logger = get_logger(
log_file=log_file, log_level=self.cfg.get('log_level', 'INFO'))
self.train_dataset = self.to_task_dataset(
train_dataset,
mode=ModeKeys.TRAIN,
@@ -949,8 +957,8 @@ class EpochBasedTrainer(BaseTrainer):
self,
data_loader,
device=self.device,
tmpdir=None,
gpu_collect=False,
tmpdir=self.cfg.evaluation.get('cache_dir', None),
gpu_collect=self.cfg.evaluation.get('gpu_collect', False),
data_loader_iters_per_gpu=self._eval_iters_per_epoch)
else:
from modelscope.trainers.utils.inference import single_gpu_test

View File

@@ -6,6 +6,9 @@ from typing import Optional
init_loggers = {}
formatter = logging.Formatter(
'%(asctime)s - %(name)s - %(levelname)s - %(message)s')
def get_logger(log_file: Optional[str] = None,
log_level: int = logging.INFO,
@@ -19,10 +22,12 @@ def get_logger(log_file: Optional[str] = None,
file_mode: Specifies the mode to open the file, if filename is
specified (if filemode is unspecified, it defaults to 'w').
"""
logger_name = __name__.split('.')[0]
logger = logging.getLogger(logger_name)
if logger_name in init_loggers:
add_file_handler_if_needed(logger, log_file, file_mode, log_level)
return logger
# handle duplicate logs to the console
@@ -49,8 +54,6 @@ def get_logger(log_file: Optional[str] = None,
file_handler = logging.FileHandler(log_file, file_mode)
handlers.append(file_handler)
formatter = logging.Formatter(
'%(asctime)s - %(name)s - %(levelname)s - %(message)s')
for handler in handlers:
handler.setFormatter(formatter)
handler.setLevel(log_level)
@@ -64,3 +67,21 @@ def get_logger(log_file: Optional[str] = None,
init_loggers[logger_name] = True
return logger
def add_file_handler_if_needed(logger, log_file, file_mode, log_level):
for handler in logger.handlers:
if isinstance(handler, logging.FileHandler):
return
if importlib.util.find_spec('torch') is not None:
from modelscope.utils.torch_utils import is_master
is_worker0 = is_master()
else:
is_worker0 = True
if is_worker0 and log_file is not None:
file_handler = logging.FileHandler(log_file, file_mode)
file_handler.setFormatter(formatter)
file_handler.setLevel(log_level)
logger.addHandler(file_handler)

View File

@@ -10,6 +10,7 @@ from modelscope.metainfo import Trainers
from modelscope.msdatasets import MsDataset
from modelscope.trainers import build_trainer
from modelscope.utils.audio.audio_utils import to_segment
from modelscope.utils.constant import DownloadMode
from modelscope.utils.hub import read_config
from modelscope.utils.test_utils import test_level
@@ -31,7 +32,9 @@ class TestANSTrainer(unittest.TestCase):
cfg.dump(self.cfg_file)
hf_ds = MsDataset.load(
'ICASSP_2021_DNS_Challenge', split='test').to_hf_dataset()
'ICASSP_2021_DNS_Challenge',
split='test',
download_mode=DownloadMode.FORCE_REDOWNLOAD).to_hf_dataset()
mapped_ds = hf_ds.map(
partial(to_segment, segment_length=SEGMENT_LENGTH_TEST),
remove_columns=['duration'],

View File

@@ -141,7 +141,6 @@ class TrainerTest(unittest.TestCase):
config_path = os.path.join(self.tmp_dir, ModelFile.CONFIGURATION)
with open(config_path, 'w') as f:
json.dump(json_cfg, f)
trainer_name = Trainers.default
kwargs = dict(
cfg_file=config_path,
@@ -157,6 +156,10 @@ class TrainerTest(unittest.TestCase):
results_files = os.listdir(self.tmp_dir)
self.assertIn(f'{trainer.timestamp}.log.json', results_files)
with open(f'{self.tmp_dir}/{trainer.timestamp}.log', 'r') as infile:
lines = infile.readlines()
self.assertTrue(len(lines) > 20)
self.assertIn(f'{trainer.timestamp}.log', results_files)
self.assertIn(f'{LogKeys.EPOCH}_1.pth', results_files)
self.assertIn(f'{LogKeys.EPOCH}_2.pth', results_files)
self.assertIn(f'{LogKeys.EPOCH}_3.pth', results_files)