恒美微站
首页
关于我们
建站服务
主题模板
案例展示
资讯中心
联系我们
用 JAX 与 Flax NNX 从零预训练 miniGPT:数据并行与张量并行的完整实战
首页
资讯中心
/
用 JAX 与 Flax NNX 从零预训练 miniGPT:数据并行与张量并行的完整实战
用 JAX 与 Flax NNX 从零预训练 miniGPT:数据并行与张量并行的完整实战
发布时间:2026/9/17 6:19:10
用 JAX 与 Flax NNX 从零预训练 miniGPT数据并行与张量并行的完整实战【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax本篇指南以 Flax 官方示例 docs_nnx/examples/minigpt.md 为核心骨架演示如何用 JAX、Flax NNX 与 Optax 在 TPU 上完成一个小型 GPTminiGPT语言模型的预训练。你将从零学会定义可自动并行化的 miniGPT 模型、用 Grain 与 Tiktoken 加载并预处理 TinyStories 数据集、编写带nnx.jit与nnx.value_and_grad的训练步骤、用 TensorBoard Profiler 对不同 batch size 与并行策略做性能剖析最终把训练好的模型保存为 Orbax checkpoint。背景用 JAX 的自动并行化做 SPMD 训练本教程的核心技术点是 JAX 的设备并行device parallelism能力它服务于 Single-Program Multi-DataSPMD这一并行编程模型。JAX 提供两种层次的并行方式二者可以自由叠加数据并行Data Parallelism把训练数据按 batch 维度切分称为 sharding到多块 GPU 或 Google TPU 上让各设备同时处理不同的数据子集。这样可以使用更大的 batch size从而显著加速训练。张量并行Tensor Parallelism把模型参数张量本身切分到多块设备上例如把一个(in_features, out_features)的权重矩阵按out_features维度切分从而突破单卡显存对模型规模的限制。在本示例中采用的是一种经典的组合策略4 路数据并行 2 路张量并行即下文jax.make_mesh((2, 1), (batch, model))所构造的 mesh对应 8 台设备。JAX 的自动并行机制会依据你给出的分片规范自动生成跨设备的计算图与通信无需手写任何pmap或xmap式的设备循环。搭建 Mesh并行拓扑的骨架要在 JAX 中切分数据首先需要创建一个jax.sharding.Mesh。mesh 本质上是 JAX 设备的多维 NumPy 数组其中每个轴都有一个名字例如x、y它封装了 TPU 资源在物理拓扑上的组织信息供编译器决定如何在设备之间分发计算。教程中使用jax.make_mesh创建带两个轴的 mesh_ jax.set_mesh(jax.make_mesh((2, 1), (batch, model)))第一个轴batch大小为 2用于数据并行第二个轴model大小为 1用于模型张量并行jax.set_mesh返回一个上下文管理器由于整个 notebook 都使用同一个 mesh这里直接忽略它即可。若只想临时设置当前 mesh可以把它用作with上下文。同时加载 GPT-2 的 BPE tokenizer来自 Tiktoken 库后续所有 tokenize / decode 都依赖它tokenizer tiktoken.get_encoding(gpt2)PartitionSpec 与*_metadata告诉编译器如何切分张量要让模型并行真正生效需要把分片意图写进模型的参数元数据中。JAX 提供jax.sharding.PartitionSpec通常简写为P来描述某个张量每个维度与 mesh 轴之间的对应关系。PartitionSpec是元组(x, y)的薄封装例如out_shardingP(x, y)表示数据的第一维沿着 mesh 的x轴切分第二维沿着y轴切分None表示该维度不切分整块复制。在 Flax NNX 中各层nnx.Linear、nnx.Embed、nnx.MultiHeadAttention、nnx.LayerNorm等的构造函数都接受形如kernel_metadata、bias_metadata、scale_metadata的参数字典。以 flax/nnx/nn/linear.py 中的nnx.Linear为例其签名包含kernel_metadata与bias_metadata在 flax/nnx/nn/attention.py#L600-L601 中MultiHeadAttention同样接收kernel_metadata、bias_metadata、out_kernel_metadata等。这些 metadata 会被写入参数变量作为初始化时的out_sharding提示JAX 编译器据此把参数放置在正确的设备分片上。对nnx.Linear而言权重kernel的形状是(in_features, out_features)kernel_metadata{out_sharding: P(None, model)}表示把out_features维度沿 mesh 的model轴切分——这正是张量并行中列并行的写法bias_metadata{out_sharding: P(model)}表示 bias 沿model轴切分与切分后的 kernel 列对齐。定义 miniGPT 模型TransformerBlock单个 Transformer 块每个 Transformer 块由多头自注意力与两层前馈网络MLP组成并遵循 Pre-LN 残差结构先 LayerNorm 再做残差相加class TransformerBlock(nnx.Module): A single Transformer block. Each Transformer block processes input sequences via self-attention and feed-forward networks. Args: embed_dim (int): Embedding dimensionality. num_heads (int): Number of attention heads. ff_dim (int): Dimensionality of the feed-forward network. rngs (flax.nnx.Rngs): A Flax NNX stream of JAX PRNG keys. rate (float): Dropout rate. Defaults to 0.1. def __init__(self, embed_dim: int, num_heads: int, ff_dim: int, *, rngs: nnx.Rngs, rate: float 0.1): self.mha nnx.MultiHeadAttention(num_headsnum_heads, in_featuresembed_dim, kernel_metadata{out_sharding: P(None, model)}, bias_metadata{out_sharding: P(model)}, rngsrngs) self.dropout1 nnx.Dropout(raterate, rngsrngs) self.layer_norm1 nnx.LayerNorm(epsilon1e-6, num_featuresembed_dim, scale_metadata{out_sharding: P(model)}, bias_metadata{out_sharding: P(model)}, rngsrngs) self.mlp nnx.Sequential( nnx.Linear(in_featuresembed_dim, out_featuresff_dim, kernel_metadata{out_sharding: P(None, model)}, bias_metadata{out_sharding: P(model)}, rngsrngs), nnx.relu, nnx.Linear(in_featuresff_dim, out_featuresembed_dim, kernel_metadata{out_sharding: P(None, model)}, bias_metadata{out_sharding: P(model)}, rngsrngs), nnx.Dropout(raterate, rngsrngs)) self.layer_norm2 nnx.LayerNorm(epsilon1e-6, num_featuresembed_dim, scale_metadata{out_sharding: P(model)}, bias_metadata{out_sharding: P(model)}, rngsrngs) # Apply the Transformer block to the input sequence. def __call__(self, inputs): # Instantiate the causal attention mask. attention_output self.mha( inputs_qinputs, is_causalTrue, decodeFalse ) attention_output self.dropout1(attention_output) out1 self.layer_norm1(inputs attention_output) ffn_output self.mlp(out1) return self.layer_norm2(out1 ffn_output)要点说明nnx.MultiHeadAttention是 Flax NNX 提供的内置多头注意力模块见 flax/nnx/nn/attention.py#L482is_causalTrue直接启用因果掩码保证第t个位置只能 attend 到前t个位置——这是自回归语言模型的关键约束decodeFalse表示非增量解码模式。nnx.Dropout需要传入rngs它会从 NNX 的随机数流中取dropout集合的 key。nnx.LayerNorm(epsilon1e-6, ...)的scale_metadata与bias_metadata都沿model轴切分与 attention 输出特征的切分方式保持一致。残差结构采用 Pre-LNinputs attention_output之后再做 LayerNorm这是 GPT-2 系列的经典布局。TokenAndPositionEmbeddingToken 与位置嵌入自注意力本身不具备位置信息因此需要把每个 token 的嵌入与其位置嵌入相加class TokenAndPositionEmbedding(nnx.Module): Combines token embeddings (words in an input sentence) with positional embeddings (the position of each word in a sentence). Args: maxlen (int): Matimum sequence length. vocal_size (int): Vocabulary size. embed_dim (int): Embedding dimensionality. rngs (flax.nnx.Rngs): A Flax NNX stream of JAX PRNG keys. def __init__(self, maxlen: int, vocab_size: int, embed_dim: int, *, rngs: nnx.Rngs): # Initialize token embeddings (using flax.nnx.Embed). # Each unique word has an embedding vector. self.token_emb nnx.Embed(num_embeddingsvocab_size, featuresembed_dim, rngsrngs) # Initialize positional embeddings (using flax.nnx.Embed). self.pos_emb nnx.Embed(num_embeddingsmaxlen, featuresembed_dim, rngsrngs) # Takes a token sequence (integers) and returns the combined token and positional embeddings. def __call__(self, x): # Generate a sequence of positions for the input tokens. positions jnp.arange(0, x.shape[1])[None, :] # Look up the positional embeddings for each position in the input sequence. position_embedding self.pos_emb(positions) # Look up the token embeddings for each token in the input sequence. token_embedding self.token_emb(x, out_shardingjax.typeof(x).sharding) # Combine token and positional embeddings. return token_embedding position_embeddingnnx.Embed对应源码 flax/nnx/nn/linear.py 中的Embed类其__call__(self, inputs, out_shardingNone)支持运行时指定输出分片。这里把 token 嵌入的输出分片设为与输入 token 相同的 shardingjax.typeof(x).sharding从而让嵌入表在 batch 维度上按数据并行切分位置嵌入基于序列位置索引跟随当前 batch 分片广播即可。MiniGPT组装完整模型并实现自回归生成class MiniGPT(nnx.Module): A miniGPT transformer model, inherits from flax.nnx.Module. Args: maxlen (int): Maximum sequence length. vocab_size (int): Vocabulary size. embed_dim (int): Embedding dimensionality. num_heads (int): Number of attention heads. feed_forward_dim (int): Dimensionality of the feed-forward network. num_transformer_blocks (int): Number of transformer blocks. Each block contains attention and feed-forward networks. rngs (nnx.Rngs): A Flax NNX stream of JAX PRNG keys. # Initialize miniGPT model components. def __init__(self, maxlen: int, vocab_size: int, embed_dim: int, num_heads: int, feed_forward_dim: int, num_transformer_blocks: int, rngs: nnx.Rngs): # Initiliaze the TokenAndPositionEmbedding that combines token and positional embeddings. self.embedding_layer TokenAndPositionEmbedding( maxlen, vocab_size, embed_dim, rngsrngs ) # Create a list of TransformerBlock instances. # Each block processes input sequences using attention and feed-forward networks. self.transformer_blocks nnx.Sequential(*[TransformerBlock( embed_dim, num_heads, feed_forward_dim, rngsrngs ) for _ in range(num_transformer_blocks)]) # Initialize the output flax.nnx.Linear layer producing logits over the vocabulary for next-token prediction. self.output_layer nnx.Linear(in_featuresembed_dim, out_featuresvocab_size, kernel_metadata{out_sharding: P(None, model)}, bias_metadata{out_sharding: P(model)}, rngsrngs) def __call__(self, inputs): # Pass the input tokens through the embedding_layer to get token embeddings. x self.embedding_layer(inputs) # Apply each transformer block sequentially to the embedded input x self.transformer_blocks(x) # Pass the output of the transformer blocks through the output layer, # and obtain logits for each token in the vocabulary (for next token prediction). return reshard(self.output_layer(x), jax.typeof(inputs).sharding) def sample_from(self, logits): logits, indices jax.lax.top_k(logits, ktop_k) logits nnx.softmax(logits) return jax.random.choice(jax.random.PRNGKey(0), indices, plogits) nnx.jit(donate_argnums(1,)) def generate_step(self, padded_tokens, sample_index): logits self(padded_tokens) next_token self.sample_from(logits[0][sample_index]) return next_token def generate_text(self, max_tokens, start_tokens): generated [] for i in range(max_tokens): sample_index len(start_tokens) len(generated) - 1 padded_tokens jnp.array((start_tokens generated [0] * (maxlen - len(start_tokens) - len(generated))))[None, :] next_token int(self.generate_step(padded_tokens, sample_index)) if next_token tokenizer.encode(|endoftext|, allowed_special{|endoftext|})[0]: break generated.append(next_token) return tokenizer.decode(start_tokens generated) # Creates the miniGPT model with 4 transformer blocks. def create_model(rngs): return MiniGPT(maxlen, vocab_size, embed_dim, num_heads, feed_forward_dim, num_transformer_blocks4, rngsrngs)值得注意的实现细节前向输出通过reshard(self.output_layer(x), jax.typeof(inputs).sharding)重新对齐到输入的数据并行分片保证后续在 batch 维度的规约与 loss 计算无需显式处理跨设备布局。generate_step用nnx.jit(donate_argnums(1,))装饰nnx.jit是 Flax NNX 的编译变换实现在 flax/nnx/transforms/compilation.py#L160donate_argnums指定把第 1 个位置参数padded_tokens的缓冲区捐献给计算从而在每次生成 token 时复用显存、避免重复分配。sample_from使用jax.lax.top_k(logits, ktop_k)只保留 logits 最大的top_k个候选再做 softmax 采样这是一种 Top-K 采样策略能显著抑制低质量候选。生成循环在遇到|endoftext|结束符时提前停止并用tokenizer.decode把 token 序列还原为文本。超参数配置vocab_size tokenizer.n_vocab num_transformer_blocks 8 maxlen 256 embed_dim 256 num_heads 8 feed_forward_dim 256 batch_size 144 num_epochs 1 top_k 10各参数含义vocab_size取 GPT-2 tokenizer 的完整词表大小num_transformer_blocks 8是堆叠的 Transformer 块数注意create_model中实际传了num_transformer_blocks4用于先快速验证生成流程maxlen 256是序列最大长度embed_dim 256是嵌入/隐藏维度num_heads 8是注意力头数feed_forward_dim 256是 MLP 中间维度batch_size 144是训练 batchnum_epochs 1只跑一个 epochtop_k 10是采样候选数。加载与预处理数据Grain Tiktoken Pandas数据加载使用 Google 的 Grain 库grain.python它是专为 JAX 生态设计的高性能数据管线。整个流程分为三步读取 TinyStories 文本 → 用TextDataset完成 tokenize 与 padding → 用pygrain.IndexSamplerpygrain.DataLoader产出 batch。dataclass class TextDataset: data: list maxlen: int def __len__(self): return len(self.data) def __getitem__(self, idx: int): # Use Tiktoken for tokenization encoding tokenizer.encode(self.data[idx], allowed_special{|endoftext|})[:self.maxlen] # Tokenize and truncate return encoding [0] * (self.maxlen - len(encoding)) # Pad to maxlen def load_and_preprocess_data(file_path, batch_size, maxlen): with open(file_path, r) as f: text f.read() stories text.split(|endoftext|) stories [story|endoftext| for story in stories if story.strip()] df pd.DataFrame({text: stories}) data df[text].dropna().tolist() dataset TextDataset(data, maxlen) sampler pygrain.IndexSampler( len(dataset), shuffleFalse, seed42, shard_optionspygrain.NoSharding(), num_epochsnum_epochs, ) dl pygrain.DataLoader( data_sourcedataset, samplersampler, operations[pygrain.Batch(batch_sizebatch_size, drop_remainderTrue)], ) return dl text_dl load_and_preprocess_data(TinyStories-train.txt, batch_size, maxlen)数据集的获取与处理说明TinyStories 数据集从 Hugging Face 下载roneneldan/TinyStories的TinyStories-train.txt本教程只使用训练集拆分。下载命令为!wget https://huggingface.co/datasets/roneneldan/TinyStories/resolve/main/TinyStories-train.txt?downloadtrue -O TinyStories-train.txtTextDataset.__getitem__先用tokenizer.encode(..., allowed_special{|endoftext|})把单条 story 编码为 token 序列|endoftext|作为特殊 token 保留截断到maxlen不足部分用0右填充0是 GPT-2 词表中|endoftext|的 id这里兼作 pad token。pygrain.IndexSampler负责按索引采样shuffleFalse、seed42保证可复现NoSharding表示当前不按多主机分片num_epochs控制遍历轮数。pygrain.DataLoader通过pygrain.Batch(batch_sizebatch_size, drop_remainderTrue)把样本组装成 batchdrop_remainderTrue丢弃末尾不足一个 batch 的样本。定义损失函数与训练步骤训练步骤是整条数据管线的心脏它把前向、反向、指标更新和参数更新全部封装进一个被nnx.jit编译的函数中# Defines the loss function using optax.softmax_cross_entropy_with_integer_labels. def loss_fn(model, batch): logits model(batch[0]) loss optax.softmax_cross_entropy_with_integer_labels(logitslogits, labelsbatch[1]).mean() return loss, logits # Define the training step with the flax.nnx.jit transformation decorator. nnx.jit(donate_argnums(0, 1, 3)) def train_step(model: MiniGPT, optimizer: nnx.Optimizer, metrics: nnx.MultiMetric, batch): grad_fn nnx.value_and_grad(loss_fn, has_auxTrue) (loss, logits), grads grad_fn(model, batch) metrics.update(lossloss, logitslogits, lablesbatch[1]) optimizer.update(model, grads)对应的底层实现loss_fn使用optax.softmax_cross_entropy_with_integer_labels(logits, labels)计算逐 token 的交叉熵labels 是整数 token id免去 one-hot再取mean()得到标量 loss同时把logits作为辅助输出返回供指标计算使用。nnx.value_and_grad(loss_fn, has_auxTrue)是 Flax NNX 的自动微分变换见 flax/nnx/transforms/autodiff.py#L440has_auxTrue表示loss_fn返回(loss, aux)元组于是得到((loss, logits), grads)。optimizer.update(model, grads)对应nnx.Optimizer.update见 flax/nnx/training/optimizer.py#L178其内部先调用tx.update(grads, opt_state)再optax.apply_updates一次性完成模型参数与优化器状态的更新。nnx.MultiMetric.update(...)原地更新各子指标源码见 flax/nnx/training/metrics.py#L414。nnx.jit(donate_argnums(0, 1, 3))把model、optimizer、batch的缓冲区全部捐献训练循环中这些对象不再被复用从而最大化显存复用率Flax 0.11 起nnx.Optimizer不再持有model属性模型需作为独立参数传入参见 flax/nnx/training/optimizer.py#L168-L175。训练循环从随机乱码到像样的微型故事首先实例化模型、优化器与指标model create_model(rngsnnx.Rngs(0))optimizer nnx.Optimizer(model, optax.adam(1e-3), wrtnnx.Param) metrics nnx.MultiMetric( lossnnx.metrics.Average(loss), ) rng jax.random.PRNGKey(0)nnx.Optimizer(model, optax.adam(1e-3), wrtnnx.Param)用 Adam学习率 1e-3优化模型的Param类变量权重与偏置wrtnnx.Param是一个 filter指明优化器只跟踪、更新参数变量参考 flax/nnx/training/optimizer.py#L131-L166 中wrt参数说明。nnx.MultiMetric(lossnnx.metrics.Average(loss))包裹一个Average指标见 flax/nnx/training/metrics.py#L58Average(loss)表示update时通过关键字loss...传入待平均的值MultiMetric.update会把关键字参数透传给各子指标compute()返回各指标的字典。正式训练前可以先做一次热身生成验证未训练模型此时输出近乎随机start_prompt Once upon a time start_tokens tokenizer.encode(start_prompt)[:maxlen] model.generate_text(maxlen, start_tokens)接着是核心训练循环。为了数据并行必须把训练数据沿batch轴切分这里用jax.device_put(..., P(batch, None))显式指定分片同时用jax.vmap变换一次性生成所有样本的 target 序列避免 Python 层逐样本循环metrics_history { train_loss: [], } prep_target_batch jax.vmap( lambda tokens: jnp.concatenate((tokens[1:], jnp.array([0]))) ) step 0 for epoch in range(num_epochs): start_time time.time() for batch in text_dl: if len(batch) % len(jax.devices()) ! 0: continue # skip the remaining elements input_batch jnp.stack(batch).T target_batch prep_target_batch(input_batch) train_step( model, optimizer, metrics, jax.device_put( (input_batch, target_batch), P(batch, None) ), ) if (step 1) % 200 0: for metric, value in metrics.compute().items(): metrics_history[ftrain_{metric}].append(value) metrics.reset() elapsed_time time.time() - start_time print( f\n\nStep {step 1}, Loss: {metrics_history[train_loss][-1]}, Elapsed Time: {elapsed_time:.2f} seconds ) start_time time.time() print(Generated text:) print(model.generate_text(maxlen, start_tokens)) step 1 # Final text generation print(Final generated text:) generated_text model.generate_text(maxlen, start_tokens)训练循环的若干细节jnp.stack(batch).T把 Grain 产出的(batch, maxlen)数组转置为(maxlen, batch)布局匹配后续jax.device_put中P(batch, None)的 batch 维度。prep_target_batch对每个输入序列做右移一位tokens[1:]拼上结束符[0]从而构造出自回归训练所需的输入-目标对。if len(batch) % len(jax.devices()) ! 0: continue跳过无法被设备数整除的尾部 batch保证每个 batch 都能被均匀切分。每 200 步输出一次训练 loss 与耗时并现场用generate_text展示模型当前的生成能力。训练完成后可视化 loss 曲线import matplotlib.pyplot as plt plt.plot(metrics_history[train_loss]) plt.title(Training Loss) plt.xlabel(Step % 200) plt.ylabel(Loss) plt.show()从实际效果看模型会从最初生成完全随机的单词逐渐过渡到能生成通顺的微型故事——本质上就是在一个很小的规模上完成了一次大语言模型的预训练。保存 CheckpointOrbax 一键持久化训练结束后用 Orbax 保存模型的全部状态nnx.state(model)提取所有变量为 pytreeimport orbax.checkpoint as orbax from pathlib import Path state nnx.state(model) checkpoint_path Path(checkpoint).resolve() checkpointer orbax.PyTreeCheckpointer() checkpointer.save(checkpoint_path, argsorbax.args.PyTreeSave(state), forceTrue)nnx.state(model)返回一个State对象其中包含模型全部参数、BatchStat、优化器需要持久化的变量等按变量类型划分可配合 filter 筛选。orbax.PyTreeCheckpointer.save配合orbax.args.PyTreeSave把该状态写入checkpoint/目录forceTrue允许覆盖已存在的 checkpoint。用 Profiling 做超参数调优在大规模训练中光看 loss 是不够的还需要剖析设备利用率与单步耗时来定位瓶颈。教程给出了完整的剖析脚手架搭建可复用的剖析流程由于模型会被反复运行以比较不同配置先做一次 warmup确保代码已 JIT 编译、TPU 已预热再开始正式 tracing以保证对比的公平性trace_dir /tmp/jax-trace/ def loop_step(batch, step): input_batch jnp.stack(batch).T target_batch prep_target_batch(input_batch) train_step(model, optimizer, metrics, jax.device_put((input_batch, target_batch), P(batch, None))) def generate_trace(): tracing_steps 30 warmup_steps 5 for current_step in range(warmup_steps tracing_steps): if current_step warmup_steps: jax.profiler.start_trace(trace_dir) with jax.profiler.StepTraceAnnotation(train, step_numcurrent_step): batch next(text_dl) loop_step(batch, current_step) jax.profiler.stop_trace()jax.profiler.start_trace/stop_trace负责采集 XLA 执行轨迹StepTraceAnnotation给每一步打上语义标签便于在 TensorBoard 的 Trace Viewer 中按 step 定位。对比不同 batch size第一组实验比较 batch size 对训练吞吐的影响重新构建数据管线会耗时几分钟trace_dir /tmp/jax-trace-batch-comparison/ batch_size 64 text_dl iter(load_and_preprocess_data(TinyStories-train.txt, batch_size, maxlen)) generate_trace() batch_size 256 text_dl iter(load_and_preprocess_data(TinyStories-train.txt, batch_size, maxlen)) generate_trace()用 TensorBoard 的 Profiler 插件对比两组 trace%tensorboard --logdir $trace_dir --port 6006对比运行结果时需要重点关注两个关键指标Framework Op Placement框架算子在设备上的放置率与Average Step Time平均单步耗时。本教程给出的实测结论是把 batch size 从 64 提升到 256FLOPS 利用率从 16% 提升到 27%平均单步耗时从 100ms 增加到 260ms但 batch size 放大了 300%换算到每个训练样本的耗时反而从 1.5ms 降到 1.02ms。也就是说更大的 batch size 同时改善了设备利用率和样本级吞吐。对比不同的并行策略接下来比较并行策略。此前使用的是 4 路数据并行 2 路模型并行另一种常见做法是纯 8 路数据并行。切换方式极其简单——只需替换 meshjax.make_mesh((8, 1), (batch, model))JAX 会自动重新计算模型与数据的分片方案其余代码无需任何改动。重新连接 TPU runtime 后跑一遍对比trace_dir /tmp/jax-trace-parallelism-comparison/ mesh Mesh(mesh_utils.create_device_mesh((4, 2)), (batch, model)) generate_trace() mesh Mesh(mesh_utils.create_device_mesh((8, 1)), (batch, model)) generate_trace()再次运行 TensorBoard%tensorboard --logdir$trace_dir实测结果显示两种策略的单步耗时几乎相同但 8 路数据并行的 FLOPS 利用率只有 13%而 4 路数据 2 路模型并行达到 27%。借助 Trace Viewer 逐块 TPU 检查算子可以发现 8 路数据并行下 TPU 大量时间处于空闲等待 host 下发数据并且耗费了大量时间在reduce_sum这类跨设备规约上——这就是利用率偏低的原因。这个例子说明通过改变超参数并对比 profile可以获得关于训练瓶颈与硬件限制的深刻洞察。batch size 与并行策略只是众多可调超参数中的两个代表学习率、序列长度、注意力头数、Transformer 层数等都会对训练速度和资源利用率产生显著影响。小结本教程完整走通了一条从数据到可部署模型的 miniGPT 预训练链路用jax.make_mesh声明并行拓扑、用PartitionSpec与各层*_metadata表达张量切分意图、用 Flax NNX 的nnx.Module体系组织模型、用 Grain Tiktoken 搭建数据管线、用nnx.jitnnx.value_and_gradnnx.Optimizer封装训练步骤最后用 Orbax 保存 checkpoint并用 TensorBoard Profiler 对 batch size 与并行策略做定量剖析。全程最令人印象深刻的是 JAX 自动并行化的魔法切换数据并行与张量并行的配比往往只需要改一行jax.make_mesh的调用。示例完整代码可在仓库 docs_nnx/examples/minigpt.md或同目录的 minigpt.ipynb中查看想深入理解其中用到的 NNX 原语可继续阅读 flax/nnx/README.md、flax/nnx/transforms/compilation.py、flax/nnx/transforms/autodiff.py、flax/nnx/training/optimizer.py 与 flax/nnx/training/metrics.py。【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考