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

如何适配WMT14等数据集到TensorFlow官方Transformer教程并完成分词

德英翻译数据集适配TensorFlow Transformer教程的解决方案

1. 对齐原有PyTorch逻辑的迁移实现

该方案完全复现你现有PyTorch代码的分词、词表构建逻辑,适配官方教程的输入格式要求:

1.1 环境依赖安装

执行以下命令安装所需工具包:

pip install spacy tensorflow tensorflow-text tensorflow-datasets
python -m spacy download de_core_news_sm
python -m spacy download en_core_web_sm

1.2 分词逻辑封装

将spacy分词逻辑包装为TensorFlow可兼容的运算操作:

import spacy
import tensorflow as tf
import tensorflow_text as tf_text
import tensorflow_datasets as tfds
from collections import Counter

# 加载spacy分词模型
spacy_de = spacy.load('de_core_news_sm')
spacy_en = spacy.load('en_core_web_sm')

# 原生分词逻辑,和PyTorch实现完全对齐
def tokenize_de(text):
    text = text.numpy().decode('utf-8')
    return [tok.text for tok in spacy_de.tokenizer(text)]

def tokenize_en(text):
    text = text.numpy().decode('utf-8')
    tokens = [tok.text for tok in spacy_en.tokenizer(text)]
    # 目标语言自动加BOS、EOS标记
    return ['<s>'] + tokens + ['</s>']

# 包装为TensorFlow可用的OP
def tf_tokenize_de(text):
    res = tf.py_function(tokenize_de, inp=[text], Tout=tf.string)
    res.set_shape([None])
    return res

def tf_tokenize_en(text):
    res = tf.py_function(tokenize_en, inp=[text], Tout=tf.string)
    res.set_shape([None])
    return res

1.3 数据集加载与过滤

加载IWSLT/WMT14数据集,过滤长度超过阈值的样本:

MAX_LEN = 100
# 加载IWSLT德英数据集,如需使用WMT14替换为"wmt14_translate/de-en"即可
dataset = tfds.load('iwslt2017/de-en', split=['train', 'validation', 'test'])
train_ds, val_ds, test_ds = dataset

def preprocess_sample(sample):
    src = tf_tokenize_de(sample['de'])
    tgt = tf_tokenize_en(sample['en'])
    return src, tgt

def filter_length(src, tgt):
    return tf.logical_and(tf.size(src) <= MAX_LEN, tf.size(tgt) <= MAX_LEN)

# 预处理训练集
train_ds = train_ds.map(preprocess_sample, num_parallel_calls=tf.data.AUTOTUNE)
train_ds = train_ds.filter(filter_length)

1.4 词表构建与批量处理

按照最小词频过滤规则构建词表,输出和官方教程完全兼容的训练batch:

MIN_FREQ = 2
BLANK_WORD = "<blank>"
special_tokens = [BLANK_WORD, '<s>', '</s>', '<unk>']

# 统计训练集词频
src_counter = Counter()
tgt_counter = Counter()
for src, tgt in train_ds.as_numpy_iterator():
    src_counter.update([t.decode('utf-8') for t in src])
    tgt_counter.update([t.decode('utf-8') for t in tgt])

# 构建词表
src_vocab = special_tokens + [tok for tok, cnt in src_counter.items() if cnt >= MIN_FREQ and tok not in special_tokens]
tgt_vocab = special_tokens + [tok for tok, cnt in tgt_counter.items() if cnt >= MIN_FREQ and tok not in special_tokens]

# 保存词表文件方便后续复用
with open('src_vocab.txt', 'w', encoding='utf-8') as f:
    f.write('\n'.join(src_vocab))
with open('tgt_vocab.txt', 'w', encoding='utf-8') as f:
    f.write('\n'.join(tgt_vocab))

# 构建词表查询器
src_lookup = tf_text.lookup.VocabLookup(
    vocabulary_file='src_vocab.txt',
    oov_value=special_tokens.index('<unk>')
)
tgt_lookup = tf_text.lookup.VocabLookup(
    vocabulary_file='tgt_vocab.txt',
    oov_value=special_tokens.index('<unk>')
)

# 批量处理逻辑,对齐官方教程输入格式
BATCH_SIZE = 32
def process_batch(src, tgt):
    src_ids = src_lookup(src)
    tgt_ids = tgt_lookup(tgt)
    # 序列padding
    src_padded = tf.keras.preprocessing.sequence.pad_sequences(src_ids, padding='post', value=0)
    tgt_padded = tf.keras.preprocessing.sequence.pad_sequences(tgt_ids, padding='post', value=0)
    # 拆分目标语言输入与标签
    return (src_padded, tgt_padded[:, :-1]), tgt_padded[:, 1:]

# 生成最终训练批次
train_batches = train_ds.padded_batch(BATCH_SIZE).map(process_batch, num_parallel_calls=tf.data.AUTOTUNE).prefetch(tf.data.AUTOTUNE)
val_batches = val_ds.map(preprocess_sample).filter(filter_length).padded_batch(BATCH_SIZE).map(process_batch).prefetch(tf.data.AUTOTUNE)

注:生成的train_batches、val_batches可直接代入官方Transformer教程的训练代码使用,无需修改模型结构


2. 基于BERT分词器的简化实现

如果不需要对齐原有spacy分词逻辑,可以直接使用多语言BERT分词器,无需手动构建词表,实现更简单:

from transformers import AutoTokenizer
tokenizer = AutoTokenizer.from_pretrained('bert-base-multilingual-cased')

def preprocess_bert(sample):
    src = tokenizer(sample['de'].numpy().decode('utf-8'), truncation=True, max_length=100)['input_ids']
    tgt = tokenizer(sample['en'].numpy().decode('utf-8'), truncation=True, max_length=100)['input_ids']
    return src, tgt

# 后续批量处理逻辑和上述1.4节完全一致

内容的提问来源于stack exchange,提问作者Dametime

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.01 23:57:03