Bert预训练掩码(MLM)
·
策略总览
遮盖策略
| 策略 | 本代码实现 | 标准BERT |
|---|---|---|
| 遮蔽概率 | 15% | 15% |
| 遮蔽中[MASK]比例 | 80% | 80% |
| 遮蔽中保留原token比例 | 10% | 10% |
| 遮蔽中随机替换比例 | 10% | 10% |
| 随机词范围 | 排除特殊标记 | 同左 |
| 标签处理 | 非遮蔽token设为-100 | 同左 |
input和output处理
output中未遮掩部分为-100,表示不参与;遮掩部分输出正确标签;
由此学习语句特征

代码实现
第一部分实现(确定掩码范围)
-
随机生成0-1的随机数,这里可以看作被遮盖的概率
rands = np.random.random(len(text_ids)) -
while遍历每个token
-
随机选择连续遮盖数
ngram = np.random.choice([1,2,3], p=[0.7,0.2,0.1]) -
通过强制小于0.15进行遮盖
while L < R and L < len(rands): rands[L] = np.random.random() * 0.15 # 强制设为小于0.15的随机数 L += 1 -
不连续遮盖
idx = R # 移动到遮蔽片段的末尾 if idx < len(rands): rands[idx] = 1 # 禁止下一个token被遮蔽第一部分完整代码如下:
input_ids, output_ids = [], [] #输入、输出对应的ids
rands = np.random.random(len(text_ids)) #动态n-gram遮蔽策略 为每个token生成一个0-1的随机数
idx=0
while idx<len(rands): #遍历每个token
if rands[idx]<0.15: #15%的概率需要mask
ngram=np.random.choice([1,2,3], p=[0.7,0.2,0.1])
if ngram==3 and len(rands)<7:#太大的gram不要应用于过短文本
ngram=2
if ngram==2 and len(rands)<4:
ngram=1
L=idx+1
R=idx+ngram #遮盖范围
while L<R and L<len(rands):
rands[L]=np.random.random()*0.15 #范围内强制mask(强制小于0.15)
L+=1
idx=R
if idx<len(rands):
rands[idx]=1 #禁止mask片段的下一个token被mask,防止连续mask
idx+=1
第二部分代码(掩码实现)
三种方式对应三个if情况
| 遮蔽中[MASK]比例 | 80% |
| 遮蔽中保留原token比例 | 10% |
| 遮蔽中随机替换比例 | 10% |
如果不掩码处理,则输出-100,loss函数将不考虑未遮掩部分(else部分)
for r, i in zip(rands, text_ids): #三种情况对应三种掩码情况
if r < 0.15 * 0.8:
input_ids.append(self.tk.mask_token_id) #输入为[MASK]
output_ids.append(i)
elif r < 0.15 * 0.9:
input_ids.append(i)
output_ids.append(i)#自己预测自己
elif r < 0.15:
input_ids.append(np.random.randint(self.spNum,self.tkNum)) #输入为随机词
output_ids.append(i)
else:
input_ids.append(i) #不遮盖的情况下,out为-100,即不让参与loss
output_ids.append(-100)#保持原样不预测
完整代码
def random_mask(self, text_ids): #不需要看代码 输入是什么? 输出是什么
input_ids, output_ids = [], [] #输入、输出对应的ids
rands = np.random.random(len(text_ids)) #动态n-gram遮蔽策略 为每个token生成一个0-1的随机数
idx=0
while idx<len(rands): #遍历每个token
if rands[idx]<0.15: #15%的概率需要mask
ngram=np.random.choice([1,2,3], p=[0.7,0.2,0.1]) #若要mask,70%单token,20%双token,10%三token
if ngram==3 and len(rands)<7:#太大的gram不要应用于过短文本
ngram=2
if ngram==2 and len(rands)<4:
ngram=1
L=idx+1
R=idx+ngram #遮盖范围
while L<R and L<len(rands):
rands[L]=np.random.random()*0.15 #范围内强制mask(强制小于0.15)
L+=1
idx=R
if idx<len(rands):
rands[idx]=1 #禁止mask片段的下一个token被mask,防止一大片连续mask
idx+=1
for r, i in zip(rands, text_ids): #三种情况对应三种掩码情况
if r < 0.15 * 0.8:
input_ids.append(self.tk.mask_token_id) #输入为[MASK]
output_ids.append(i)
elif r < 0.15 * 0.9:
input_ids.append(i)
output_ids.append(i)#自己预测自己
elif r < 0.15:
input_ids.append(np.random.randint(self.spNum,self.tkNum)) #输入为随机词
output_ids.append(i)
else:
input_ids.append(i) #不遮盖的情况下,out为-100,即不让参与loss
output_ids.append(-100)#保持原样不预测
return input_ids, output_ids
更多推荐



所有评论(0)