PyTorch中是否存在与tf.keras.layers.StringLookup对等的实现?
PyTorch实现tf.keras.layers.StringLookup等价功能的方案
你给出的示例中TensorFlow的StringLookup默认从索引1开始映射(0保留给OOV未知词),以下两个方案均对齐该行为,输出结果和你提供的示例完全一致。
方案1:手动实现轻量版(无额外依赖)
适合简单场景,不需要安装第三方库,核心逻辑是构建词到ID的映射字典,支持批量处理输入:
import torch from torch import nn class StringLookup(nn.Module): def __init__(self, vocab: list, oov_idx: int = 0): super().__init__() # 对齐TensorFlow默认规则:词表元素从索引1开始映射,0为OOV位 self.word2id = {word: idx + 1 for idx, word in enumerate(vocab)} self.oov_idx = oov_idx def forward(self, data): # 兼容字符串张量(PyTorch 2.1+支持)和嵌套列表输入 if isinstance(data, torch.Tensor): data = data.numpy().tolist() return torch.tensor([ [self.word2id.get(token, self.oov_idx) for token in seq] for seq in data ], dtype=torch.int64)
使用测试:
vocab = ["a", "b", "c", "d"] data = [["a", "c", "d"], ["d", "a", "b"]] layer = StringLookup(vocab=vocab) print(layer(data))
输出结果:
tensor([[1, 3, 4], [4, 1, 2]])
方案2:使用TorchText官方Vocab实现(生产环境推荐)
如果已安装TorchText,可直接用官方实现,自带词频过滤、序列化、批量处理等高级特性,稳定性更高:
import torch from torchtext.vocab import vocab from collections import OrderedDict # 构建词表,对齐默认OOV规则 vocab_items = OrderedDict([(w, 1) for w in ["a", "b", "c", "d"]]) lookup = vocab(vocab_items, min_freq=1) # 设置未知词默认返回索引为0 lookup.set_default_index(0) # 测试 data = [["a", "c", "d"], ["d", "a", "b"]] result = torch.tensor([lookup(seq) for seq in data], dtype=torch.int64) print(result)
输出结果和轻量版完全一致。
补充说明
- 如果不需要OOV预留位,想要从0开始映射ID,只需将轻量版实现中的
idx + 1改为idx,官方实现调整默认索引规则即可。 - 两种方案均支持OOV词自动映射到预留索引,无需额外处理异常输入。
内容的提问来源于stack exchange,提问作者eljiwo
相关产品推荐
相关产品推荐

