news 2026/9/17 19:01:36

使用 Flax NNX 构建英西机器翻译的 Encoder-Decoder Transformer 实战教程

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
使用 Flax NNX 构建英西机器翻译的 Encoder-Decoder Transformer 实战教程

使用 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变为nnxops变为jnp,并在 JAX 中一步步训练一个英西翻译模型。原始 Notebook 位于 docs_nnx/examples/machine_translation.ipynb,本教程对应的 Markdown 版本为 docs_nnx/examples/machine_translation.md。

1. 环境准备与依赖安装

首先安装所需依赖:tiktoken(OpenAI 分词器)、grain(Google 数据加载库)、flaxoptax(优化器),另外本教程还会用到requestsmatplotlibtqdm

# !pip install tiktoken grain flax optax

标准库与数据处理导入:

import pathlib import random import string import re import numpy as np

JAX、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.Modulennx.MultiHeadAttentionnnx.Embednnx.Linearnnx.LayerNormnnx.Dropoutnnx.Rngsnnx.jitnnx.value_and_gradnnx.viewnnx.MultiMetricnnx.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 = 20

4. 文本标准化、分词与填充

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_inputstarget_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换成jnpkeras/layers换成nnx;一些模块专属参数随之增减,例如新版本中绝大多数模块都携带rngsMultiHeadAttention的调用中带有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层在源语句与目标语句之间是共享复用的。前向过程:

  1. 对英文(encoder_inputs)token 做嵌入与位置编码,送入编码器;
  2. 对西语(decoder_inputs)token 做嵌入与位置编码,连同编码器输出一起送入解码器,并施加 dropout;
  3. 用最后一层线性投影把解码器输出映射到词表上的分布。
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_inputsdecoder_inputstarget_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 遍历一轮);
  • DataLoaderworker_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_epochevaluate_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 loss

eval_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 = 10

6.3 用 nnx.view 切换训练 / 推理模式

训练时需要完整的 dropout 随机化,而评估模型时不希望有 dropout,通过deterministic=True标志控制。这里对同一个模型创建两个视图(train_modeleval_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”(“去吃朋友”),因为hacercomer在此上下文中被混淆;
  • 推理采用贪心 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.*/Modelnnx.Module子类__init__中实例化子层,在__call__中定义前向
ops.*jnp.*所有张量运算改用 JAX numpy
keras.layers.MultiHeadAttentionnnx.MultiHeadAttention需要num_headsin_features,可传decodemask
Embeddingnnx.Embednum_embeddings/features命名不同
层内随机性(keras 自动管理)rngs: nnx.RngsNNX 要求每个含随机性的模块显式传入rngs
model.compile+model.fit手写train_step+ 外层 epoch 循环训练步用@nnx.jitnnx.value_and_grad显式定义
model.evaluate行为开关nnx.view(model, deterministic=...)同一份参数,两种前向行为
keras内置 metricsnnx.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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/17 18:59:22

80V/8A异步降压控制器实战:从BUCK原理到48V转12V模块设计

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/17 18:57:27

编译原理词法分析器实战:Token扫描、最长匹配与错误恢复

如果你正在上编译原理这门课&#xff0c;大概率会在开课三五周之后收到第一份实验任务&#xff1a;实现一个简单的词法分析器。很多人拿到题目第一反应是"这不就是个字符串切割吗"&#xff0c;然后花一个晚上写了两百行 if-else&#xff0c;跑通课本上那几行示例就交…

作者头像 李华
网站建设 2026/9/17 18:54:48

AD域架构详解:父域、子域、树域、林域的区别与实战部署

很多企业IT管理员第一次接触Windows域环境时&#xff0c;都会被“父域、子域、树域、林域”这四个概念绕晕。我在给客户做架构评估时经常遇到这种状况&#xff1a;域控已经搭了几个&#xff0c;名字看着像一家子&#xff0c;但你要问他们林子里的信任关系是怎么走的、子域和树域…

作者头像 李华
网站建设 2026/9/17 18:52:32

DeepSeek教育行业落地指南:从API接入到LoRA微调与推理优化

简介&#xff1a;面向教育行业AI应用开发者、方案架构师与高校技术团队&#xff0c;这份969页PDF系统讲解基于DeepSeek大模型的智能助教完整方案&#xff0c;重点解决教育场景下对话式辅导与课程设计自动化的落地痛点。全篇共65个大章节&#xff0c;从DeepSeek API接入与本地部…

作者头像 李华
网站建设 2026/9/17 18:49:02

Python+Gurobi求解VRPTW:从数学建模到代码实现详解

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华