基于MindSpore实现Vision Transformer图像分类
收藏回复举报
基于MindSpore实现Vision Transformer图像分类
发表于2024-12-07 15:21:13
0 查看

前言

在了解Vision Transformer之前,我们需要先了解一下Transformer,Transformer最开始是应用在NLP领域的,拿过来用到Vision中就叫Vision Transformer。而这里要提到的,就是Transformer中的self-Attention(自注意力)和Multiple-Head Attention(多头注意力)。

用在NLP领域中用到的注意力机制举例,一般为Encoder-Decoder框架,比如中英翻译,输入的英文是Source,我们要获取到的是Target(中文翻译),Attention机制就发生在Target的元素Query和Source中的所有元素之间,其同时关注自身和目标值。

而这里说的自注意力机制只关注自身,比如Source中会有一个注意力机制,Target中会有一个注意力机制,他两是没有关系的。

还是用中英翻译举例,注意力机制的查询和键分别来自于英文和中文,通过查询(Query)英文单词,去匹配中文汉字的键(Key),自注意力机制只关注自己一个语言,可以理解为:”我喜欢“后面可以跟”你“,也可以跟”吃饭“。

1)如果查询和键是同一组内的特征,并且相互做注意力机制,则称为自注意力机制或内部注意力机制。 2)多头注意力机制的多头表示对每个Query和所有的Key-Value做多次注意力机制。做两次,就是两头,做三次,就是三头。这样做的意义在于获取每个Query和所有的Key-Value的不同的依赖关系。 3)自注意力机制的优缺点简记为【优点:感受野大。缺点:需要大数据。】

以下是关于这两个自注意力机制的官方公式,很复杂也很难理解,但现在别盯着他不放,先慢慢往下看,这篇就是说明这个公式及其过程:

true

Self-Attention

true

我们先说明白这里面这些符号都是干啥的,或者求出来用来干啥的,避免看半天还一头雾水:

q代表query,后续会去和每一个k进行匹配

k 代表key,后续会被每个q匹配

v 代表从a中提取得到的信息,后续会和q和k的乘积进行运算

d是k的维度

后续q 和k匹配的过程可以理解成计算两者的相关性,相关性越大对应v的权重也就越大

简单来说,最初的输入向量首先会经过Embedding层映射成Q(Query),K(Key),V(Value)三个向量,由于是并行操作,所以代码中是映射成为dim x 3的向量然后进行分割,换言之,如果你的输入向量为一个向量序列(𝑥1,𝑥2,𝑥3),其中的𝑥1,𝑥2,𝑥3都是一维向量,那么每一个一维向量都会经过Embedding层映射出Q,K,V三个向量,只是Embedding矩阵不同,矩阵参数也是通过学习得到的。这里大家可以认为,Q,K,V三个矩阵是发现向量之间关联信息的一种手段,需要经过学习得到,至于为什么是Q,K,V三个,主要是因为需要两个向量点乘以获得权重,又需要另一个向量来承载权重向加的结果,所以,最少需要3个矩阵。

后续我们要用q*k得到v的权重,然后进行一定缩放(除以根号d),再乘上v,就是第一个公式。

从数值上理解

行内公式显示有问题,这里直接贴图 true

true

true

true

从维度上进行理解

我们假设载入的$x_1$经过Embedding后变为$a_1$维度为1X4,$W^q$的维度为4X3,两者进行叉乘运算后就得到了维度为1X3的Query,k和v同理

然后我们吧a1和a2并行起来

true

然后把公式中的式子也换成维度:

true

整个过程放在一张图上可以这么看:

true

true

true

true

接下来就要将每个head的结果进行拼接,此时还是以两个head举例:

true

这个图里面的b大家可能忘了,这个b就是Self-Attention中求得的最后结果,在多头注意力这边,这个结果还要再进行计算。

true

true

true

true

true

模型网络结构

Vision Transformer(ViT)模型主要由三个模块组成,以下是模型框架:

  • Linear Projection of Flattened Patches(Embedding层)
  • Transformer Encoder(图右侧有给出更加详细的结构)
  • MLP Head(最终用于分类的层结构)

true 首先,Transformer是从NLP领域学习改进过来的嘛,所以人家训练的认识的是字,也就是一串token(向量)序列,而我们直接给模型输入一个图片,人认识都不认识,更别说预测了。

Embedding层(Linear Projection of Flattened Patches)结构

对于图像数据而言,其数据格式为[H,W,C],我们就要先通过Embedding层(图中的Linear Projection of Flattened Patches)给他转换一下,尝试把图片分为一堆小的Patches,此处把输入图片按照16*16的Patch进行划分(就图中左下角画的九宫格这个意思),每个Patch数据shape会变为[16,16,3],通过映射得到一个长度为768的向量(token),即[16,16,3]->[768]。

代码实现中直接使用卷积层,ViT也有好几种类型,此处以ViT-B/16为例,直接使用shape=16X16,stride=16,个数768的卷积层,原输入图像通过这层卷积后维度由[224,224,3]变成[14,14,768],然后把H和W两个维度展平(两个14),即[14,14,768]->[196,768],此时这个二维矩阵正是Transformer想要的。

在输入Transformer Encoder之前要加上[class]token以及Position Embedding。这个[class]token是用于分类的,是一个可训练的参数,数据格式和上面得到的token一样都是一个向量,此处就是一个长度为768的向量,维度为[1,768],与上面的token拼接在一起就是[197,768],此外,还要添加一个Position Embedding用于定位,定位拆出来的这个块在原图的什么位置,如果没有这个Position Embedding,这就是一堆乱的拼图,有时候我们不知道原图的情况下,玩儿4*4的拼图都费劲,更别说让机器啥都不知道来拼16X16的拼图了,这个Position Embedding原理如下:

true

图片右侧有色条,最上面最黄色的部分就是相似度最高的,可以看图中左上角第一张图,他的左上角(1,1)的位置就是它本身嘛,肯定就是最像自己的地方,所以在颜色表现上就是最黄的,第一行和第一列都是和他相似度比较高的,所以颜色都在色条的上半部分。就通过这个表现记录了其位置。

Position Embedding是直接在原来的token上进行相加,所以shape应该与原来的token保持一致,为[197,768],在这个加法过程中维度不会发生变化,是直接相加。

true

Transformer Encoder

Transformer Encoder其实就是重复堆叠Encoder Block L次,下图是大佬绘制的Encoder Block,主要由以下几部分组成:

  • Layer Norm,这种Normalization方法主要是针对NLP领域提出的,这里是对每个token进行Norm处理
  • Dropout/DropPath,在原论文的代码中是直接使用的Dropout层,在但rwightman实现的代码中使用的是DropPath(stochastic depth),可能后者会更好一点。
  • MLP Block,如图右侧所示,就是全连接+GELU激活函数+Dropout组成也非常简单,需要注意的是第一个全连接层会把输入节点个数翻4倍[197, 768] -> [197, 3072],第二个全连接层会还原回原节点个数[197, 3072] -> [197, 768]

true

MLP Head

在Transformer Encoder中,输出的shape和输入得shape是一样的,输出的还是[197,768],出来之后我们要添加一个Layer Norm层,把我们之前添加进去的class[token]拿出来,这么费工夫不就是为了最后得到这个分类信息嘛,然后通过MLP Head得到最终的分类结果,在训练自己的数据集时,这一层只需要一个简单的Linear,在原论文训练数据集上较为复杂,由Linear+tanh激活函数+Linear组成。

true

实战

接下来我们具体通过代码来学习:

from mindspore import nn, ops


class Attention(nn.Cell):
   def __init__(self,
                dim: int,
                num_heads: int = 8,
                keep_prob: float = 1.0,
                attention_keep_prob: float = 1.0):
       super(Attention, self).__init__()

       self.num_heads = num_heads
       head_dim = dim // num_heads
       self.scale = ms.Tensor(head_dim ** -0.5)

       self.qkv = nn.Dense(dim, dim * 3)
       self.attn_drop = nn.Dropout(p=1.0-attention_keep_prob)
       self.out = nn.Dense(dim, dim)
       self.out_drop = nn.Dropout(p=1.0-keep_prob)
       self.attn_matmul_v = ops.BatchMatMul()
       self.q_matmul_k = ops.BatchMatMul(transpose_b=True)
       self.softmax = nn.Softmax(axis=-1)

   def construct(self, x):
       """Attention construct."""
       b, n, c = x.shape
       qkv = self.qkv(x)
       qkv = ops.reshape(qkv, (b, n, 3, self.num_heads, c // self.num_heads))
       qkv = ops.transpose(qkv, (2, 0, 3, 1, 4))
       q, k, v = ops.unstack(qkv, axis=0)
       attn = self.q_matmul_k(q, k)
       attn = ops.mul(attn, self.scale)
       attn = self.softmax(attn)
       attn = self.attn_drop(attn)
       out = self.attn_matmul_v(attn, v)
       out = ops.transpose(out, (0, 2, 1, 3))
       out = ops.reshape(out, (b, n, c))
       out = self.out(out)
       out = self.out_drop(out)

       return out

Transformer Encoder

from typing import Optional, Dict


class FeedForward(nn.Cell):
   def __init__(self,
                in_features: int,
                hidden_features: Optional[int] = None,
                out_features: Optional[int] = None,
                activation: nn.Cell = nn.GELU,
                keep_prob: float = 1.0):
       super(FeedForward, self).__init__()
       out_features = out_features or in_features
       hidden_features = hidden_features or in_features
       self.dense1 = nn.Dense(in_features, hidden_features)
       self.activation = activation()
       self.dense2 = nn.Dense(hidden_features, out_features)
       self.dropout = nn.Dropout(p=1.0-keep_prob)

   def construct(self, x):
       """Feed Forward construct."""
       x = self.dense1(x)
       x = self.activation(x)
       x = self.dropout(x)
       x = self.dense2(x)
       x = self.dropout(x)

       return x


class ResidualCell(nn.Cell):
   def __init__(self, cell):
       super(ResidualCell, self).__init__()
       self.cell = cell

   def construct(self, x):
       """ResidualCell construct."""
       return self.cell(x) + x

下面的代码我们将TransformerEncoder结构和一个多层感知器(MLP)结合,就构成了ViT模型的backbone部分。

class TransformerEncoder(nn.Cell):
   def __init__(self,
                dim: int,
                num_layers: int,
                num_heads: int,
                mlp_dim: int,
                keep_prob: float = 1.,
                attention_keep_prob: float = 1.0,
                drop_path_keep_prob: float = 1.0,
                activation: nn.Cell = nn.GELU,
                norm: nn.Cell = nn.LayerNorm):
       super(TransformerEncoder, self).__init__()
       layers = []

       for _ in range(num_layers):
           normalization1 = norm((dim,))
           normalization2 = norm((dim,))
           attention = Attention(dim=dim,
                                 num_heads=num_heads,
                                 keep_prob=keep_prob,
                                 attention_keep_prob=attention_keep_prob)

           feedforward = FeedForward(in_features=dim,
                                     hidden_features=mlp_dim,
                                     activation=activation,
                                     keep_prob=keep_prob)

           layers.append(
               nn.SequentialCell([
                   ResidualCell(nn.SequentialCell([normalization1, attention])),
                   ResidualCell(nn.SequentialCell([normalization2, feedforward]))
               ])
           )
       self.layers = nn.SequentialCell(layers)

   def construct(self, x):
       """Transformer construct."""
       return self.layers(x)

构建一个完整的Vit模型

from mindspore.common.initializer import Normal
from mindspore.common.initializer import initializer
from mindspore import Parameter


def init(init_type, shape, dtype, name, requires_grad):
   """Init."""
   initial = initializer(init_type, shape, dtype).init_data()
   return Parameter(initial, name=name, requires_grad=requires_grad)


class ViT(nn.Cell):
   def __init__(self,
                image_size: int = 224,
                input_channels: int = 3,
                patch_size: int = 16,
                embed_dim: int = 768,
                num_layers: int = 12,
                num_heads: int = 12,
                mlp_dim: int = 3072,
                keep_prob: float = 1.0,
                attention_keep_prob: float = 1.0,
                drop_path_keep_prob: float = 1.0,
                activation: nn.Cell = nn.GELU,
                norm: Optional[nn.Cell] = nn.LayerNorm,
                pool: str = 'cls') -> None:
       super(ViT, self).__init__()

       self.patch_embedding = PatchEmbedding(image_size=image_size,
                                             patch_size=patch_size,
                                             embed_dim=embed_dim,
                                             input_channels=input_channels)
       num_patches = self.patch_embedding.num_patches

       self.cls_token = init(init_type=Normal(sigma=1.0),
                             shape=(1, 1, embed_dim),
                             dtype=ms.float32,
                             name='cls',
                             requires_grad=True)

       self.pos_embedding = init(init_type=Normal(sigma=1.0),
                                 shape=(1, num_patches + 1, embed_dim),
                                 dtype=ms.float32,
                                 name='pos_embedding',
                                 requires_grad=True)

       self.pool = pool
       self.pos_dropout = nn.Dropout(p=1.0-keep_prob)
       self.norm = norm((embed_dim,))
       self.transformer = TransformerEncoder(dim=embed_dim,
                                             num_layers=num_layers,
                                             num_heads=num_heads,
                                             mlp_dim=mlp_dim,
                                             keep_prob=keep_prob,
                                             attention_keep_prob=attention_keep_prob,
                                             drop_path_keep_prob=drop_path_keep_prob,
                                             activation=activation,
                                             norm=norm)
       self.dropout = nn.Dropout(p=1.0-keep_prob)
       self.dense = nn.Dense(embed_dim, num_classes)

   def construct(self, x):
       """ViT construct."""
       x = self.patch_embedding(x)
       cls_tokens = ops.tile(self.cls_token.astype(x.dtype), (x.shape[0], 1, 1))
       x = ops.concat((cls_tokens, x), axis=1)
       x += self.pos_embedding

       x = self.pos_dropout(x)
       x = self.transformer(x)
       x = self.norm(x)
       x = x[:, 0]
       if self.training:
           x = self.dropout(x)
       x = self.dense(x)

       return x

模型训练

from mindspore.nn import LossBase
from mindspore.train import LossMonitor, TimeMonitor, CheckpointConfig, ModelCheckpoint
from mindspore import train

# define super parameter
epoch_size = 10
momentum = 0.9
num_classes = 1000
resize = 224
step_size = dataset_train.get_dataset_size()

# construct model
network = ViT()

# load ckpt
vit_url = "https://download.mindspore.cn/vision/classification/vit_b_16_224.ckpt"
path = "./ckpt/vit_b_16_224.ckpt"

vit_path = download(vit_url, path, replace=True)
param_dict = ms.load_checkpoint(vit_path)
ms.load_param_into_net(network, param_dict)

# define learning rate
lr = nn.cosine_decay_lr(min_lr=float(0),
                       max_lr=0.00005,
                       total_step=epoch_size * step_size,
                       step_per_epoch=step_size,
                       decay_epoch=10)

# define optimizer
network_opt = nn.Adam(network.trainable_params(), lr, momentum)


# define loss function
class CrossEntropySmooth(LossBase):
   """CrossEntropy."""

   def __init__(self, sparse=True, reduction='mean', smooth_factor=0., num_classes=1000):
       super(CrossEntropySmooth, self).__init__()
       self.onehot = ops.OneHot()
       self.sparse = sparse
       self.on_value = ms.Tensor(1.0 - smooth_factor, ms.float32)
       self.off_value = ms.Tensor(1.0 * smooth_factor / (num_classes - 1), ms.float32)
       self.ce = nn.SoftmaxCrossEntropyWithLogits(reduction=reduction)

   def construct(self, logit, label):
       if self.sparse:
           label = self.onehot(label, ops.shape(logit)[1], self.on_value, self.off_value)
       loss = self.ce(logit, label)
       return loss


network_loss = CrossEntropySmooth(sparse=True,
                                 reduction="mean",
                                 smooth_factor=0.1,
                                 num_classes=num_classes)

# set checkpoint
ckpt_config = CheckpointConfig(save_checkpoint_steps=step_size, keep_checkpoint_max=100)
ckpt_callback = ModelCheckpoint(prefix='vit_b_16', directory='./ViT', config=ckpt_config)

# initialize model
# "Ascend + mixed precision" can improve performance
ascend_target = (ms.get_context("device_target") == "Ascend")
if ascend_target:
   model = train.Model(network, loss_fn=network_loss, optimizer=network_opt, metrics={"acc"}, amp_level="O2")
else:
   model = train.Model(network, loss_fn=network_loss, optimizer=network_opt, metrics={"acc"}, amp_level="O0")

# train model
model.train(epoch_size,
           dataset_train,
           callbacks=[ckpt_callback, LossMonitor(125), TimeMonitor(125)],
           dataset_sink_mode=False,)

true

模型验证

dataset_val = ImageFolderDataset(os.path.join(data_path, "val"), shuffle=True)

trans_val = [
   transforms.Decode(),
   transforms.Resize(224 + 32),
   transforms.CenterCrop(224),
   transforms.Normalize(mean=mean, std=std),
   transforms.HWC2CHW()
]

dataset_val = dataset_val.map(operations=trans_val, input_columns=["image"])
dataset_val = dataset_val.batch(batch_size=16, drop_remainder=True)

# construct model
network = ViT()

# load ckpt
param_dict = ms.load_checkpoint(vit_path)
ms.load_param_into_net(network, param_dict)

network_loss = CrossEntropySmooth(sparse=True,
                                 reduction="mean",
                                 smooth_factor=0.1,
                                 num_classes=num_classes)

# define metric
eval_metrics = {'Top_1_Accuracy': train.Top1CategoricalAccuracy(),
               'Top_5_Accuracy': train.Top5CategoricalAccuracy()}

if ascend_target:
   model = train.Model(network, loss_fn=network_loss, optimizer=network_opt, metrics=eval_metrics, amp_level="O2")
else:
   model = train.Model(network, loss_fn=network_loss, optimizer=network_opt, metrics=eval_metrics, amp_level="O0")

# evaluate model
result = model.eval(dataset_val)
print(result)

模型推理

dataset_infer = ImageFolderDataset(os.path.join(data_path, "infer"), shuffle=True)

trans_infer = [
   transforms.Decode(),
   transforms.Resize([224, 224]),
   transforms.Normalize(mean=mean, std=std),
   transforms.HWC2CHW()
]

dataset_infer = dataset_infer.map(operations=trans_infer,
                                 input_columns=["image"],
                                 num_parallel_workers=1)
dataset_infer = dataset_infer.batch(1)

调用predict方法进行推理

import os
import pathlib
import cv2
import numpy as np
from PIL import Image
from enum import Enum
from scipy import io


class Color(Enum):
   """dedine enum color."""
   red = (0, 0, 255)
   green = (0, 255, 0)
   blue = (255, 0, 0)
   cyan = (255, 255, 0)
   yellow = (0, 255, 255)
   magenta = (255, 0, 255)
   white = (255, 255, 255)
   black = (0, 0, 0)


def check_file_exist(file_name: str):
   """check_file_exist."""
   if not os.path.isfile(file_name):
       raise FileNotFoundError(f"File `{file_name}` does not exist.")


def color_val(color):
   """color_val."""
   if isinstance(color, str):
       return Color[color].value
   if isinstance(color, Color):
       return color.value
   if isinstance(color, tuple):
       assert len(color) == 3
       for channel in color:
           assert 0 <= channel <= 255
       return color
   if isinstance(color, int):
       assert 0 <= color <= 255
       return color, color, color
   if isinstance(color, np.ndarray):
       assert color.ndim == 1 and color.size == 3
       assert np.all((color >= 0) & (color <= 255))
       color = color.astype(np.uint8)
       return tuple(color)
   raise TypeError(f'Invalid type for color: {type(color)}')


def imread(image, mode=None):
   """imread."""
   if isinstance(image, pathlib.Path):
       image = str(image)

   if isinstance(image, np.ndarray):
       pass
   elif isinstance(image, str):
       check_file_exist(image)
       image = Image.open(image)
       if mode:
           image = np.array(image.convert(mode))
   else:
       raise TypeError("Image must be a `ndarray`, `str` or Path object.")

   return image


def imwrite(image, image_path, auto_mkdir=True):
   """imwrite."""
   if auto_mkdir:
       dir_name = os.path.abspath(os.path.dirname(image_path))
       if dir_name != '':
           dir_name = os.path.expanduser(dir_name)
           os.makedirs(dir_name, mode=777, exist_ok=True)

   image = Image.fromarray(image)
   image.save(image_path)


def imshow(img, win_name='', wait_time=0):
   """imshow"""
   cv2.imshow(win_name, imread(img))
   if wait_time == 0:  # prevent from hanging if windows was closed
       while True:
           ret = cv2.waitKey(1)

           closed = cv2.getWindowProperty(win_name, cv2.WND_PROP_VISIBLE) < 1
           # if user closed window or if some key pressed
           if closed or ret != -1:
               break
   else:
       ret = cv2.waitKey(wait_time)


def show_result(img: str,
               result: Dict[int, float],
               text_color: str = 'green',
               font_scale: float = 0.5,
               row_width: int = 20,
               show: bool = False,
               win_name: str = '',
               wait_time: int = 0,
               out_file: Optional[str] = None) -> None:
   """Mark the prediction results on the picture."""
   img = imread(img, mode="RGB")
   img = img.copy()
   x, y = 0, row_width
   text_color = color_val(text_color)
   for k, v in result.items():
       if isinstance(v, float):
           v = f'{v:.2f}'
       label_text = f'{k}: {v}'
       cv2.putText(img, label_text, (x, y), cv2.FONT_HERSHEY_COMPLEX,
                   font_scale, text_color)
       y += row_width
   if out_file:
       show = False
       imwrite(img, out_file)

   if show:
       imshow(img, win_name, wait_time)


def index2label():
   """Dictionary output for image numbers and categories of the ImageNet dataset."""
   metafile = os.path.join(data_path, "ILSVRC2012_devkit_t12/data/meta.mat")
   meta = io.loadmat(metafile, squeeze_me=True)['synsets']

   nums_children = list(zip(*meta))[4]
   meta = [meta[idx] for idx, num_children in enumerate(nums_children) if num_children == 0]

   _, wnids, classes = list(zip(*meta))[:3]
   clssname = [tuple(clss.split(', ')) for clss in classes]
   wnid2class = {wnid: clss for wnid, clss in zip(wnids, clssname)}
   wind2class_name = sorted(wnid2class.items(), key=lambda x: x[0])

   mapping = {}
   for index, (_, class_name) in enumerate(wind2class_name):
       mapping[index] = class_name[0]
   return mapping


# Read data for inference
for i, image in enumerate(dataset_infer.create_dict_iterator(output_numpy=True)):
   image = image["image"]
   image = ms.Tensor(image)
   prob = model.predict(image)
   label = np.argmax(prob.asnumpy(), axis=1)
   mapping = index2label()
   output = {int(label): mapping[int(label)]}
   print(output)
   show_result(img="./dataset/infer/n01440764/ILSVRC2012_test_00000279.JPEG",
               result=output,
               out_file="./dataset/infer/ILSVRC2012_test_00000279.JPEG")

true

本帖最后由 匿名用户2025/04/30 14:01:30 编辑

我要发帖子