add training args support and image classification fintune example

design doc: https://yuque.antfin.com/pai/rwqgvl/khy4uw5dgi39s6ke

usage:
```python

    from modelscope.trainers.training_args import (ArgAttr, MSArgumentParser,
                                               training_args)


    training_args.topk = ArgAttr(cfg_node_name=['train.evaluation.metric_options.topk',
                                                'evaluation.metric_options.topk'],
                                 default=(1,), help='evaluation using topk, tuple format, eg (1,), (1,5)')
    training_args.train_data = ArgAttr(type=str, default='tany0699/cats_and_dogs', help='train dataset')
    training_args.validation_data = ArgAttr(type=str, default='tany0699/cats_and_dogs', help='validation dataset')
    training_args.model_id = ArgAttr(type=str, default='damo/cv_vit-base_image-classification_ImageNet-labels', help='model name')

    parser = MSArgumentParser(training_args)
    cfg_dict = parser.get_cfg_dict()
    args = parser.args
    
    train_dataset = create_dataset(args.train_data, split='train')
    val_dataset = create_dataset(args.validation_data, split='validation')

    def cfg_modify_fn(cfg):
        cfg.merge_from_dict(cfg_dict)
        return cfg

    kwargs = dict(
        model=args.model_id,          # model id
        train_dataset=train_dataset,  # training dataset
        eval_dataset=val_dataset,     # validation dataset
        cfg_modify_fn=cfg_modify_fn     # callback to modify configuration
    )

    trainer = build_trainer(name=Trainers.image_classification, default_args=kwargs)
    # start to train
    trainer.train()
```
        Link: https://code.alibaba-inc.com/Ali-MaaS/MaaS-lib/codereview/11225071
This commit is contained in:
wenmeng.zwm
2022-12-30 07:35:15 +08:00
committed by yingda.chen
parent 1e4a3dcbe5
commit b8ec677739
4 changed files with 421 additions and 0 deletions

View File

@@ -0,0 +1,67 @@
import os
from modelscope.metainfo import Trainers
from modelscope.msdatasets.ms_dataset import MsDataset
from modelscope.trainers.builder import build_trainer
from modelscope.trainers.training_args import ArgAttr, CliArgumentParser, training_args
def define_parser():
training_args.num_classes = ArgAttr(cfg_node_name=['model.mm_model.head.num_classes',
'model.mm_model.train_cfg.augments.0.num_classes',
'model.mm_model.train_cfg.augments.1.num_classes'],
type=int, help='number of classes')
training_args.train_batch_size.default = 16
training_args.train_data_worker.default = 1
training_args.max_epochs.default = 1
training_args.optimizer.default = 'AdamW'
training_args.lr.default = 1e-4
training_args.warmup_iters = ArgAttr('train.lr_config.warmup_iters', type=int, default=1, help='number of warmup epochs')
training_args.topk = ArgAttr(cfg_node_name=['train.evaluation.metric_options.topk',
'evaluation.metric_options.topk'],
default=(1,), help='evaluation using topk, tuple format, eg (1,), (1,5)')
training_args.train_data = ArgAttr(type=str, default='tany0699/cats_and_dogs', help='train dataset')
training_args.validation_data = ArgAttr(type=str, default='tany0699/cats_and_dogs', help='validation dataset')
training_args.model_id = ArgAttr(type=str, default='damo/cv_vit-base_image-classification_ImageNet-labels', help='model name')
parser = CliArgumentParser(training_args)
return parser
def create_dataset(name, split):
namespace, dataset_name = name.split('/')
return MsDataset.load(dataset_name, namespace=namespace,
subset_name='default',
split=split)
def train(parser):
cfg_dict = parser.get_cfg_dict()
args = parser.args
train_dataset = create_dataset(args.train_data, split='train')
val_dataset = create_dataset(args.validation_data, split='validation')
def cfg_modify_fn(cfg):
cfg.merge_from_dict(cfg_dict)
return cfg
kwargs = dict(
model=args.model_id, # model id
train_dataset=train_dataset, # training dataset
eval_dataset=val_dataset, # validation dataset
cfg_modify_fn=cfg_modify_fn # callback to modify configuration
)
# in distributed training, specify pytorch launcher
if 'MASTER_ADDR' in os.environ:
kwargs['launcher'] = 'pytorch'
trainer = build_trainer(name=Trainers.image_classification, default_args=kwargs)
# start to train
trainer.train()
if __name__ == '__main__':
parser = define_parser()
train(parser)

View File

@@ -0,0 +1,5 @@
PYTHONPATH=. python -m torch.distributed.launch --nproc_per_node=2 \
examples/pytorch/finetune_image_classification.py \
--num_classes 2 \
--train_data 'tany0699/cats_and_dogs' \
--validation_data 'tany0699/cats_and_dogs'

View File

@@ -0,0 +1,270 @@
# Copyright (c) Alibaba, Inc. and its affiliates.
import dataclasses
from argparse import Action, ArgumentDefaultsHelpFormatter, ArgumentParser
from typing import Any, Dict, List, Union
from addict import Dict as Adict
@dataclasses.dataclass
class ArgAttr():
""" Attributes for each arg
Args:
cfg_node_name (str or list[str]): if set empty, it means a normal arg for argparse, otherwise it means
this arg value correspond to those nodes in configuration file, and will replace them for training.
default: default value for current argument.
type: type for current argument.
choices (list of str): choices of value for this argument.
help (str): help str for this argument.
Examples:
```python
# define argument train_batch_size which corresponds to train.dataloader.batch_size_per_gpu
training_args = Adict(
train_batch_size=ArgAttr(
'train.dataloader.batch_size_per_gpu',
default=16,
type=int,
help='training batch size')
)
# num_classes which will modify three places in configuration
training_args = Adict(
num_classes = ArgAttr(
['model.mm_model.head.num_classes',
'model.mm_model.train_cfg.augments.0.num_classes',
'model.mm_model.train_cfg.augments.1.num_classes'],
type=int,
help='number of classes')
)
```
# a normal argument which has no relation with configuration
training_args = Adict(
local_rank = ArgAttr(
'',
default=1,
type=int,
help='local rank for current training process')
)
"""
cfg_node_name: Union[str, List[str]] = ''
default: Any = None
type: type = None
choices: List[str] = None
help: str = ''
training_args = Adict(
train_batch_size=ArgAttr(
'train.dataloader.batch_size_per_gpu',
default=16,
type=int,
help='training batch size'),
train_data_worker=ArgAttr(
'train.dataloader.workers_per_gpu',
default=8,
type=int,
help='number of data worker used for training'),
eval_batch_size=ArgAttr(
'evaluation.dataloader.batch_size_per_gpu',
default=16,
type=int,
help='training batch size'),
max_epochs=ArgAttr(
'train.max_epochs',
default=10,
type=int,
help='max number of training epoch'),
work_dir=ArgAttr(
'train.work_dir',
default='./work_dir',
type=str,
help='training directory to save models and training logs'),
lr=ArgAttr(
'train.optimizer.lr',
default=0.001,
type=float,
help='initial learning rate'),
optimizer=ArgAttr(
'train.optimizer.type',
default='SGD',
type=str,
choices=[
'Adadelta', 'Adagrad', 'Adam', 'AdamW', 'Adamax', 'ASGD',
'RMSprop', 'Rprop'
'SGD'
],
help='optimizer type'),
local_rank=ArgAttr(
'', default=0, type=int, help='local rank for this process'))
class CliArgumentParser(ArgumentParser):
""" Argument Parser to define and parse command-line args for training.
Args:
arg_dict (dict of `ArgAttr` or list of them): dict or list of dict which defines different
paramters for training.
"""
def __init__(self, arg_dict: Union[Dict[str, ArgAttr],
List[Dict[str, ArgAttr]]], **kwargs):
if 'formatter_class' not in kwargs:
kwargs['formatter_class'] = ArgumentDefaultsHelpFormatter
super().__init__(**kwargs)
self.arg_dict = arg_dict if isinstance(
arg_dict, Dict) else self._join_args(arg_dict)
self.define_args()
def _join_args(self, arg_dict_list: List[Dict[str, ArgAttr]]):
total_args = arg_dict_list[0].copy()
for args in arg_dict_list[1:]:
total_args.update(args)
return total_args
def define_args(self):
for arg_name, arg_attr in self.arg_dict.items():
name = f'--{arg_name}'
kwargs = dict(type=arg_attr.type, help=arg_attr.help)
if arg_attr.default is not None:
kwargs['default'] = arg_attr.default
else:
kwargs['required'] = True
if arg_attr.choices is not None:
kwargs['choices'] = arg_attr.choices
kwargs['action'] = SingleAction
self.add_argument(name, **kwargs)
def get_cfg_dict(self, args=None):
"""
Args:
args (default None):
List of strings to parse. The default is taken from sys.argv. (same as argparse.ArgumentParser)
Returns:
cfg_dict (dict of config): each key is a config node name such as 'train.max_epochs', this cfg_dict
should be used with function `cfg.merge_from_dict` to update config object.
"""
self.args, remainning = self.parse_known_args(args)
args_dict = vars(self.args)
cfg_dict = {}
for k, v in args_dict.items():
if k not in self.arg_dict or self.arg_dict[k].cfg_node_name == '':
continue
cfg_node = self.arg_dict[k].cfg_node_name
if isinstance(cfg_node, list):
for node in cfg_node:
cfg_dict[node] = v
else:
cfg_dict[cfg_node] = v
return cfg_dict
class DictAction(Action):
"""
argparse action to split an argument into KEY=VALUE form
on the first = and append to a dictionary. List options can
be passed as comma separated values, i.e 'KEY=V1,V2,V3', or with explicit
brackets, i.e. 'KEY=[V1,V2,V3]'. It also support nested brackets to build
list/tuple values. e.g. 'KEY=[(V1,V2),(V3,V4)]'
"""
@staticmethod
def parse_int_float_bool_str(val):
try:
return int(val)
except ValueError:
pass
try:
return float(val)
except ValueError:
pass
if val.lower() in ['true', 'false']:
return val.lower() == 'true'
if val == 'None':
return None
return val
@staticmethod
def parse_iterable(val):
"""Parse iterable values in the string.
All elements inside '()' or '[]' are treated as iterable values.
Args:
val (str): Value string.
Returns:
list | tuple: The expanded list or tuple from the string.
Examples:
>>> DictAction._parse_iterable('1,2,3')
[1, 2, 3]
>>> DictAction._parse_iterable('[a, b, c]')
['a', 'b', 'c']
>>> DictAction._parse_iterable('[(1, 2, 3), [a, b], c]')
[(1, 2, 3), ['a', 'b'], 'c']
"""
def find_next_comma(string):
"""Find the position of next comma in the string.
If no ',' is found in the string, return the string length. All
chars inside '()' and '[]' are treated as one element and thus ','
inside these brackets are ignored.
"""
assert (string.count('(') == string.count(')')) and (
string.count('[') == string.count(']')), \
f'Imbalanced brackets exist in {string}'
end = len(string)
for idx, char in enumerate(string):
pre = string[:idx]
# The string before this ',' is balanced
if ((char == ',') and (pre.count('(') == pre.count(')'))
and (pre.count('[') == pre.count(']'))):
end = idx
break
return end
# Strip ' and " characters and replace whitespace.
val = val.strip('\'\"').replace(' ', '')
is_tuple = False
if val.startswith('(') and val.endswith(')'):
is_tuple = True
val = val[1:-1]
elif val.startswith('[') and val.endswith(']'):
val = val[1:-1]
elif ',' not in val:
# val is a single value
return DictAction.parse_int_float_bool_str(val)
values = []
while len(val) > 0:
comma_idx = find_next_comma(val)
element = DictAction.parse_iterable(val[:comma_idx])
values.append(element)
val = val[comma_idx + 1:]
if is_tuple:
values = tuple(values)
return values
def __call__(self, parser, namespace, values, option_string):
options = {}
for kv in values:
key, val = kv.split('=', maxsplit=1)
options[key] = self.parse_iterable(val)
setattr(namespace, self.dest, options)
class SingleAction(DictAction):
""" Argparse action to convert value to tuple or list or nested structure of
list and tuple, i.e 'V1,V2,V3', or with explicit brackets, i.e. '[V1,V2,V3]'.
It also support nested brackets to build list/tuple values. e.g. '[(V1,V2),(V3,V4)]'
"""
def __call__(self, parser, namespace, value, option_string):
if isinstance(value, str):
setattr(namespace, self.dest, self.parse_iterable(value))
else:
setattr(namespace, self.dest, value)

View File

@@ -0,0 +1,79 @@
# Copyright (c) Alibaba, Inc. and its affiliates.
import glob
import os
import shutil
import tempfile
import unittest
import cv2
import json
import numpy as np
import torch
from modelscope.trainers.training_args import (ArgAttr, CliArgumentParser,
training_args)
from modelscope.utils.test_utils import test_level
class TrainingArgsTest(unittest.TestCase):
def setUp(self):
print(('Testing %s.%s' % (type(self).__name__, self._testMethodName)))
def tearDown(self):
super().tearDown()
@unittest.skipUnless(test_level() >= 0, 'skip test in current test level')
def test_define_args(self):
myparser = CliArgumentParser(training_args)
input_args = [
'--max_epochs', '100', '--work_dir', 'ddddd', '--train_batch_size',
'8', '--unkown', 'unkown'
]
args, remainning = myparser.parse_known_args(input_args)
myparser.print_help()
self.assertTrue(args.max_epochs == 100)
self.assertTrue(args.work_dir == 'ddddd')
self.assertTrue(args.train_batch_size == 8)
@unittest.skipUnless(test_level() >= 0, 'skip test in current test level')
def test_new_args(self):
training_args.num_classes = ArgAttr(
'model.mm_model.head.num_classes',
type=int,
help='number of classes')
training_args.mean = ArgAttr(
'train.data.mean', help='3-dim mean vector')
training_args.flip = ArgAttr('train.data.flip', help='flip or not')
training_args.img_size = ArgAttr(
'train.data.img_size', help='image size')
myparser = CliArgumentParser(training_args)
input_args = [
'--max_epochs', '100', '--work_dir', 'ddddd', '--train_batch_size',
'8', '--num_classes', '10', '--mean', '[125.0,125.0,125.0]',
'--flip', 'false', '--img_size', '(640,640)'
]
args, remainning = myparser.parse_known_args(input_args)
myparser.print_help()
self.assertTrue(args.max_epochs == 100)
self.assertTrue(args.work_dir == 'ddddd')
self.assertTrue(args.train_batch_size == 8)
self.assertTrue(args.num_classes == 10)
self.assertTrue(len(args.mean) == 3)
self.assertTrue(not args.flip)
self.assertAlmostEqual(args.mean[0], 125.0)
self.assertAlmostEqual(args.img_size, (640, 640))
cfg_dict = myparser.get_cfg_dict(args=input_args)
self.assertTrue(cfg_dict['model.mm_model.head.num_classes'] == 10)
self.assertAlmostEqual(cfg_dict['train.data.mean'],
[125.0, 125.0, 125.0])
self.assertTrue(not cfg_dict['train.data.flip'])
self.assertEqual(cfg_dict['train.dataloader.batch_size_per_gpu'], 8)
self.assertEqual(cfg_dict['train.work_dir'], 'ddddd')
self.assertEqual(cfg_dict['train.max_epochs'], 100)
self.assertEqual(cfg_dict['train.data.img_size'], (640, 640))
if __name__ == '__main__':
unittest.main()