mirror of
https://github.com/modelscope/modelscope.git
synced 2026-09-01 19:49:03 +02:00
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:
67
examples/pytorch/finetune_image_classification.py
Normal file
67
examples/pytorch/finetune_image_classification.py
Normal 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)
|
||||
5
examples/pytorch/run_train.sh
Normal file
5
examples/pytorch/run_train.sh
Normal 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'
|
||||
270
modelscope/trainers/training_args.py
Normal file
270
modelscope/trainers/training_args.py
Normal 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)
|
||||
79
tests/trainers/test_training_args.py
Normal file
79
tests/trainers/test_training_args.py
Normal 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()
|
||||
Reference in New Issue
Block a user