model_args.py 2.0 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667686970
  1. import abc
  2. import typing as T
  3. from chainer_addons.links import PoolingType
  4. from chainer_addons.models import PrepareType
  5. from cvargparse import Arg
  6. from cvargparse import BaseParser
  7. from cvfinetune.parser.utils import parser_extender
  8. from cvmodelz.models import ModelFactory
  9. @parser_extender
  10. def add_model_args(parser: BaseParser, *,
  11. model_modules: T.Optional[T.List[str]] = None) -> None:
  12. if model_modules is None:
  13. model_modules = ["chainercv2", "cvmodelz"]
  14. choices = ModelFactory.get_models(model_modules)
  15. _args = [
  16. Arg("--model_type", "-mt",
  17. required=True,
  18. choices=choices,
  19. help="type of the model"),
  20. Arg("--pretrained_on", "-pt",
  21. default="imagenet",
  22. choices=["imagenet", "inat"],
  23. help="type of model pre-training"),
  24. Arg("--input_size", type=int, nargs="+", default=0,
  25. help="overrides default input size of the model, if greater than 0"),
  26. Arg("--parts_input_size", type=int, nargs="+", default=0,
  27. help="overrides default input part size of the model, if greater than 0"),
  28. PrepareType.as_arg("prepare_type",
  29. help_text="type of image preprocessing"),
  30. PoolingType.as_arg("pooling",
  31. help_text="type of pre-classification pooling"),
  32. Arg("--load", type=str,
  33. help="ignore weights and load already fine-tuned model (classifier will NOT be re-initialized and number of classes will be unchanged)"),
  34. Arg("--weights", type=str,
  35. help="ignore default weights and load already pre-trained model (classifier will be re-initialized and number of classes will be changed)"),
  36. Arg("--headless", action="store_true",
  37. help="ignores classifier layer during loading"),
  38. Arg("--load_strict", action="store_true",
  39. help="load weights in a strict mode"),
  40. Arg("--load_path", type=str, default="",
  41. help="load path within the weights archive"),
  42. ]
  43. parser.add_args(_args, group_name="Model arguments")
  44. class ModelParserMixin(abc.ABC):
  45. def __init__(self, *args, **kwargs):
  46. super(ModelParserMixin, self).__init__(*args, **kwargs)
  47. add_model_args(self)
  48. __all__ = [
  49. "ModelParserMixin",
  50. "add_model_args"
  51. ]