如何适配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
相关产品推荐
相关产品推荐

