diff --git a/examples/pytorch/finetune_image_classification.py b/examples/pytorch/finetune_image_classification.py new file mode 100644 index 00000000..cb30c32c --- /dev/null +++ b/examples/pytorch/finetune_image_classification.py @@ -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) diff --git a/examples/pytorch/run_train.sh b/examples/pytorch/run_train.sh new file mode 100644 index 00000000..2093fa09 --- /dev/null +++ b/examples/pytorch/run_train.sh @@ -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' diff --git a/modelscope/trainers/training_args.py b/modelscope/trainers/training_args.py new file mode 100644 index 00000000..c387e7b8 --- /dev/null +++ b/modelscope/trainers/training_args.py @@ -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) diff --git a/tests/trainers/test_training_args.py b/tests/trainers/test_training_args.py new file mode 100644 index 00000000..0aad9ddc --- /dev/null +++ b/tests/trainers/test_training_args.py @@ -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()