Files
modelscope-scepter/docs/zh_cn/scepter/model.md
T
2024-04-18 15:53:37 +08:00

5.3 KiB
Raw Blame History

模型模块 (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下自定义;