Python生成无连续False值的随机掩码数组 实现文本单词均匀掩码
非连续随机掩码优化实现
需求背景
你需要对给定的单词列表做随机掩码处理,要求掩码占比约为设定值,且不会出现两个连续被掩码的单词(即不会输出连续的_)。
前置校验
首先要确认掩码数量的合理性:对于长度为LEN的单词列表,最多可掩码的单词数为(LEN + 1) // 2,如果按照maskrate计算出来的待掩码数量nbmask超过这个上限,需要自动调整到最大值,否则无法满足无连续掩码的要求。
优化后代码实现
import numpy as np # 待处理单词列表 words = ['First', 'Citizen:', 'Before', 'we', 'proceed', 'any', 'further,', 'hear', 'me', 'speak.', 'All:', 'Speak,', 'speak.', 'First', 'Citizen:', 'You', 'are', 'all', 'resolved', 'rather', 'to', 'die', 'than', 'to', 'famish?', 'All:', 'Resolved.', 'resolved.', 'First', 'Citizen:', 'First,', 'you', 'know', 'Caius', 'Marcius', 'is', 'chief', 'enemy', 'to', 'the', 'people.'] LEN = len(words) maskrate = 0.2 nbmask = int(np.floor(LEN * maskrate)) # 校验掩码数量上限,避免逻辑异常 max_nbmask = (LEN + 1) // 2 if nbmask > max_nbmask: nbmask = max_nbmask # 生成无相邻的掩码数组 mask = np.ones(LEN, dtype=bool) selected_pos = set() # 随机采样不相邻的掩码位置 while len(selected_pos) < nbmask: pos = np.random.randint(0, LEN) # 校验当前位置和前后位置都未被选为掩码位 if pos not in selected_pos and pos-1 not in selected_pos and pos+1 not in selected_pos: selected_pos.add(pos) # 标记掩码位 for pos in selected_pos: mask[pos] = False # 生成最终掩码后的文本 masked_words = [] for word, is_keep in zip(words, mask): masked_words.append(word if is_keep else '_') print("掩码数组:", mask) print("掩码后结果:", masked_words)
实现说明
- 核心逻辑是随机采样掩码位置时,同时校验该位置的前后一位都没有被选中为掩码位,从根源避免连续掩码的情况
- 加入了掩码数量的边界校验,避免因为设定的掩码占比过高导致逻辑死循环
- 最终生成的掩码分布完全随机,不会有固定的间隔规律,满足均匀分布的要求
内容的提问来源于stack exchange,提问作者kiriloff
相关产品推荐
相关产品推荐

