File size: 2,279 Bytes
9e14838 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 | #!/usr/bin/python
# -*- coding: UTF-8 -*-
""" Plugin loader for extract, training and model tasks """
from utils import logger
import os
from importlib import import_module
from typing import Type
from trainer._base import TrainerBase
from torch.utils.data import Dataset
from model._base import ModelBase
class PluginLoader():
"""
Plugin loader for extract, training and model tasks
function: get_{model_type}
args: {model_name}
will return the a class named model_type under model_type.model_name.py
as it return a class you should also annotate the returning classtype to make
code linting avaliable in some IDE
"""
@staticmethod
def get_classifier(name) -> Type[ModelBase]:
""" Return requested attribute encoder plugin """
return PluginLoader._import("model.classifier", name)
@staticmethod
def get_trainer(name) -> Type[TrainerBase]:
""" Return requested trainer plugin """
return PluginLoader._import("trainer", name)
@staticmethod
def get_dataset(name) -> Type[Dataset]:
""" Return requested trainer plugin """
return PluginLoader._import("dataset", name)
@staticmethod
def _import(attr, name):
""" Import the plugin's module """
name = name.replace("-", "_")
ttl = attr.split(".")[-1].title()
logger.info("Loading %s from %s plugin...", ttl, name.title())
attr = "model" if attr == "Trainer" else attr.lower()
mod = ".".join((attr, name))
module = import_module(mod)
logger.info(str(module) + str(ttl))
return getattr(module, ttl)
@staticmethod
def get_available_trainer():
""" Return a list of available models """
modelpath = os.path.join(os.path.dirname(__file__), "trainer")
models = sorted(item.name.replace(".py", "").replace("_", "-")
for item in os.scandir(modelpath)
if not item.name.startswith("_")
and item.name.endswith(".py"))
return models
@staticmethod
def get_default_model():
""" Return the default model """
models = PluginLoader.get_available_models()
return 'original' if 'original' in models else models[0]
|