You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.29 17:45:01