ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

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

使用 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.nnxNew 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 NNXkeras/layers变为nnxops变为jnp并在 JAX 中一步步训练一个英西翻译模型。原始 Notebook 位于 docs_nnx/examples/machine_translation.ipynb本教程对应的 Markdown 版本为 docs_nnx/examples/machine_translation.md。1. 环境准备与依赖安装首先安装所需依赖tiktokenOpenAI 分词器、grainGoogle 数据加载库、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_baseGPT-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]格式完好。同时记录分词器词表大小并把所有输入的最大序列长度固定为 20strip_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换成jnpkeras/layers换成nnx一些模块专属参数随之增减例如新版本中绝大多数模块都携带rngsMultiHeadAttention的调用中带有decodeFalse。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_headsnum_heads, in_featuresembed_dim, decodeFalse, rngsrngs) self.dense_proj nnx.Sequential( nnx.Linear(embed_dim, dense_dim, rngsrngs), nnx.relu, nnx.Linear(dense_dim, embed_dim, rngsrngs), ) self.layernorm_1 nnx.LayerNorm(embed_dim, rngsrngs) self.layernorm_2 nnx.LayerNorm(embed_dim, rngsrngs) 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.pynnx.MultiHeadAttention(num_heads, in_features, qkv_features, decode, rngs)是仓库中多注意力头的官方实现示例即用nnx.MultiHeadAttention(num_heads8, in_features5, qkv_features16, decodeFalse, rngsnnx.Rngs(0))初始化。decodeFalse表示训练/并行编码阶段使用完整的注意力矩阵而非逐 token 增量解码。前馈部分使用的nnx.Linear(in_features, out_features, rngsrngs)与nnx.Sequential分别定义于 flax/nnx/nn/linear.pyLinear作用于输入最后一维与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_embeddingsvocab_size, featuresembed_dim, rngsrngs) self.position_embeddings nnx.Embed(num_embeddingssequence_length, featuresembed_dim, rngsrngs) 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, maskNone): if mask is None: return None else: return jnp.not_equal(inputs, 0)源码对照flax/nnx/nn/linear.pynnx.Embed(num_embeddings, features, rngsrngs)是仓库官方嵌入模块示例nnx.Embed(num_embeddings5, features3, rngsnnx.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_headsnum_heads, in_featuresembed_dim, decodeFalse, rngsrngs) self.attention_2 nnx.MultiHeadAttention(num_headsnum_heads, in_featuresembed_dim, decodeFalse, rngsrngs) self.dense_proj nnx.Sequential( nnx.Linear(embed_dim, latent_dim, rngsrngs), nnx.relu, nnx.Linear(latent_dim, embed_dim, rngsrngs), ) self.layernorm_1 nnx.LayerNorm(embed_dim, rngsrngs) self.layernorm_2 nnx.LayerNorm(embed_dim, rngsrngs) self.layernorm_3 nnx.LayerNorm(embed_dim, rngsrngs) def __call__(self, inputs, encoder_outputs): causal_mask nnx.make_causal_mask(inputs[:,:,0]) attention_output_1 self.attention_1( inputs_qinputs, inputs_vinputs, inputs_kinputs, maskcausal_mask ) out_1 self.layernorm_1(inputs attention_output_1) attention_output_2 self.attention_2( inputs_qout_1, inputs_vencoder_outputs, inputs_kencoder_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.pynnx.make_causal_mask(x, extra_batch_dims0, dtypejnp.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_inputstoken 做嵌入与位置编码送入编码器对西语decoder_inputstoken 做嵌入与位置编码连同编码器输出一起送入解码器并施加 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, rngsrngs) self.positional_embedding PositionalEmbedding(sequence_length, vocab_size, embed_dim, rngsrngs) self.decoder TransformerDecoder(embed_dim, latent_dim, num_heads, rngsrngs) self.dropout nnx.Dropout(ratedropout_rate, rngsrngs) self.dense nnx.Linear(embed_dim, vocab_size, rngsrngs) 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, deterministicFalse) logits self.dense(decoder_outputs) return logits源码对照flax/nnx/nn/stochastic.pynnx.Dropout(rate, rngs)是官方 dropout 层其 docstring 明确说明使用 dropout 需调用train()方法或在构造/调用时传入deterministicFalse。本教程在模型__call__中显式传deterministicFalse训练期推理期则通过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), shuffleTrue, seed12, # Seed for reproducibility shard_optionsgrain.NoSharding(), # No sharding since its a single-device setup num_epochs1, # Iterate over the dataset for one epoch ) val_sampler grain.IndexSampler( len(val_data), shuffleFalse, seed12, shard_optionsgrain.NoSharding(), num_epochs1, ) train_loader grain.DataLoader( data_sourcetrain_data, samplertrain_sampler, # Sampler to determine how to access the data worker_count4, # Number of child processes launched to parallelize the transformations worker_buffer_size2, # Count of output batches to produce in advance per worker operations[ CustomPreprocessing(), grain.Batch(batch_sizebatch_size, drop_remainderTrue), ] ) val_loader grain.DataLoader( data_sourceval_data, samplerval_sampler, worker_count4, worker_buffer_size2, operations[ CustomPreprocessing(), grain.Batch(batch_sizebatch_size), ] )参数说明IndexSampler的三个关键参数shuffle是否打乱、seed复现种子、shard_options单设备用grain.NoSharding()、num_epochs一个 epoch 遍历一轮DataLoader的worker_count4表示启动 4 个子进程并行做变换worker_buffer_size2表示每个 worker 预生成 2 个输出 batchgrain.Batch(batch_size, drop_remainderTrue)在训练时丢弃末尾不足一个 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(logitslogits, labelslabels) 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_metricsnnx.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( lossloss, logitslogits, labelslabels, )nnx.MultiMetric负责跨 batch 累积 loss 与 accuracy两个独立的 history 字典记录每个 epoch 的数值供后续绘图eval_metrics nnx.MultiMetric( lossnnx.metrics.Average(loss), accuracynnx.metrics.Accuracy(), ) train_metrics_history { train_loss: [], } eval_metrics_history { test_loss: [], test_accuracy: [], }源码对照flax/nnx/training/metrics.py 与 flax/nnx/training/metrics.pynnx.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通过deterministicTrue标志控制。这里对同一个模型创建两个视图train_model与eval_model各持一种标志设置。两个视图共享同一份底层参数——因此通过train_model更新权重后eval_model立即可见model TransformerModel(sequence_length, vocab_size, embed_dim, latent_dim, num_heads, dropout_rate, rngsrng) train_model nnx.view(model, deterministicFalse) eval_model nnx.view(model, deterministicTrue) optimizer nnx.ModelAndOptimizer(model, optax.adamw(learning_rate))源码对照flax/nnx/module.pynnx.view(node, **kwargs)创建一个属性按 kwargs 更新、但 JAX 数组引用与原始节点共享的新节点若 kwargs 中的属性在任何模块中都找不到会抛出 ValueError。这正是 NNX 实现“同一份参数、不同行为开关”的官方机制。优化器方面nnx.ModelAndOptimizer(model, tx)flax/nnx/training/optimizer.py是模型与优化器的便捷组合类内部代理到nnx.Optimizerflax/nnx/training/optimizer.py单 Optax 优化器的通用训练状态。注意ModelAndOptimizer已被官方标记为 deprecated新代码建议直接使用nnx.Optimizer(model, tx)其update(grads)方法会把梯度应用到模型参数并推进 Optax 优化器内部状态。6.4 单 epoch 训练与验证函数train_one_epoch遍历训练集所有 batch对每个 batch 调用train_step并记录逐步 lossevaluate_model先重置累积指标再对完整验证集跑eval_step最后打印并记录 epoch 级 loss 与 accuracybar_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( descf[train] epoch: {epoch}/{num_epochs}, , totaltrain_total_steps, bar_formatbar_format, leaveTrue, ) 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[ftest_{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_remainderTrue保持一致保证进度条总数与实际 batch 数吻合。7. 开始训练数据加载器、模型、优化器与 epoch 训练/验证函数都就绪现在正式开跑。在 RTX 3090 上本配置约占用 19GB 显存batch_size512时每个 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], labelLoss value during the training) plt.yscale(log) plt.legend();验证集上的 loss 与 accuracyfig, 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注意是deterministicTrue的视图取位置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]: Were 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 theres anything I can do. [Translation]: [start] déjame saber si hay cualquier cosa que yo pueda hacer [end] [Input]: Lets 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]: Im 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]观察与局限多数译文语法正确、语义忠实如 “Were 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.*/Modelnnx.Module子类在__init__中实例化子层在__call__中定义前向ops.*jnp.*所有张量运算改用 JAX numpykeras.layers.MultiHeadAttentionnnx.MultiHeadAttention需要num_heads、in_features可传decode与maskEmbeddingnnx.Embednum_embeddings/features命名不同层内随机性keras 自动管理rngs: nnx.RngsNNX 要求每个含随机性的模块显式传入rngsmodel.compilemodel.fit手写train_step 外层 epoch 循环训练步用nnx.jit与nnx.value_and_grad显式定义model.evaluate行为开关nnx.view(model, deterministic...)同一份参数两种前向行为keras内置 metricsnnx.MultiMetricnnx.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),仅供参考
返回列表