Translation Dataset新增boundaries键后配置调用的技术问题
解决TranslationDataset中词边界数据用于SelfAttention掩码的问题
看起来你踩了框架数据加载逻辑的一个常见坑——当你在配置里写data:boundaries时,框架默认会把它当成一个独立的数据集文件(比如boundaries.dev)去加载,而不是你已经嵌入到_data映射里的那个键。添加KeyMap映射只会强化这个误解,让框架更执着地去找对应的文件,自然会报错。
下面是一步步的解决方案,帮你正确获取每行对应的词边界数据,并用它来限制SelfAttention的注意力范围:
1. 确保数据集正确返回包含boundaries的样本字典
首先,检查你的TranslationDataset的__getitem__方法,确保它把boundaries作为样本字典的一个键返回,而不是只存在内部的_data映射里。比如:
class TranslationDataset(Dataset): def __init__(self, src_path, tgt_path, boundary_path): # 加载源数据、目标数据和词边界数据 self.src_data = self.load_data(src_path) self.tgt_data = self.load_data(tgt_path) self.boundaries_data = self.load_boundaries(boundary_path) # 每行对应一个boundaries列表 def __getitem__(self, idx): return { "src": self.src_data[idx], "tgt": self.tgt_data[idx], "boundaries": self.boundaries_data[idx] # 关键:把boundaries放到返回的字典里 } def load_boundaries(self, path): # 自定义加载逻辑,比如每行解析成[(start1, end1), (start2, end2)...]的列表 boundaries = [] with open(path, "r", encoding="utf-8") as f: for line in f: line = line.strip() if not line: boundaries.append([]) continue pairs = [tuple(map(int, p.split(","))) for p in line.split(";")] boundaries.append(pairs) return boundaries
2. 批量处理时正确对齐boundaries数据
如果你的数据加载器用了collate_fn,记得要处理boundaries的批量对齐。比如,因为每个样本的词数量可能不同,你可以用列表保存批量的boundaries,或者用pad的方式统一长度(根据你的模型需求选择):
def collate_fn(batch): src_batch = torch.tensor([item["src"] for item in batch], dtype=torch.long) tgt_batch = torch.tensor([item["tgt"] for item in batch], dtype=torch.long) # 直接保存每个样本的boundaries列表,后续在模型里处理 boundaries_batch = [item["boundaries"] for item in batch] return { "src": src_batch, "tgt": tgt_batch, "boundaries": boundaries_batch }
3. 在模型/自定义SelfAttention层中提取boundaries并生成掩码
不需要在配置文件里做任何额外的data或KeyMap配置,直接在模型的forward方法里从输入字典中拿到boundaries,然后生成对应的注意力掩码:
class CustomSelfAttention(nn.Module): def __init__(self, embed_dim, num_heads): super().__init__() self.self_attn = nn.MultiheadAttention(embed_dim, num_heads) def generate_boundary_mask(self, boundaries_batch, seq_len): # 生成形状为[batch_size, seq_len, seq_len]的掩码矩阵 batch_size = len(boundaries_batch) mask = torch.ones(batch_size, seq_len, seq_len, dtype=torch.bool, device=self.self_attn.out_proj.weight.device) for batch_idx in range(batch_size): for start, end in boundaries_batch[batch_idx]: # 同一个词内的位置允许互相注意力,所以把对应位置的mask设为False(PyTorch中True表示被掩码) mask[batch_idx, start:end+1, start:end+1] = False return mask def forward(self, src, boundaries_batch): seq_len = src.size(0) # 假设src的形状是[seq_len, batch_size, embed_dim](MultiheadAttention的默认输入格式) attn_mask = self.generate_boundary_mask(boundaries_batch, seq_len) # 传入掩码到SelfAttention层 attn_output, _ = self.self_attn(src, src, src, attn_mask=attn_mask) return attn_output
然后在你的主模型中调用这个自定义注意力层:
class TranslationModel(nn.Module): def __init__(self, src_vocab_size, tgt_vocab_size, embed_dim, num_heads): super().__init__() self.src_embedding = nn.Embedding(src_vocab_size, embed_dim) self.custom_attn = CustomSelfAttention(embed_dim, num_heads) # 其他层... def forward(self, inputs): src = inputs["src"].transpose(0, 1) # 转换为[seq_len, batch_size] src_emb = self.src_embedding(src) boundaries_batch = inputs["boundaries"] attn_output = self.custom_attn(src_emb, boundaries_batch) # 后续的解码、输出处理... return outputs
关键思路总结
- 不要试图让框架把
boundaries当成独立的输入序列来加载,而是让它作为每个样本的辅助元数据存在于样本字典中。 - 配置文件里不需要为
boundaries做任何额外配置,因为数据集已经把它包含在返回的样本里了。 - 在模型内部直接从输入字典提取
boundaries,然后根据需求生成注意力掩码,限制SelfAttention的范围。
内容的提问来源于stack exchange,提问作者ben1806
相关产品推荐
相关产品推荐

