使用Mindspore模型复现Big Transfer(BiT)模型

1. 背景介绍

本文是对笔者参加的华为昇腾AI创新大赛2022-昇思赛道比赛的一个总结和回顾,昇思赛道的题目即是使用Mindspore框架复现经典论文中的模型,笔者选做了其中一道赛题即是“利用MindSpore实现Big Transfer图像分类网络”。关于比赛的详情可以参见链接:昇腾AI创新大赛2022-昇思赛道

2. BiT模型介绍

论文链接: Big Transfer (BiT): General Visual Representation Learning

Big Transfer模型是基于Resnet-v2模型建立的,为了达到更好的迁移效果对Resnet-v2做了如下更改:①将所有的Batch Normalization替换为Group Normalization;②在所有卷积层中都使用了Weight Standardization;③Big Transfer论文使用的一些模型将Resnet-v2中网络层宽度做了加倍处理;④在迁移到下游任务进行微调时,对中等规模的数据集会使用mixup正则。

3. 实现过程

笔者编写代码使用的Mindspore版本为1.5.1,使用华为昇腾910处理器。模型构建主要参照了论文的官方实现代码,官方代码库中有pytorch,tensorflow多个框架下的实现版本,笔者在实现过程中主要参照了pytorch框架下的代码。

关于网络模型的构建,BiT模型只是在Resnet-v2模型的基础上做了一点点改进。如上图所示,左边的图即为原始Resnet网络一个单元模块的结构,Resnetv2相较于原始的网络将Batch Normalization层和ReLU激活层移到了卷积层的前面,其他结构均与原网络一致。BiT模型中将所有的BN层修改成了Group Normalization层,同时在所有的卷积层中添加了Weight Standardization对权重进行归一化操作。其中Group Normalization可直接使用mindspore的nn.GroupNorm()层,Weight Standardization的实现如下:

class StdConv2d(nn.Conv2d):
    def __init__(self, in_channels, out_channels, kernel_size, stride=1, pad_mode='same', padding=0, dilation=1, group=1, has_bias=False, weight_init='normal', bias_init='zeros', data_format='NCHW'):
        super().__init__(in_channels, out_channels, kernel_size, stride, pad_mode, padding, dilation, group, has_bias, weight_init, bias_init, data_format)
        self.ops_conv2d = ops.Conv2D(self.out_channels, self.kernel_size,
                                 pad_mode=self.pad_mode, pad=self.padding, stride=self.stride,)
        self.reduce_mean = ops.ReduceMean(keep_dims=True)

    def construct(self, x):
        w = self.weight
        m = self.reduce_mean(w, (1, 2, 3))
        w = w - m
        v = self.reduce_mean(ops.square(w), (1, 2, 3))
        w = w / ops.sqrt(v + 1e-10)
        # bias does not matter when combined with GN
        #"dilation","data_format" maybe can be removed
        ops_conv2d = ops.Conv2D(self.out_channels, self.kernel_size, pad_mode=self.pad_mode,
                                pad=self.padding, stride=self.stride, dilation=self.dilation,
                                group=self.group, data_format=self.format)
        return  ops_conv2d(x, w)

这样我们就可以构建出模型的一个block单元,代码如下所示:

class PreActBottleneck(nn.Cell):
    def __init__(self, cin, cout=None, cmid=None, stride=1):
        super().__init__()
        cout = cout or cin
        cmid = cmid or cout//4

        self.gn1 = nn.GroupNorm(32, cin)
        self.conv1 = conv1x1(cin, cmid)
        self.gn2 = nn.GroupNorm(32, cmid)
        self.conv2 = conv3x3(cmid, cmid, stride)  # Original code has it on conv1!!
        self.gn3 = nn.GroupNorm(32, cmid)
        self.conv3 = conv1x1(cmid, cout)
        self.relu = nn.ReLU()
        self.downsample = None

        if (stride != 1 or cin != cout):
            # Projection also with pre-activation according to paper.
            self.downsample = conv1x1(cin, cout, stride)
    
    def construct(self, x):
        out = self.relu(self.gn1(x))

        # residual branch
        residual = x
        if self.downsample is not None:
            residual = self.downsample(out)
        
        # Unit's branch
        out = self.conv1(out)
        out = self.conv2(self.relu(self.gn2(out)))
        out = self.conv3(self.relu(self.gn3(out)))

        return out + residual

最后只需要按照不同的网络深度将对应数量的这些block堆叠起来就可以了,另外BiT模型中使用了不同宽度的网络模型如ResNet152x4,这个只对每一层网络通道数进行了加倍,ResNet152x4模型即是对ResNet152模型每层网络通道数加宽到4倍。

模型使用的数据集为Cifar10数据集,可以直接Mindspore.dataset中的接口载入数据,在这个过程中也可以mindspore.dataset.vision.c_transforms模块对数据进行一些预处理和数据增强操作。

def mktrainval(args, logger):
    """Returns train and validation datasets."""
    precrop, crop = bit_hyperrule.get_resolution_from_dataset(args.dataset)
    # ========transforms========
    train_transform = [
        c_trans.Resize((precrop, precrop)),
        c_trans.RandomCrop((crop, crop)),
        c_trans.RandomHorizontalFlip(),
        c_trans.Normalize((0.5*255, 0.5*255, 0.5*255), (0.5*255, 0.5*255, 0.5*255)),
        c_trans.HWC2CHW()
    ]
    val_transform = [
        c_trans.Resize((crop, crop)),
        c_trans.Normalize((0.5*255, 0.5*255, 0.5*255), (0.5*255, 0.5*255, 0.5*255)),
        c_trans.HWC2CHW()
    ]
    type_cast_op = c2_trans.TypeCast(ms.int32) # cast label to int32 datatype

    # ========load data========
    if args.dataset == "cifar10":
        # "/home/ma-user/work/workspace_lc/dist_test/dataset/cifar-10-batches-bin"
        train_set = ds.Cifar10Dataset(args.datadir, "train", shuffle=True)
        valid_set = ds.Cifar10Dataset(args.datadir, "test", shuffle=True, num_samples=1000)
    else:
        # other datasets to be completed
        pass

    num_train, num_val = train_set.get_dataset_size(), valid_set.get_dataset_size()
    logger.info(f"Using a training set with {num_train} images.")
    logger.info(f"Using a validation set with {num_val} images.")

    micro_batch_size = args.batch // args.batch_split

    train_set = train_set.map(operations=train_transform, input_columns=["image"])
    train_set = train_set.map(operations=type_cast_op, input_columns=["label"])
    train_set = train_set.batch(batch_size=micro_batch_size)
    valid_set = valid_set.map(operations=val_transform, input_columns=["image"])
    valid_set = valid_set.map(operations=type_cast_op, input_columns=["label"])
    valid_set = valid_set.batch(batch_size=micro_batch_size)

    return train_set, valid_set, num_train, num_val

由于时间和算力的限制,比赛中直接使用了官方提供的预训练权重文件,因为Mindpore的权重文件为ckpt格式的,所以需要先将权重文件转换为ckpt格式。官方代码库中提供了npz格式的预训练权重,下面代码实现了将BiT-M-R50x1模型的权重文件转化为ckpt格式。

def npz2ckpt(zero_head):
    weights = np.load("BiT-M-R50x1.npz")
    # print(type(weights))
    save_list = []
    save_list.append({"name": "root.conv.weight",
                      "data": ms.Parameter(ms.Tensor(weights["resnet/root_block/standardized_conv2d/kernel"]))})
    save_list.append({"name": "head.gn.gamma",
                      "data": ms.Parameter(ms.Tensor(weights["resnet/group_norm/gamma"]))})
    save_list.append({"name": "head.gn.beta",
                      "data": ms.Parameter(ms.Tensor(weights["resnet/group_norm/beta"]))})
    
    if zero_head == False:
        save_list.append({"name": "head.conv.weight",
                        "data": ms.Parameter(ms.Tensor(weights["resnet/head/conv2d/kernel"]))})
        save_list.append({"name": "head.conv.bias",
                        "data": ms.Parameter(ms.Tensor(weights["resnet/head/conv2d/bias"]))})
    
    convname = 'standardized_conv2d'
    for block_id in range(1,5):
        bname = f"block{block_id}"
        num_unit = [3, 4, 6, 3]
        for unit_id in range(1, num_unit[block_id-1]+1):
            prefix_sc = f"resnet/{bname}/unit{unit_id:02d}/"
            prefix_tg = f"body.{bname}.unit{unit_id:02d}."
            save_list.append({"name": prefix_tg + "conv1.weight",
                            "data": ms.Parameter(ms.Tensor(weights[prefix_sc+f"a/{convname}/kernel"]))})
            save_list.append({"name": prefix_tg + "conv2.weight",
                            "data": ms.Parameter(ms.Tensor(weights[prefix_sc+f"b/{convname}/kernel"]))})
            save_list.append({"name": prefix_tg + "conv3.weight",
                            "data": ms.Parameter(ms.Tensor(weights[prefix_sc+f"c/{convname}/kernel"]))})
            save_list.append({"name": prefix_tg + "gn1.gamma",
                            "data": ms.Parameter(ms.Tensor(weights[prefix_sc+f"a/group_norm/gamma"]))})
            save_list.append({"name": prefix_tg + "gn2.gamma",
                            "data": ms.Parameter(ms.Tensor(weights[prefix_sc+f"b/group_norm/gamma"]))})
            save_list.append({"name": prefix_tg + "gn3.gamma",
                            "data": ms.Parameter(ms.Tensor(weights[prefix_sc+f"c/group_norm/gamma"]))})
            save_list.append({"name": prefix_tg + "gn1.beta",
                            "data": ms.Parameter(ms.Tensor(weights[prefix_sc+f"a/group_norm/beta"]))})
            save_list.append({"name": prefix_tg + "gn2.beta",
                            "data": ms.Parameter(ms.Tensor(weights[prefix_sc+f"b/group_norm/beta"]))})
            save_list.append({"name": prefix_tg + "gn3.beta",
                            "data": ms.Parameter(ms.Tensor(weights[prefix_sc+f"c/group_norm/beta"]))})
            
            if(unit_id == 1):
                save_list.append({"name": prefix_tg + "downsample.weight",
                                "data": ms.Parameter(ms.Tensor(weights[prefix_sc+f"a/proj/{convname}/kernel"]))})
        
        ms.save_checkpoint(save_list, "BiT-M-R50x1.ckpt")

另外笔者也实现了分布式多卡训练,主要通过编写bash脚本实现,具体实现方法可以参考官方教程文档,这里提供笔者编写的dist_train.sh脚本供参考:

#!/usr/bin/bash

echo "=============================================================================================================="
echo "Please run the script as: "
echo "bash dist_train.sh RANK_SIZE [optional arguments]"
echo "For example: bash dist_train.sh 8 --name dist_152x2 --model BiT-M-R152x2 --logdir ./bit_logs --dataset cifar10 --datadir /home/ma-user/work/data/cifar-10-batches-bin --base_lr 0.003 --batch 64 --eval_every 100 --num_classes 10"
echo "It is better to use the absolute path."
echo "=============================================================================================================="
RANK_SIZE=$1

EXEC_PATH=$(pwd)

export PYTHONPATH="$EXEC_PATH":$PYTHONPATH

for((i=1;i<${RANK_SIZE};i++))
do
    rm -rf device$i
    mkdir device$i
    cp ./train_ms_graph_dist.py ./device$i
    cd ./device$i
    export DEVICE_ID=$i
    export RANK_ID=$i
    echo "start training for device $i"
    env > env$i.log
    python ./train_ms_graph_dist.py ${@:2} > train.log$i 2>&1 &
    cd ../
done
rm -rf device0
mkdir device0
cp ./train_ms_graph_dist.py ./device0
cd ./device0
export DEVICE_ID=0
export RANK_ID=0
echo "start training for device 0"
env > env0.log
python ./train_ms_graph_dist.py ${@:2} > train.log0 2>&1
if [ $? -eq 0 ];then
    echo "training success"
else
    echo "training failed"
    exit 2
fi
cd ../

笔者实现的完整代码开源在github上,代码编写可能不太美观,仅供参考,链接为:https://github.com/Vincent-luo/BiT_mindspore/

4. 总结和体会

在使用Mindspore框架时感觉代码的编写与pytorch还是挺相似的,不过要想自定义网络的训练方式时Mindspore会显得稍微麻烦一点,需要定义一个nn.TrainOneStepCell类,并且网络的梯度计算是通过GradOperation算子进行获取的。另外,最大的一点体会就是Mindspore框架会为分动态图和静态图两种模式,动态图方便编写和调试程序,但是运行速度较慢;静态图则需要在运行前进行编译,相对动态图代码编写上会有一些限制,但是运行速度较快。笔者在实现过程中先编写了动态图下的代码,发现训练速度有点太长了,需要十几个小时有点超出了能接受的范围,修改为静态图后速度才缩短到一到两个小时。不知道是否是自己实现的问题,静态图版本训练的模型仍然会比pytorch下的版本要慢,并且精度也有些微的下降。

总体上通过这次比赛,自己也对Mindspore框架有了初步的认识,当然自己肯定有很多地方能理解的不是太清楚,不过听说Mindspore框架后续会提高动态图的运行性能,进一步改善模型的易用性,所以还是很看好Mindspore的发展的,希望未来能发展得越来越好,让我们在选择深度学习框架时多出一个好的选择。

posted @ 2022-11-01 16:21  Vincent-luo  阅读(412)  评论(0)    收藏  举报