• CogView中的Transformer


    入门小菜鸟,希望像做笔记记录自己学的东西,也希望能帮助到同样入门的人,更希望大佬们帮忙纠错啦~侵权立删。

    目录

    一、原理

    1、总体介绍

    2、具体实现

    (1)不采取稀疏处理(默认)

    (2)采取稀疏训练

    ​(3)稀疏推断

    二、代码解析

    1、__init__

    (1)参数设定

    (2)存储激活检查点标志

    (3)定义输出层初始化方法

    (4)Position embedding

    (5)窗口定义

    (6)Transformer layers设置

    (7)将 num_layer 个 transformer layer打包在一起,以列表形式保存

    (8)output层的LayerNorm处理

    (9)激活点检查

    2、forward

    (1)获取最终的输入层的相关信息

    (2)attention mask建立

    (3)稀疏训练or推断准备

    (4)对输入层的处理

    (5)这次是否有产生记忆模块

    (6)获取下一层的输入——分为是否采取检查点激活两种情况来分析

    (7)最后一层norm

    (8)记忆模块更新

    (9)返回这一层的输出结果和记忆模块


    一、原理

    1、总体介绍

    将n个的 transformer blocks 打包在一起,即 n * transformer layer + final layernorm 两部分组成

    2、具体实现

    (1)不采取稀疏处理(默认)

     (2)采取稀疏训练

     新建的rmask(k为输入的总列数;w为窗口大小;t为调整窗口数量所用)

     (3)稀疏推断


    二、代码解析

    1、__init__

    (1)参数设定

    1. class GPT2ParallelTransformer(torch.nn.Module):
    2. """GPT-2 transformer.
    3. This module takes input from embedding layer and it's output can
    4. be used directly by a logit layer. It consists of L (num-layers)
    5. blocks of:
    6. layer norm
    7. self attention
    8. residual connection
    9. layer norm
    10. mlp
    11. residual connection
    12. followed by a final layer norm.
    13. Arguments:
    14. num_layers: Number of transformer layers.
    15. hidden_size: The hidden size of the self attention.
    16. num_attention_heads: number of attention head in the self
    17. attention.
    18. attention_dropout_prob: dropout probability of the attention
    19. score in self attention.
    20. output_dropout_prob: dropout probability for the outputs
    21. after self attention and final output.
    22. checkpoint_activations: if True, checkpoint activations.
    23. checkpoint_num_layers: number of layers to checkpoint. This
    24. is basically the chunk size in checkpoitning.
    25. layernorm_epsilon: epsilon used in layernorm to avoid
    26. division by zero.
    27. init_method_std: standard deviation of the init method which has
    28. the form N(0, std).
    29. use_scaled_init_for_output_weights: If Ture use 1/sqrt(2*num_layers)
    30. scaling for the output weights (
    31. output of self attention and mlp).
    32. """
    33. def __init__(self,
    34. num_layers,
    35. hidden_size,
    36. num_attention_heads,
    37. max_sequence_length,
    38. max_memory_length,
    39. embedding_dropout_prob,
    40. attention_dropout_prob,
    41. output_dropout_prob,
    42. checkpoint_activations,
    43. checkpoint_num_layers=1,
    44. layernorm_epsilon=1.0e-5,
    45. init_method_std=0.02,
    46. use_scaled_init_for_output_weights=True,
    47. query_window=128,
    48. key_window_times=6,
    49. num_pivot=768
    50. ):
    51. super(GPT2ParallelTransformer, self).__init__()
    • num_layers:transformer层的数量;
    • hidden_size:自我注意力模块的隐藏大小(嵌入向量的维度);
    • num_attention_heads:自我注意力模块中attention head的数量;
    • max_sequence_length:词典大小;
    • max_memory_length:最大记忆长度;
    • embedding_dropout_prob:嵌入层(该模块的输入部分)中元素被dropout的概率(为了解决过拟合问题而随机丢弃一部分元素);
    • attention_dropout_prob:同样道理,注意力模块中注意力得分被dropout的概率;
    • output_dropout_prob:同理,输出层后的输出被dropout的概率;
    • checkpoint_activations:是否执行检查点激活;
    • checkpoint_num_layers:检查点的层数。这基本上是checkpoitning中的块大小;
    • layernorm_epsilon:在layernform中用于避免被零除的ε(用于防止分母为0);
    • init_method_std:初始化方法(使用让权重呈现正态分布的方法)中正态分布的方差;
    • use_scaled_init_for_output_weights:是否对自注意力和mlp的输出的权重调用scaled_init_method进行初始化;
    • query_window:稀疏处理中的窗口大小;
    • key_window_times:用于调整窗口数量;
    • num_pivot:transformer里图像token和文本token的总和数量

    (2)存储激活检查点标志

    1. # Store activation checkpoiting flag.
    2. #首先先记录是否执行检查点激活,检查点的层数,最大记忆长度和最大序列长度信息
    3. self.checkpoint_activations = checkpoint_activations
    4. self.checkpoint_num_layers = checkpoint_num_layers
    5. self.max_memory_length = max_memory_length
    6. self.max_sequence_length = max_sequence_length

    (3)定义输出层初始化方法

    由use_scaled_init_for_output_weights决定,若为False则不进行初始化缩放,若为true则调用scaled_init_method进行初始化

    1. #输出层初始化方法定义——由use_scaled_init_for_output_weights决定,若为False则不进行初始化缩放,若为true则调用scaled_init_method进行初始化
    2. output_layer_init_method = None
    3. if use_scaled_init_for_output_weights:
    4. output_layer_init_method = scaled_init_method(init_method_std,
    5. num_layers)

    scaled_init_method函数——返回初始化方法:初始权重呈均值为0,方差为init_method_std//sqrt(2*num_layers)的正态分布

    1. def scaled_init_method(sigma, num_layers):
    2. """Init method based on N(0, sigma/sqrt(2*num_layers)."""
    3. std = sigma / math.sqrt(2.0 * num_layers)
    4. def init_(tensor):
    5. return torch.nn.init.normal_(tensor, mean=0.0, std=std)
    6. return init_

    (4)Position embedding

    先进行嵌入层的dropout(防止过拟合),然后调用torch.nn.Embedding()方法按词典大小max_sequence_length和嵌入向量的维度hidden_size来定义词向量格式,然后将词向量的值初始化为呈以0为均值,以init_method_std为方差的正态分布。

    1. # Embeddings dropout嵌入层dropout
    2. self.embedding_dropout = torch.nn.Dropout(embedding_dropout_prob)
    3. # Position embedding (serial).初始化含位置信息的词向量方法
    4. self.position_embeddings = torch.nn.Embedding(max_sequence_length,
    5. hidden_size)#随机以max_sequence_length为词典的大小(词的个数),以hidden_size来嵌入向量的维度(即用多少维来表示一个符号)初始化词向量,默认词向量值在正态分布N(0,1)中随机取值
    6. # Initialize the position embeddings.词向量值在正态分布N(0,init_method_std)中随机取值
    7. torch.nn.init.normal_(self.position_embeddings.weight, mean=0.0, std=init_method_std)

    (5)窗口定义

    1. self.query_window = query_window
    2. self.key_window_times = key_window_times
    3. self.num_pivot = num_pivot

    (6)Transformer layers设置

    首先定义了一个get_layer()函数来获得对应层id的网络层(transformer layer)

    1. #获得对应层id的网络层
    2. def get_layer(layer_id):
    3. return GPT2ParallelTransformerLayer(
    4. hidden_size,
    5. num_attention_heads,
    6. attention_dropout_prob,
    7. output_dropout_prob,
    8. layernorm_epsilon,
    9. unscaled_init_method(init_method_std),
    10. output_layer_init_method=output_layer_init_method,
    11. query_window=query_window,
    12. key_window_times=key_window_times,
    13. scale_normalization=True
    14. )

    这里调用了GPT2ParallelTransformerLayer类

    (7)将 num_layer 个 transformer layer打包在一起,以列表形式保存

    1. # Transformer layers.
    2. self.layers = torch.nn.ModuleList(
    3. [get_layer(layer_id) for layer_id in range(num_layers)])

    (8)output层的LayerNorm处理

    1. # Final layer norm before output.
    2. self.final_layernorm = LayerNorm(hidden_size, eps=layernorm_epsilon)

    (9)激活点检查

    1. if deepspeed.checkpointing.is_configured():
    2. global get_cuda_rng_tracker, checkpoint
    3. get_cuda_rng_tracker = deepspeed.checkpointing.get_cuda_rng_tracker
    4. checkpoint = deepspeed.checkpointing.checkpoint
    5. self.rmask = None#是否进行稀疏处理

    2、forward

    1. def forward(self, hidden_states, position_ids, attention_mask, txt_indices_bool, img_indices_bool, is_sparse=0, *mems):
    2. '''''
    3. hidden_states:输入的网络层;
    4. position_ids:位置编码;
    5. attention_mask;
    6. txt_indices_bool:选取文本token有效的索引矩阵
    7. img_indices_bool:选取图像token有效的索引矩阵
    8. is_sparse:是否稀疏处理,稀疏训练,稀疏推断
    9. mems:记忆模块;
    10. '''''

    (1)获取最终的输入层的相关信息

    获取b,s和最终的输入列数(hidden_states和记忆模块的concat的结果)

    1. batch_size, query_length = hidden_states.size()[:2]#获取batchsize(b)和读取的序列长度(s)
    2. memory_length = mems[0].size(1) if mems else 0#获取记忆模块的序列长度(模块列数)
    3. key_length = query_length + memory_length#得到最终的序列长度(类似concat维数增加)

    (2)attention mask建立

    最终shape[1,1,s,s](无记忆模块情况下,有记忆为[1,1,s,s+m],m为memory_length)

    1. # conventional transformer
    2. #建立常规transformer的attention mask
    3. def build_mask_matrix(query_length, key_length, sep):
    4. m = torch.ones((1, query_length, key_length), device=hidden_states.device, dtype=hidden_states.dtype)#初始化为全一矩阵
    5. assert query_length <= key_length
    6. m[0, :, -query_length:] = torch.tril(m[0, :, -query_length:])#返回m[0, :, -query_length:]区域(最后两维)是下三角矩阵的矩阵
    7. m[0, :, :sep + (key_length - query_length)] = 1#注意力标记
    8. m = m.unsqueeze(1)#[1,s,s+m]->[1,1,s,s+m]
    9. return m
    10. #生成attention_mask,无记忆模块是[1,1,s,s],有记忆是[1,1,s,s+m]
    11. attention_mask = build_mask_matrix(query_length, key_length, sep)

    (3)稀疏训练or推断准备

    ✨获取稀疏训练的rmask

    1. #启用稀疏训练生成rmask
    2. if is_sparse == 1 and (self.rmask is None):
    3. w, times = self.query_window, self.key_window_times#滑动窗口大小+窗口数的减少量获取
    4. g = key_length // w#获取全局attention窗口个数
    5. tmp = torch.ones((g-times+1, w , w), device=hidden_states.device, dtype=hidden_states.dtype)#初始化rmask(可理解为g-times+1个窗口)
    6. tmp = torch.tril(1 - torch.block_diag(*tmp))#*将三维矩阵变成二维矩阵列表;torch.block_diag将g-times+1个w*w矩阵组合成一个块对角矩阵,1-使得中间块为0,其余为1;torch.tril返回下三角矩阵。shape为((g-times+1)*w,(g-times+1)*w)
    7. self.rmask = torch.nn.functional.pad(tmp, (0, (times-1)*w, (times-1)*w, 0)) # pad (left, right, top, bottom),这四个元素的位置代表了填充的位置,大小为填充的行数,默认填0,所以最终shape为(g*w,g*w),左下角为一个((g-times+1)*w,(g-times+1)*w)大小的下三角矩阵

    ✨获取左边界和支点

    1. if is_sparse == 2:#稀疏推断
    2. left_boundary = max(0, key_length - self.key_window_times * self.query_window)#获取左边界(将key_length分为n份query_window的块块,做除法后的余数部分为左边界
    3. window_idx = torch.arange(left_boundary, key_length, device=hidden_states.device, dtype=torch.long).expand(batch_size, -1)#torch.arange获得[left_boundary,...,key_length-1];expand(batch_size, -1)获得batchsize条[left_boundary,...,key_length-1],获得shape为(batchsize*key_length-left_boundary)
    4. elif is_sparse == 1:#稀疏训练
    5. left_boundary = key_length#获取左边界
    6. num_pivot = self.num_pivot#transformer里图像token和文本token的总和数量获取

    ✨选取每个batch中对应有效的index的image token和txt token

    1. #选取每个batch中对应有效的index的image token和txt token
    2. if is_sparse: # 1 or 2
    3. # select out the real indices for sampling
    4. img_indices = [img_indices_bool[i][:left_boundary].nonzero(as_tuple=False).view(-1) for i in range(batch_size)]#.nonzero(as_tuple=False)取出非0元素的索引(即取出有效索引);.view(-1)将其展平
    5. txt_indices = [txt_indices_bool[i][:left_boundary].nonzero(as_tuple=False).view(-1) for i in range(batch_size)]

    ✨稀疏推断支点数目设定

    1. #稀疏推断支点数目设定(总token数量增加)
    2. if is_sparse == 2:
    3. ratio = self.num_pivot / self.max_sequence_length#支点比例获取
    4. max_text_num = max(len(text_idx) for text_idx in txt_indices)#获取batch中最长的有效文本token长度
    5. num_pivot = max_text_num + int((left_boundary - max_text_num) * ratio)#支点数目更新

    (4)对输入层的处理

    给输入层加入初始化的位置信息词向量并且进行dropout操作

    1. #对输入层的处理
    2. position_embeddings = self.position_embeddings(position_ids)#对位置信息position_ids进行词向量的初始化
    3. hidden_states = hidden_states + position_embeddings#输入层加入初始化的位置信息词向量
    4. hidden_states = self.embedding_dropout(hidden_states)#对输入层进行dropout

    (5)这次是否有产生记忆模块

    若拥有最大记忆长度,则产生的记忆模块是输入层,但不需要计算其梯度

    1. #这次是否有产生记忆模块
    2. if self.max_memory_length > 0:#若拥有最大记忆长度,
    3. mem_layers = [hidden_states.detach()]#记忆模块赋为输入层,但不需要计算其梯度
    4. else:#否则没有记忆模块
    5. mem_layers = []

    然后保存一下attention mask

            attention_mask_saved = attention_mask#保存attention mask

    (6)获取下一层的输入——分为是否采取检查点激活两种情况来分析

    (都要利用get_layer来实现,所以都要先获取相应的参数输入才可调用)

    ✨采取检查点激活

    ①首先是必要的初始化和参数获取

    1. l = 0#初始化start层id
    2. num_layers = len(self.layers)#Transformer layers的数量获取
    3. chunk_length = self.checkpoint_num_layers#检查点的层数

    循环获取层

                while l < num_layers:

    ②稀疏训练or推断情况下获取下一层的输入的参数

                    if is_sparse > 0:#稀疏训练or推断

    🌳获取pivot的索引(pivot即随机抽取的token,用于代表全局整幅图片)

    1. # ===================== Pivot Mask ======================== #
    2. pivot_idx = torch.stack([
    3. torch.cat((
    4. text_idx,
    5. img_indices[i][
    6. torch.tensor(random.sample(range(len(img_indices[i])), k=num_pivot - len(text_idx)), dtype=torch.long, device=text_idx.device)
    7. ]
    8. ), dim=0)
    9. for i, text_idx in enumerate(txt_indices)
    10. ])
    11. #首先由random.sample随机抽取(预设支点数量-该batch的有效文本token长度)=该batch的有效图像token长度个图像token索引,并且将文本token和图像token拼接在一起

    🌳然后对于稀疏训练:获取pivot_attention_mask,进而获取输入所需的参数列表

    1. if is_sparse == 1: # sparse training
    2. assert key_length == query_length#断言最终的序列长度和读取的序列长度(s)是否相同
    3. b, s = batch_size, key_length
    4. pivot_attention_mask = self.rmask.expand(b, s, s).gather(dim=-1, index=pivot_idx.unsqueeze(1).expand(b, s, self.num_pivot))#生成针对随机选取的token的注意力矩阵——pivot attention mask
    5. #expand()函数扩展维度,其余不变。
    6. # 相当于先由b个原来的s*s(s=g*w)大小的rmask(即每个batch里都有rmask)拼成一个大小为(b,s,s)的矩阵;再由gather函数根据 index 参数(即是索引)返回矩阵里面对应位置的值(即挑出随机选中的token对应索引值的rmask值)——针对的是每个batch的s*s的rmask;最后再由expand函数展成大小为(b,s,随机选取的token数量)的矩阵
    7. args = [hidden_states, pivot_attention_mask, pivot_idx, torch.tensor(is_sparse)]#参数列表记录

    🌳然后对于稀疏推理:获取全部需要注意的token的idx,并形成参数列表

    1. elif is_sparse == 2: # sparse inference
    2. pw_idx = torch.cat((pivot_idx, window_idx), dim=-1)#获取随机选取的token的idx矩阵与额外标记注意的窗口的idx矩阵concat后的需要attention的idx矩阵
    3. args = [hidden_states, attention_mask_saved, pw_idx, torch.tensor(is_sparse)]#参数列表记录

    🌳错误提示

    1. else:
    2. raise NotImplementedError

    ③非稀疏处理情况下参数列表获取

    1. else:
    2. args = [hidden_states, attention_mask_saved]#非稀疏处理的参数列表记录(输入层和attention mask)

    ④记忆模块对参数列表的补充

    1. #对于记忆模块的参数补充
    2. if mems:
    3. args += mems[l: l + chunk_length]

    ⑤获得下一层的输入并进行检查,且start层idx(l)更新

    1. #检查点激活并得到下一层的输入层
    2. hidden_states = checkpoint(custom(l, l + chunk_length), *args)
    3. #start为第l层,end为第l + chunk_length层(共检查点层数数量)
    4. l += chunk_length#下一个检查点的开始层数

    这里调用custom函数——用于获取下一层的输入

    1. def custom(start, end):
    2. def custom_forward(*inputs):
    3. layers_ = self.layers[start:end]#获取对应的层序列
    4. x_, inputs = inputs[0], inputs[1:]#将他们分成两份(头和其余)
    5. if is_sparse > 0:#稀疏处理
    6. inputs, mems_ = inputs[:3], inputs[3:]#输入为前3层,其余为记忆模块
    7. else:#不采取稀疏处理
    8. inputs, mems_ = inputs[:1], inputs[1:]#输入为第1层,其余为记忆模块
    9. for i, layer in enumerate(layers_):
    10. mem_i_ = mems_[i] if mems_ else None#获取第i层的记忆模块
    11. x_ = layer(x_, *inputs, mem=mem_i_)#调用get_layer中GPT2ParallelTransformerLayer的forward——x_对应hidden_states(输入), inputs对应ltor_mask(attention mask)
    12. if self.max_memory_length > 0:
    13. mem_layers.append(x_.detach())#记忆模块添加(不参与梯度计算)
    14. return x_
    15. return custom_forward

    ✨不采取检查点激活

    思路和上面的检查点激活类似,只是不考虑了检查点层数和checkpoint

    1. else:#不进行检查点激活
    2. assert is_sparse != 1, 'Please use checkpoint_activations for sparse attention training.'
    3. for i, layer in enumerate(self.layers):#遍历Transformer layers
    4. if is_sparse == 0:#非稀疏处理——获取下一步传入的参数列表
    5. args = [hidden_states, attention_mask_saved]
    6. elif is_sparse == 2:#稀疏推断
    7. pivot_idx = torch.stack([
    8. torch.cat((
    9. text_idx,
    10. img_indices[i][
    11. torch.tensor(random.sample(range(len(img_indices[i])), k=num_pivot - len(text_idx)), dtype=torch.long, device=text_idx.device)
    12. ]
    13. ), dim=0)
    14. for i, text_idx in enumerate(txt_indices)
    15. ])#首先由random.sample随机抽取(预设支点数量-该batch的有效文本token长度)=该batch的有效图像token长度个图像token索引,并且将文本token和图像token拼接在一起
    16. pw_idx = torch.cat((pivot_idx, window_idx), dim=-1)#获取随机选取的token的idx矩阵与额外标记注意的窗口的idx矩阵concat后的需要attention的idx矩阵
    17. args = [hidden_states, attention_mask_saved, pw_idx, torch.tensor(is_sparse)]#参数列表记录
    18. mem_i = mems[i] if mems else None#对应层的记忆模块
    19. hidden_states = layer(*args, mem=mem_i)#下一层的输入层获取
    20. if self.max_memory_length > 0:#记忆层添加
    21. mem_layers.append(hidden_states.detach())

    (7)最后一层norm

    作Layernorm操作

    1. # Final layer norm.
    2. output = self.final_layernorm(hidden_states)#即对下一层的输入(这一层的输出)做一个LayerNorm规范

    (8)记忆模块更新

    1. #更新记忆模块
    2. if self.max_memory_length > 0:
    3. mem_layers = self.update_mems(mem_layers, mems)

    这里调用update_mems进行更新

    1. def update_mems(self, hiddens, mems):
    2. memory_length = mems[0].size(1) if mems else 0#原记忆模块的长度(列数)
    3. query_length = hiddens[0].size(1)#新待加入的记忆模块长度(列数)
    4. new_memory_length = min(self.max_memory_length, memory_length + query_length)#新的记忆模块的长度确定
    5. new_mems = []
    6. with torch.no_grad():
    7. for i in range(len(hiddens)):
    8. if new_memory_length <= query_length:#说明选中的是self.max_memory_length(记忆模块完全为新的记忆层组成)。取每一层的每一行的后new_memory_length组成新的记忆矩阵
    9. new_mems.append(hiddens[i][:, -new_memory_length:])
    10. else:#说明选中的是memory_length + query_length。取原来的记忆模块和新加入的进行拼接(沿列拼接)
    11. new_mems.append(torch.cat((mems[i][:, -new_memory_length+query_length:], hiddens[i]), dim=1))
    12. return new_mems

    (9)返回这一层的输出结果和记忆模块

            return (output, *mem_layers)#返回下一层的输入(这一层的输出结果)和记忆模块

    欢迎大家在评论区批评指正,谢谢~

  • 相关阅读:
    LNMP架构之搭建Discuz论坛
    SSM整合(四)
    Map,List,Set 等集合以及底层数据结构
    0动态规划+二分查找中等 LeetCode2008. 出租车的最大盈利
    技术为业务赋能:深度剖析开发与业务的紧密结合
    在Ubuntu Linux Desktop上构建matter开发环境
    chatgpt赋能python:Python如何快速取出所有元素?
    【算法】快速排序
    Java OpenJDK 8u345 Windows Installer
    easy_see
  • 原文地址:https://blog.csdn.net/weixin_55073640/article/details/126531746