mmaction2 指南 (5) 自定义新模块

本文档介绍了如何在mmaction2中自定义新模块,包括优化器、新组件如backbone、head和损失函数,以及添加学习率调度策略。通过实例展示了在模块代码和配置文件中的实现方法。

mmaction2 指南

5. 自定义新模块

自定义优化器

参考 CopyOfSGD
例子:
mmaction/core/optimizer/my_optimizer.py

from .registry import OPTIMIZERS
from torch.optim import Optimizer

@OPTIMIZERS.register_module()
class MyOptimizer(Optimizer):

    def __init__(self, a, b, c):

模块的 __init__.py中写入from .my_optimizer import MyOptimizer,这样 registry will find the new module and add it:
配置文件中写:

optimizer = dict(type='SGD', lr=0.02, momentum=0.9, weight_decay=0.0001)
optimizer = dict(type='Adam', lr=0.0003, weight_decay=0.0001)
optimizer = dict(type='MyOptimizer', a=a_value, b=b_value, c=c_value)

自定义优化器 Constructor

有些 fine-grained parameter,比如只想在BN层weight decay,需要写 optimizer constructor,继承自DefaultOptimizerConstructor, 重写 add_params(self, params, module) method.

TSM的优化器构造例子:
TSMOptimizerConstructor

自定义的话:mmaction/core/optimizer/my_optimizer_constructor.py

from mmcv.runner import OPTIMIZER_BUILDERS, DefaultOptimizerConstructor

@OPTIMIZER_BUILDERS.register_module()
class MyOptimizerConstructor(DefaultOptimizerConstructor):

mmaction/core/optimizer/__init__.py 中加入from .my_optimizer_constructor import MyOptimizerConstructor

配置中加入

# optimizer
optimizer = dict(
    type='SGD',
    constructor='MyOptimizerConstructor',
    paramwise_cfg=dict(fc_lr5=True),
    lr=0.02,
    momentum=0.9,
    weight_decay=0.0001)

自定义新组件

新的backbone

新建 mmaction/models/backbones/resnet.py

import torch.nn as nn

from ..registry import BACKBONES

@BACKBONES.register_module()
class ResNet(nn.Module):

    def __init__(self, arg1, arg2):
        pass

    def forward(self, x):  # should return a tuple
        pass

    def init_weights(self, pretrained=None):
        pass

mmaction/models/backbones/__init__.py 加入 from .resnet import ResNet

配置中写入

model = dict(
    ...
    backbone=dict(
        type='ResNet',
        arg1=xxx,
        arg2=xxx),
)
新的head

创建mmaction/models/heads/tsn_head.py,继承BaseHead 重写 init_weights(self)forward(self, x)

from ..registry import HEADS
from .base import BaseHead


@HEADS.register_module()
class TSNHead(BaseHead):

    def __init__(self, arg1, arg2):
        pass

    def forward(self, x):
        pass

    def init_weights(self):
        pass

mmaction/models/heads/__init__.py中加入from .tsn_head import TSNHead
配置中写入

model = dict(
    ...
    cls_head=dict(
        type='TSNHead',
        num_classes=400,
        in_channels=2048,
        arg1=xxx,
        arg2=xxx),
新的损失

mmaction/models/losses/my_loss.py

import torch
import torch.nn as nn

from ..builder import LOSSES

def my_loss(pred, target):
    assert pred.size() == target.size() and target.numel() > 0
    loss = torch.abs(pred - target)
    return loss


@LOSSES.register_module()
class MyLoss(nn.Module):

    def forward(self, pred, target):
        loss = my_loss(pred, target)
        return loss

mmaction/models/losses/__init__.py中加入from .my_loss import MyLoss, my_loss
配置中写入 loss_bbox=dict(type='MyLoss'))

添加新的学习策略(learning rate scheduler (updater))

lr_config = dict(policy='step', step=[20, 40])

1.文件 $MMAction2/mmaction/core/lr 中写入 LrUpdaterHook, 继承自 mmcv.LrUpdaterHook

@HOOKS.register_module()
# Register it here
class RelativeStepLrUpdaterHook(LrUpdaterHook):
    # You should inheritate it from mmcv.LrUpdaterHook
    def __init__(self, runner, steps, lrs, **kwargs):
        super().__init__(**kwargs)
        assert len(steps) == (len(lrs))
        self.steps = steps
        self.lrs = lrs

    def get_lr(self, runner, base_lr):
        # Only this function is required to override
        # This function is called before each training epoch, return the specific learning rate here.
        progress = runner.epoch if self.by_epoch else runner.iter
        for i in range(len(self.steps)):
            if progress < self.steps[i]:
                return self.lrs[i]

配置中写入 lr_config = dict(policy='RelativeStep', steps=[20, 40, 60], lrs=[0.1, 0.01, 0.001])

评论
添加红包

请填写红包祝福语或标题

红包个数最小为10个

红包金额最低5元

当前余额3.43前往充值 >
需支付:10.00
成就一亿技术人!
领取后你会自动成为博主和红包主的粉丝 规则
hope_wisdom
发出的红包
实付
使用余额支付
点击重新获取
扫码支付
钱包余额 0

抵扣说明:

1.余额是钱包充值的虚拟货币,按照1:1的比例进行支付金额的抵扣。
2.余额无法直接购买下载,可以购买VIP、付费专栏及课程。

余额充值