使用 Flax NNX 构建英西机器翻译的 Encoder-Decoder Transformer 实战教程
【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax
本文以 Flax 的现代 API ——flax.nnx(New NNX)为主线,完整复刻并深度讲解一个可直接运行的英西(English→Spanish)机器翻译任务:从数据下载、tiktoken 分词、Transformer 编码器-解码器模型的逐层搭建,到基于 grain 的数据加载、训练循环、指标跟踪、推理与效果分析。读完本文,你将掌握nnx.Module的面向对象建模方式、nnx.jit/nnx.value_and_grad的显式训练步、nnx.view的确定性开关、nnx.MultiMetric指标管理,以及用nnx.ModelAndOptimizer(或nnx.Optimizer)驱动 Optax 优化器完成端到端训练的完整工程链路。
本教程改编自 Keras 官方文档的英文到西班牙语序列到序列 Transformer 翻译示例(该示例又源自 François Chollet《Deep Learning with Python》第二版),但在实现上完全切换到 JAX + Flax NNX:keras/layers变为nnx,ops变为jnp,并在 JAX 中一步步训练一个英西翻译模型。原始 Notebook 位于 docs_nnx/examples/machine_translation.ipynb,本教程对应的 Markdown 版本为 docs_nnx/examples/machine_translation.md。
1. 环境准备与依赖安装
首先安装所需依赖:tiktoken(OpenAI 分词器)、grain(Google 数据加载库)、flax与optax(优化器),另外本教程还会用到requests、matplotlib与tqdm:
# !pip install tiktoken grain flax optax标准库与数据处理导入:
import pathlib import random import string import re import numpy as npJAX、Flax 与训练框架导入:
import jax.numpy as jnp import optax from flax import nnx分词器(tiktoken)、数据加载器(grain)与进度条(tqdm)导入:
import tiktoken import grain.python as grain import tqdm说明:本文后续所用到的
nnx.Module、nnx.MultiHeadAttention、nnx.Embed、nnx.Linear、nnx.LayerNorm、nnx.Dropout、nnx.Rngs、nnx.jit、nnx.value_and_grad、nnx.view、nnx.MultiMetric与nnx.Optimizer等 API 均可在仓库 flax/nnx/ 目录下找到对应源码实现。
2. 下载与提取 spa-eng 平行语料
获取数据的方式有很多,本教程为简单直观起见,将数据下载到临时目录、解压后读入 Python 对象并进行处理。
import requests import zipfile import tempfile url = "http://storage.googleapis.com/download.tensorflow.org/data/spa-eng.zip" with tempfile.TemporaryDirectory() as temp_dir: temp_path = pathlib.Path(temp_dir) zip_file_path = temp_path / "spa-eng.zip" response = requests.get(url) zip_file_path.write_bytes(response.content) with zipfile.ZipFile(zip_file_path, "r") as zip_ref: zip_ref.extractall(temp_path) text_file = temp_path / "spa-eng" / "spa.txt" with open(text_file) as f: lines = f.read().split("\n")[:-1] text_pairs = [] for line in lines: eng, spa = line.split("\t") spa = "[start] " + spa + " [end]" text_pairs.append((eng, spa))要点说明:
spa-eng.zip是 TensorFlow 公开的英西平行语料,解压后得到spa-eng/spa.txt,每行形如英文\t西班牙文。- 西班牙语句子在处理时被加上
[start] ... [end]标记,作为解码器端的起始符与结束符。 - 使用
TemporaryDirectory保证数据使用完毕后自动清理,不污染工作区。
3. 划分训练 / 验证 / 测试集
与原教程保持一致,方便对照“哪些相同、哪些不同”;一个早期差异是:本教程选用现成的 tiktoken 编码器cl100k_base(GPT-4 同族的分词方案),它对多种语言理解广泛且速度快。
random.shuffle(text_pairs) num_val_samples = int(0.15 * len(text_pairs)) num_train_samples = len(text_pairs) - 2 * num_val_samples train_pairs = text_pairs[:num_train_samples] val_pairs = text_pairs[num_train_samples : num_train_samples + num_val_samples] test_pairs = text_pairs[num_train_samples + num_val_samples :] print(f"{len(text_pairs)} total pairs") print(f"{len(train_pairs)} training pairs") print(f"{len(val_pairs)} validation pairs") print(f"{len(test_pairs)} test pairs")数据集按 70% / 15% / 15% 划分:先留出 15% 作为验证集,再留出 15% 作为测试集,其余为训练集。随机打乱由random.shuffle完成(注意这是 Python 标准库的洗牌,与后续 grain 采样器的seed相互独立)。
随后创建分词器并确定两个全局常量:
tokenizer = tiktoken.get_encoding("cl100k_base")为保持简单并贴合原教程,我们去除标点;但保留[与],以保证[start]、[end]格式完好。同时记录分词器词表大小,并把所有输入的最大序列长度固定为 20:
strip_chars = string.punctuation + "¿" strip_chars = strip_chars.replace("[", "") strip_chars = strip_chars.replace("]", "") vocab_size = tokenizer.n_vocab sequence_length = 204. 文本标准化、分词与填充
custom_standardization将字符串转为小写,并删除上面定义的标点字符(“¿”是西语特有的倒问号,一并剔除):
def custom_standardization(input_string): lowercase = input_string.lower() return re.sub(f"[{re.escape(strip_chars)}]", "", lowercase)tokenize_and_pad将字符串编码为 token ID 序列,超长则截断到max_length,不足则用 0 填充,使每个样本都是定长:
def tokenize_and_pad(text, tokenizer, max_length): tokens = tokenizer.encode(text)[:max_length] padded = tokens + [0] * (max_length - len(tokens)) if len(tokens) < max_length else tokens ##assumes list-like - (https://github.com/openai/tiktoken/blob/main/tiktoken/core.py#L81 current tiktoken out) return padded注:
tiktoken.encode返回的是list[int],因此可以直接用tokens + [0] * n拼接填充。
format_dataset对英西两个字符串做标准化与分词,然后返回三个数组:
encoder_inputs—— 完整分词后的英文句子,喂给编码器;decoder_inputs—— 西语句子整体右移一位(去掉最后一个 token),作为解码器每一步的输入提示;target_output—— 西语句子整体左移一位(去掉第一个 token),作为每一步要预测的目标。
def format_dataset(eng, spa, tokenizer, sequence_length): eng = custom_standardization(eng) spa = custom_standardization(spa) eng = tokenize_and_pad(eng, tokenizer, sequence_length) spa = tokenize_and_pad(spa, tokenizer, sequence_length) return { "encoder_inputs": eng, "decoder_inputs": spa[:-1], "target_output": spa[1:], }对每个划分应用预处理,得到最终的内存数据集:
train_data = [format_dataset(eng, spa, tokenizer, sequence_length) for eng, spa in train_pairs] val_data = [format_dataset(eng, spa, tokenizer, sequence_length) for eng, spa in val_pairs] test_data = [format_dataset(eng, spa, tokenizer, sequence_length) for eng, spa in test_pairs]此时数据已完成提取、格式化、分词与填充,train/validate/test 各自包含字典条目,形如:
## data selection example print(train_data[135])输出大致如下(encoder_inputs中尾部大量 0 是填充位;decoder_inputs比target_output整体左移一位,正是“右移的输入、左移的目标”这一翻译范式):
{'encoder_inputs': [9514, 265, 3339, 264, 2466, 16930, 1618, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0], 'decoder_inputs': [29563, 60, 1826, 7206, 71086, 37116, 653, 16109, 1493, 54189, 510, 408, 60, 0, 0, 0, 0, 0, 0], 'target_output': [60, 1826, 7206, 71086, 37116, 653, 16109, 1493, 54189, 510, 408, 60, 0, 0, 0, 0, 0, 0, 0]}5. 定义 Transformer 组件:Encoder、Decoder、Positional Embedding
从结构上看,本实现与原教程高度相似,区别在于:ops换成jnp,keras/layers换成nnx;一些模块专属参数随之增减,例如新版本中绝大多数模块都携带rngs,MultiHeadAttention的调用中带有decode=False。
5.1 TransformerEncoder:自注意力 + 前馈投影
TransformerEncoder实现标准的编码器块:对输入序列做自注意力,随后接两层前馈投影;每个子层之后都有残差连接与层归一化。
class TransformerEncoder(nnx.Module): def __init__(self, embed_dim: int, dense_dim: int, num_heads: int, rngs: nnx.Rngs, **kwargs): self.embed_dim = embed_dim self.dense_dim = dense_dim self.num_heads = num_heads self.attention = nnx.MultiHeadAttention(num_heads=num_heads, in_features=embed_dim, decode=False, rngs=rngs) self.dense_proj = nnx.Sequential( nnx.Linear(embed_dim, dense_dim, rngs=rngs), nnx.relu, nnx.Linear(dense_dim, embed_dim, rngs=rngs), ) self.layernorm_1 = nnx.LayerNorm(embed_dim, rngs=rngs) self.layernorm_2 = nnx.LayerNorm(embed_dim, rngs=rngs) def __call__(self, inputs): attention_output = self.attention( inputs_q = inputs, inputs_k = inputs, inputs_v = inputs, decode = False ) proj_input = self.layernorm_1(inputs + attention_output) proj_output = self.dense_proj(proj_input) return self.layernorm_2(proj_input + proj_output)源码对照(flax/nnx/nn/attention.py):nnx.MultiHeadAttention(num_heads, in_features, qkv_features, decode, rngs)是仓库中多注意力头的官方实现,示例即用nnx.MultiHeadAttention(num_heads=8, in_features=5, qkv_features=16, decode=False, rngs=nnx.Rngs(0))初始化。decode=False表示训练/并行编码阶段使用完整的注意力矩阵而非逐 token 增量解码。前馈部分使用的nnx.Linear(in_features, out_features, rngs=rngs)与nnx.Sequential分别定义于 flax/nnx/nn/linear.py(Linear作用于输入最后一维)与nnx.Sequential容器。
5.2 PositionalEmbedding:可学习的词嵌入 + 位置嵌入
PositionalEmbedding组合两张可学习嵌入表:一张把 token ID 映射为向量,另一张把位置索引(0, 1, 2, …)映射为向量;两者相加,使每个 token 同时具备语义表示与位置表示。compute_mask返回布尔数组,标记非填充 token(即 token ID 不为 0 的位置)。
class PositionalEmbedding(nnx.Module): def __init__(self, sequence_length: int, vocab_size: int, embed_dim: int, rngs: nnx.Rngs, **kwargs): self.token_embeddings = nnx.Embed(num_embeddings=vocab_size, features=embed_dim, rngs=rngs) self.position_embeddings = nnx.Embed(num_embeddings=sequence_length, features=embed_dim, rngs=rngs) self.sequence_length = sequence_length self.vocab_size = vocab_size self.embed_dim = embed_dim def __call__(self, inputs): length = inputs.shape[1] positions = jnp.arange(0, length)[None, :] embedded_tokens = self.token_embeddings(inputs) embedded_positions = self.position_embeddings(positions) return embedded_tokens + embedded_positions def compute_mask(self, inputs, mask=None): if mask is None: return None else: return jnp.not_equal(inputs, 0)源码对照(flax/nnx/nn/linear.py):nnx.Embed(num_embeddings, features, rngs=rngs)是仓库官方嵌入模块,示例nnx.Embed(num_embeddings=5, features=3, rngs=nnx.Rngs(0))说明其构造签名。这里词表规模即tokenizer.n_vocab(约 10 万量级),位置表规模为sequence_length。
5.3 TransformerDecoder:因果自注意力 + 交叉注意力
TransformerDecoder实现带两层注意力的解码器块:
attention_1是目标序列上的带掩码自注意力。因果掩码(causal mask)防止每个位置attend到未来 token,这是自回归生成“只用已生成 token 预测下一个 token”的关键;attention_2是交叉注意力:query 来自解码器,key/value 来自编码器输出,让解码器在每一步都能关注整句源文。
每个注意力层之后都跟随残差连接、层归一化与共享的前馈投影。
class TransformerDecoder(nnx.Module): def __init__(self, embed_dim: int, latent_dim: int, num_heads: int, rngs: nnx.Rngs, **kwargs): self.embed_dim = embed_dim self.latent_dim = latent_dim self.num_heads = num_heads self.attention_1 = nnx.MultiHeadAttention(num_heads=num_heads, in_features=embed_dim, decode=False, rngs=rngs) self.attention_2 = nnx.MultiHeadAttention(num_heads=num_heads, in_features=embed_dim, decode=False, rngs=rngs) self.dense_proj = nnx.Sequential( nnx.Linear(embed_dim, latent_dim, rngs=rngs), nnx.relu, nnx.Linear(latent_dim, embed_dim, rngs=rngs), ) self.layernorm_1 = nnx.LayerNorm(embed_dim, rngs=rngs) self.layernorm_2 = nnx.LayerNorm(embed_dim, rngs=rngs) self.layernorm_3 = nnx.LayerNorm(embed_dim, rngs=rngs) def __call__(self, inputs, encoder_outputs): causal_mask = nnx.make_causal_mask(inputs[:,:,0]) attention_output_1 = self.attention_1( inputs_q=inputs, inputs_v=inputs, inputs_k=inputs, mask=causal_mask ) out_1 = self.layernorm_1(inputs + attention_output_1) attention_output_2 = self.attention_2( inputs_q=out_1, inputs_v=encoder_outputs, inputs_k=encoder_outputs ) out_2 = self.layernorm_2(out_1 + attention_output_2) proj_output = self.dense_proj(out_2) return self.layernorm_3(out_2 + proj_output)源码对照(flax/nnx/nn/attention.py):nnx.make_causal_mask(x, extra_batch_dims=0, dtype=jnp.float32)对形状[batch..., len]的输入生成形状[batch..., 1, len, len]的因果掩码,内部用jnp.arange广播成坐标矩阵后以jnp.greater_equal比较,保证位置 i 只能看到位置 ≤ i 的 token。本教程在__call__中以inputs[:,:,0]取解码器输入的最后维(embedding 维)作为掩码构造依据——这是该 Notebook 特有的写法,等价地也可直接用inputs.shape[-1]或单独传入长度信息构造掩码。
5.4 TransformerModel:组装完整编码器-解码器
TransformerModel把所有组件串成完整的编码器-解码器架构。值得注意:positional_embedding层在源语句与目标语句之间是共享复用的。前向过程:
- 对英文(
encoder_inputs)token 做嵌入与位置编码,送入编码器; - 对西语(
decoder_inputs)token 做嵌入与位置编码,连同编码器输出一起送入解码器,并施加 dropout; - 用最后一层线性投影把解码器输出映射到词表上的分布。
class TransformerModel(nnx.Module): def __init__(self, sequence_length: int, vocab_size: int, embed_dim: int, latent_dim: int, num_heads: int, dropout_rate: float, rngs: nnx.Rngs): self.sequence_length = sequence_length self.vocab_size = vocab_size self.embed_dim = embed_dim self.latent_dim = latent_dim self.num_heads = num_heads self.dropout_rate = dropout_rate self.encoder = TransformerEncoder(embed_dim, latent_dim, num_heads, rngs=rngs) self.positional_embedding = PositionalEmbedding(sequence_length, vocab_size, embed_dim, rngs=rngs) self.decoder = TransformerDecoder(embed_dim, latent_dim, num_heads, rngs=rngs) self.dropout = nnx.Dropout(rate=dropout_rate, rngs=rngs) self.dense = nnx.Linear(embed_dim, vocab_size, rngs=rngs) def __call__(self, encoder_inputs: jnp.array, decoder_inputs: jnp.array): x = self.positional_embedding(encoder_inputs) encoder_outputs = self.encoder(x) x = self.positional_embedding(decoder_inputs) decoder_outputs = self.decoder(x, encoder_outputs) decoder_outputs = self.dropout(decoder_outputs, deterministic=False) logits = self.dense(decoder_outputs) return logits源码对照(flax/nnx/nn/stochastic.py):nnx.Dropout(rate, rngs)是官方 dropout 层,其 docstring 明确说明:使用 dropout 需调用train()方法,或在构造/调用时传入deterministic=False。本教程在模型__call__中显式传deterministic=False(训练期),推理期则通过nnx.view切换(见第 7 节)。
6. 构建数据加载器与训练定义
数据加载阶段用 pygrain 可以更省计算资源,但这里用最直白的方式把每一步展示清楚:数据对进、一组组jnp数组出,与前面构造的字典一一对应(encoder_inputs、decoder_inputs、target_output)。
batch_size = 512 #set here for the loader and model train later on class CustomPreprocessing(grain.MapTransform): def __init__(self): pass def map(self, data): return { "encoder_inputs": np.array(data["encoder_inputs"]), "decoder_inputs": np.array(data["decoder_inputs"]), "target_output": np.array(data["target_output"]), } train_sampler = grain.IndexSampler( len(train_data), shuffle=True, seed=12, # Seed for reproducibility shard_options=grain.NoSharding(), # No sharding since it's a single-device setup num_epochs=1, # Iterate over the dataset for one epoch ) val_sampler = grain.IndexSampler( len(val_data), shuffle=False, seed=12, shard_options=grain.NoSharding(), num_epochs=1, ) train_loader = grain.DataLoader( data_source=train_data, sampler=train_sampler, # Sampler to determine how to access the data worker_count=4, # Number of child processes launched to parallelize the transformations worker_buffer_size=2, # Count of output batches to produce in advance per worker operations=[ CustomPreprocessing(), grain.Batch(batch_size=batch_size, drop_remainder=True), ] ) val_loader = grain.DataLoader( data_source=val_data, sampler=val_sampler, worker_count=4, worker_buffer_size=2, operations=[ CustomPreprocessing(), grain.Batch(batch_size=batch_size), ] )参数说明:
IndexSampler的三个关键参数:shuffle(是否打乱)、seed(复现种子)、shard_options(单设备用grain.NoSharding())、num_epochs(一个 epoch 遍历一轮);DataLoader的worker_count=4表示启动 4 个子进程并行做变换,worker_buffer_size=2表示每个 worker 预生成 2 个输出 batch;grain.Batch(batch_size, drop_remainder=True)在训练时丢弃末尾不足一个 batch 的样本,验证时保留全部(drop_remainder默认 False)。
6.1 损失函数与训练 / 评估步
Optax 没有原教程使用的完全相同的损失函数,但这里的 softmax 交叉熵完全够用——如果你不用带_with_integer_labels后缀的版本,可以自行 one-hot 编码标签。
def compute_loss(logits, labels): loss = optax.softmax_cross_entropy_with_integer_labels(logits=logits, labels=labels) return jnp.mean(loss)原教程中模型与训练的多数细节藏在 keras 内部,这里我们全部显式写出为 step 函数,稍后用于train_one_epoch与evaluate_model。
train_step执行一次前向、计算损失、用nnx.value_and_grad求梯度,并通过优化器更新参数;整体用@nnx.jit编译以提升性能:
@nnx.jit def train_step(model, optimizer, batch): def loss_fn(model, train_encoder_input, train_decoder_input, train_target_input): logits = model(train_encoder_input, train_decoder_input) loss = compute_loss(logits, train_target_input) return loss grad_fn = nnx.value_and_grad(loss_fn) loss, grads = grad_fn(model, jnp.array(batch["encoder_inputs"]), jnp.array(batch["decoder_inputs"]), jnp.array(batch["target_output"])) optimizer.update(grads) return losseval_step只做前向、不更新权重,把 loss 与 accuracy 累积进eval_metrics:
@nnx.jit def eval_step(model, batch, eval_metrics): logits = model(jnp.array(batch["encoder_inputs"]), jnp.array(batch["decoder_inputs"])) loss = compute_loss(logits, jnp.array(batch["target_output"])) labels = jnp.array(batch["target_output"]) eval_metrics.update( loss=loss, logits=logits, labels=labels, )nnx.MultiMetric负责跨 batch 累积 loss 与 accuracy;两个独立的 history 字典记录每个 epoch 的数值,供后续绘图:
eval_metrics = nnx.MultiMetric( loss=nnx.metrics.Average('loss'), accuracy=nnx.metrics.Accuracy(), ) train_metrics_history = { "train_loss": [], } eval_metrics_history = { "test_loss": [], "test_accuracy": [], }源码对照(flax/nnx/training/metrics.py 与 flax/nnx/training/metrics.py):nnx.metrics.Average是最基础的均值指标,nnx.metrics.Accuracy继承自Average且无需传入字符串(它内部知道如何从 logits 与 labels 计算正确率),MultiMetric把多个指标打包,一次update同时更新全部指标。
6.2 关键超参数
模型与训练的关键超参数:
embed_dim—— token 与位置嵌入的维度;latent_dim—— 每个编码/解码块内部前馈投影的宽度;num_heads—— 注意力头数;dropout_rate—— 训练期间丢弃激活的比例;learning_rate/num_epochs—— AdamW 步长与训练轮数。
## Hyperparameters rng = nnx.Rngs(0) embed_dim = 256 latent_dim = 2048 num_heads = 8 dropout_rate = 0.5 vocab_size = tokenizer.n_vocab sequence_length = 20 learning_rate = 1.5e-3 num_epochs = 106.3 用 nnx.view 切换训练 / 推理模式
训练时需要完整的 dropout 随机化,而评估模型时不希望有 dropout,通过deterministic=True标志控制。这里对同一个模型创建两个视图(train_model与eval_model),各持一种标志设置。两个视图共享同一份底层参数——因此通过train_model更新权重后,eval_model立即可见:
model = TransformerModel(sequence_length, vocab_size, embed_dim, latent_dim, num_heads, dropout_rate, rngs=rng) train_model = nnx.view(model, deterministic=False) eval_model = nnx.view(model, deterministic=True) optimizer = nnx.ModelAndOptimizer(model, optax.adamw(learning_rate))源码对照(flax/nnx/module.py):nnx.view(node, **kwargs)创建一个属性按 kwargs 更新、但 JAX 数组引用与原始节点共享的新节点;若 kwargs 中的属性在任何模块中都找不到会抛出 ValueError。这正是 NNX 实现“同一份参数、不同行为开关”的官方机制。
优化器方面,nnx.ModelAndOptimizer(model, tx)(flax/nnx/training/optimizer.py)是模型与优化器的便捷组合类,内部代理到nnx.Optimizer(flax/nnx/training/optimizer.py,单 Optax 优化器的通用训练状态)。注意:ModelAndOptimizer已被官方标记为 deprecated,新代码建议直接使用nnx.Optimizer(model, tx);其update(grads)方法会把梯度应用到模型参数并推进 Optax 优化器内部状态。
6.4 单 epoch 训练与验证函数
train_one_epoch遍历训练集所有 batch,对每个 batch 调用train_step并记录逐步 loss;evaluate_model先重置累积指标,再对完整验证集跑eval_step,最后打印并记录 epoch 级 loss 与 accuracy:
bar_format = "{desc}[{n_fmt}/{total_fmt}]{postfix} [{elapsed}<{remaining}]" train_total_steps = len(train_data) // batch_size def train_one_epoch(epoch): with tqdm.tqdm( desc=f"[train] epoch: {epoch}/{num_epochs}, ", total=train_total_steps, bar_format=bar_format, leave=True, ) as pbar: for batch in train_loader: loss = train_step(train_model, optimizer, batch) train_metrics_history["train_loss"].append(loss.item()) pbar.set_postfix({"loss": loss.item()}) pbar.update(1) def evaluate_model(epoch): # Compute the metrics on the train and val sets after each training epoch. eval_metrics.reset() # Reset the eval metrics for val_batch in val_loader: eval_step(eval_model, val_batch, eval_metrics) for metric, value in eval_metrics.compute().items(): eval_metrics_history[f'test_{metric}'].append(value) print(f"[test] epoch: {epoch + 1}/{num_epochs}") print(f"- total loss: {eval_metrics_history['test_loss'][-1]:0.4f}") print(f"- Accuracy: {eval_metrics_history['test_accuracy'][-1]:0.4f}")注意eval_metrics.reset()必须在每个 epoch 验证前调用,否则指标会跨 epoch 累积导致数值失真;train_total_steps = len(train_data) // batch_size与 loader 的drop_remainder=True保持一致,保证进度条总数与实际 batch 数吻合。
7. 开始训练
数据加载器、模型、优化器与 epoch 训练/验证函数都就绪,现在正式开跑。在 RTX 3090 上,本配置约占用 19GB 显存,batch_size=512时每个 epoch 约 18 秒:
for epoch in range(num_epochs): train_one_epoch(epoch) evaluate_model(epoch)训练 loss 曲线可以用对数坐标绘制——1000 步之后线性坐标很难看出进展:
import matplotlib.pyplot as plt plt.plot(train_metrics_history["train_loss"], label="Loss value during the training") plt.yscale('log') plt.legend();验证集上的 loss 与 accuracy:
fig, axs = plt.subplots(1, 2) axs[0].set_title("Log loss value on eval set") axs[0].plot(np.log(eval_metrics_history["test_loss"])) axs[1].set_title("Accuracy on eval set") axs[1].plot(eval_metrics_history["test_accuracy"]) plt.tight_layout();从训练统计看,accuracy 确实在持续上升,但大约从第 5 个 epoch 之后上升变得艰难——可以合理判断模型在第 5 个 epoch 之后开始过拟合。这也是小规模平行语料 + 大词表(tiktokencl100k_base)训练场景下的常见现象。
8. 用训练好的模型做推理
训练的全部意义在于得到一个可保存、可加载的推理模型。用过近期 LLM 的读者会对这个模式很熟悉:输入句子被分词成数组,然后逐个 token 计算“下一个 token”。与当前主流的 decoder-only LLM 不同,这是一个把“英西互译”模式直接烘焙进架构的 encoder-decoder 模型。
相比原教程的use函数,这里有几处改动——由于使用了 tiktoken 分词器,[start]与[end]不再各是一个 token:[start]被拆成[29563, 60](即"[start"+"]"),[end]被拆成[58308, 60](即"[end"+"]")。因此推理初始只以单个 token[start开头,也无法只用last_token = "[end]"判断结束。另一个主要改动是:输入被假定为单句,而非批量推理。
def decode_sequence(input_sentence): input_sentence = custom_standardization(input_sentence) tokenized_input_sentence = tokenize_and_pad(input_sentence, tokenizer, sequence_length) decoded_sentence = "[start" for i in range(sequence_length): tokenized_target_sentence = tokenize_and_pad(decoded_sentence, tokenizer, sequence_length)[:-1] predictions = eval_model(jnp.array([tokenized_input_sentence]), jnp.array([tokenized_target_sentence])) sampled_token_index = np.argmax(predictions[0,i, :]).item(0) sampled_token = tokenizer.decode([sampled_token_index]) decoded_sentence += "" + sampled_token if decoded_sentence[-5:] == "[end]": break return decoded_sentence逐行拆解:
- 先对输入句做
custom_standardization与定长 token 化; decoded_sentence以"[start"起步;- 循环内把已生成的字符串重新 token 化并去掉最后一位(与训练时
decoder_inputs = spa[:-1]对齐),调用eval_model(注意是deterministic=True的视图); - 取位置
i上 logits 的 argmax 作为当前步采样 token,解码后拼接到decoded_sentence; - 当尾部出现
"[end]"(5 个字符)时提前终止。
随后从测试集随机抽取句子做 10 次翻译:
test_eng_texts = [pair[0] for pair in test_pairs]test_result_pairs = [] for _ in range(10): input_sentence = random.choice(test_eng_texts) translated = decode_sequence(input_sentence) test_result_pairs.append(f"[Input]: {input_sentence} [Translation]: {translated}")9. 测试结果与分析
就模型与数据而言,效果已经相当不错——翻译结果“确实很西语”。不过要提醒:在“交朋友”这件事上,别把hacer(去做)和comer(去吃)搞混了。
for i in test_result_pairs: print(i)示例输出:
[Input]: We're going to have a baby. [Translation]: [start] nosotros vamos a tener un bebé [end] [Input]: You drive too fast. [Translation]: [start] conducís demasiado rápido [end] [Input]: Let me know if there's anything I can do. [Translation]: [start] déjame saber si hay cualquier cosa que yo pueda hacer [end] [Input]: Let's go to the kitchen. [Translation]: [start] vayamos a la cocina [end] [Input]: Tom gasped. [Translation]: [start] tom se quedó sin aliento [end] [Input]: I was just hanging out with some of my friends. [Translation]: [start] estaba escquieto con algunos de mi amigos [end] [Input]: Tom is in the bathroom. [Translation]: [start] tom está en el cuarto de baño [end] [Input]: I feel safe here. [Translation]: [start] me siento segura [end] [Input]: I'm going to need you later. [Translation]: [start] me voy a necesitar después [end] [Input]: A party is a good place to make friends with other people. [Translation]: [start] una fiesta es un buen lugar de comer amigos con otras personas [end]观察与局限:
- 多数译文语法正确、语义忠实,如 “We're going to have a baby.” → “nosotros vamos a tener un bebé”;
- 个别译文暴露了小语料 + 自回归贪心解码的典型缺陷:最后一句把 “make friends” 翻成 “comer amigos”(“去吃朋友”),因为
hacer与comer在此上下文中被混淆; - 推理采用贪心 argmax,没有 beam search 或 temperature sampling;想要更高质量可引入 beam search,或参考仓库 examples/wmt/(WMT 翻译示例)、examples/gemma/(自回归采样实现)等更完整的翻译/采样管线。
10. 从 Keras 到 Flax NNX 的迁移要点回顾
结合本文实现与原 Keras 教程,可总结出 Keras → Flax NNX 的核心迁移心法:
| 原 Keras 概念 | Flax NNX 对应物 | 说明 |
|---|---|---|
keras.layers.*/Model | nnx.Module子类 | 在__init__中实例化子层,在__call__中定义前向 |
ops.* | jnp.* | 所有张量运算改用 JAX numpy |
keras.layers.MultiHeadAttention | nnx.MultiHeadAttention | 需要num_heads、in_features,可传decode与mask |
Embedding | nnx.Embed | num_embeddings/features命名不同 |
| 层内随机性(keras 自动管理) | rngs: nnx.Rngs | NNX 要求每个含随机性的模块显式传入rngs |
model.compile+model.fit | 手写train_step+ 外层 epoch 循环 | 训练步用@nnx.jit与nnx.value_and_grad显式定义 |
model.evaluate行为开关 | nnx.view(model, deterministic=...) | 同一份参数,两种前向行为 |
keras内置 metrics | nnx.MultiMetric+nnx.metrics.Average/Accuracy | 显式reset/update/compute |
本文对应的 Notebook 与 Markdown 位于 docs_nnx/examples/machine_translation.ipynb 与 docs_nnx/examples/machine_translation.md,可与仓库中其他 NNX 示例(如 docs_nnx/examples/minigpt.md、docs_nnx/examples/vit_training.md)对照学习;NNX 各核心模块源码在 flax/nnx/ 下可直接查阅。
【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考