5.3 KiB
模型模块 (Model)
Overview
模型模块分为backbone、neck、head、loss、metric、network、tokenizer、tuner;
- backbone/neck:一般为提取feature主要模块(necks非必有);
- head:根据不同任务类型,输入backbone提取的feature,输出下游任务所需logit;
- loss:用于计算不同类型loss;
- metric:用于计算各类评测指标;
- tokenizer:用于分词;
- tuner:用于创建微调模块;
- network:train和test模块,对数据集输入的batch整合上述模块进行最终loss和指标计算;
backbone/neck/head/loss/tuner
Basic Usage
子类注册:
from scepter.modules.model.registry import BACKBONES
from scepter.modules.model.base_model import BaseModel
@BACKBONES.register_class("ResNet")
class ResNet(BaseModel):
def __init__(self, cfg, logger=None):
super(ResNet, self).__init__(cfg, logger=logger)
from scepter.modules.model.registry import NECKS
from scepter.modules.model.base_model import BaseModel
@NECKS.register_class()
class GlobalAveragePooling(BaseModel):
def __init__(self, cfg, logger=None):
super(GlobalAveragePooling, self).__init__(cfg, logger=logger)
from scepter.modules.model.registry import HEADS
from scepter.modules.model.base_model import BaseModel
@HEADS.register_class()
class ClassifierHead(BaseModel):
def __init__(self, cfg, logger=None):
super(ClassifierHead, self).__init__(cfg, logger=logger)
from scepter.modules.model.registry import LOSSES
import torch.nn as nn
@LOSSES.register_class()
class CrossEntropy(nn.Module):
def __init__(self, cfg, logger=None):
super(CrossEntropy, self).__init__(cfg, logger=logger)
实际调用:
from scepter.modules.model.registry import BACKBONES, NECKS, HEADS, LOSSES, TUNERS
backbone = BACKBONES.build(cfg.BACKBONE, logger=logger)
neck = NECKS.build(cfg.NECK, logger=logger)
head = HEADS.build(cfg.HEAD, logger=logger)
loss = LOSSES.build(cfg.LOSS, logger=logger)
tuner = TUNERS.build(cfg.TUNER, logger=logger)
function __init__()
Parameters(输入参数)
(cfg: scepter.modules.utils.config.Config, logger = None) -> None
主要用于初始化模型各layer;
function forward()
根据需要具体实现;
metric
Basic Usage
子类注册:
from scepter.modules.model.metrics.registry import METRICS
from scepter.modules.model.metrics.base_metric import BaseMetric
@METRICS.register_class("AccuracyMetric")
class AccuracyMetric(BaseMetric):
def __init__(self, cfg, logger=None):
super(CrossEntropy, self).__init__(cfg, logger=logger)
实际用法:
from scepter.modules.model.metrics.registry import METRICS
metric = METRICS.build(cfgs, logger)
function __init__()
Parameters(输入参数)
(cfg: scepter.modules.utils.config.Config, logger = None) -> None
初始化计算metric所需超参,例如topk等系数;
function __call__()
@torch.no_grad()
通常输入logit和label以及其他所需要的变量,输出计算指标;
tokenizer
Basic Usage
子类注册:
from scepter.modules.model.registry import TOKENIZERS
from scepter.modules.model.tokenizers import BaseTokenizer
@TOKENIZERS.register_class()
class BaseBertTokenizer(BaseTokenizer):
def __init__(self, cfg, logger=None):
super(BaseBertTokenizer, self).__init__(cfg, logger=logger)
实际用法:
from scepter.modules.model.registry import TOKENIZERS
tokenizer = TOKENIZERS.build(cfgs, logger)
function __init__()
Parameters(输入参数)
(cfg: scepter.modules.utils.config.Config, logger = None) -> None 用于初始化加载tokenizer对象,例如BertTokenizer;
function tokenize()
输入需要分词的text list,输出分词后转换的token id sequence以及其他所需的attention mask/tpye id list/position id list等;
network
Basic Usage
子类注册:
from scepter.modules.model.registry import MODELS
from scepter.modules.model.networks.train_module import TrainModule
@MODELS.register_class()
class Classifier(TrainModule):
def __init__(self, cfg, logger=None):
super(Classifier, self).__init__(cfg, logger=logger)
实际用法:
from scepter.modules.model.registry import MODELS
model = MODELS.build(self.cfg.MODEL, logger=self.logger)
function __init__()
Parameters(输入参数)
(cfg: scepter.modules.utils.config.Config, logger = None) -> None
结合上述build方法,初始化训练所需的backbone、neck、head、loss、metric、tokenizer模块;
function forward_train()
输入训练batch的数据,经backbone、neck、head、loss计算相关loss;
function forward_test()
输入测试batch的数据,经backbone、neck、head、metrics计算相关指标;
function forward
实际调用接口,用于分发任务至forward_train()/forward_test();
其他训练/测试所需函数可在network下自定义;