Jax/Flax实现GRU前向传播速度远低于PyTorch的问题咨询
背景
最近使用Jax实现了一个双层GRU网络,性能极差无法正常使用,因此设计了简单测试与PyTorch进行速度对比,以下是测试相关信息:
最小可复现示例
测试基于Google Colab GPU运行环境
import flax.linen as jnn import jax import torch import torch.nn as tnn import numpy as np import jax.numpy as jnp def keyGen(seed): key1 = jax.random.PRNGKey(seed) while True: key1, key2 = jax.random.split(key1) yield key2 key = keyGen(1) hidden_size=200 seq_length = 1000 in_features = 6 out_features = 4 batch_size = 8 class RNN_jax(jnn.Module): @jnn.compact def __call__(self, x, carry_gru1, carry_gru2): carry_gru1, x = jnn.GRUCell()(carry_gru1, x) carry_gru2, x = jnn.GRUCell()(carry_gru2, x) x = jnn.Dense(4)(x) x = x/jnp.linalg.norm(x) return x, carry_gru1, carry_gru2 class RNN_torch(tnn.Module): def __init__(self, batch_size, hidden_size, in_features, out_features): super().__init__() self.gru = tnn.GRU( input_size=in_features, hidden_size=hidden_size, num_layers=2 ) self.dense = tnn.Linear(hidden_size, out_features) self.init_carry = torch.zeros((2, batch_size, hidden_size)) def forward(self, X): X, final_carry = self.gru(X, self.init_carry) X = self.dense(X) return X/X.norm(dim=-1).unsqueeze(-1).repeat((1, 1, 4)) rnn_jax = RNN_jax() rnn_torch = RNN_torch(batch_size, hidden_size, in_features, out_features) Xj = jax.random.normal(next(key), (seq_length, batch_size, in_features)) Yj = jax.random.normal(next(key), (seq_length, batch_size, out_features)) Xt = torch.from_numpy(np.array(Xj)) Yt = torch.from_numpy(np.array(Yj)) initial_carry_gru1 = jnp.zeros((batch_size, hidden_size)) initial_carry_gru2 = jnp.zeros((batch_size, hidden_size)) params = rnn_jax.init(next(key), Xj[0], initial_carry_gru1, initial_carry_gru2) def forward(params, X): carry_gru1, carry_gru2 = initial_carry_gru1, initial_carry_gru2 Yhat = [] for x in X: # x.shape = (batch_size, in_features) yhat, carry_gru1, carry_gru2 = rnn_jax.apply(params, x, carry_gru1, carry_gru2) Yhat.append(yhat) # y.shape = (batch_size, out_features) #return jnp.concatenate(Y, axis=0) jitted_forward = jax.jit(forward)
测试结果
# 未编译Jax版本 %time forward(params, Xj)
CPU耗时: 用户态7分17秒, 系统态8.18秒, 总耗时7分25秒 墙上时间:7分17秒
# 编译耗时 %time jitted_forward(params, Xj)
CPU耗时: 用户态8分9秒, 系统态4.46秒, 总耗时8分13秒 墙上时间:8分12秒
# 编译后Jax版本 %timeit jitted_forward(params, Xj)
最慢运行耗时是最快的204.20倍,推测存在中间结果缓存。10000轮测试,5轮最优值为单轮115 µs
# PyTorch版本 %timeit lambda: rnn_torch(Xt)
10000000轮测试,5轮最优值为单轮65.7 ns
问题解答
1. Jax版本运行速度慢的原因
- 实现逻辑存在明显差异:PyTorch侧调用的是
tnn.GRU,是针对全序列运算做了底层融合、并行优化的成熟算子;JAX侧是手动用Python原生for循环逐时间步调用GRUCell,未开启JIT时就是纯Python解释器执行每一步运算,没有任何优化,因此耗时极高。 - PyTorch测速逻辑完全错误:测试代码中
%timeit lambda: rnn_torch(Xt)只是统计创建lambda匿名函数的耗时,根本没有实际执行模型推理,65.7ns的结果没有任何参考价值,修正为%timeit rnn_torch(Xt)后实际运行耗时会和JAX编译后的结果处于同一量级。 - 额外冗余开销:forward函数未设置返回值,但JAX在计算图追踪阶段仍会记录所有中间操作,没有剪枝优化的空间。
2. Jax编译耗时长的原因
JAX的JIT编译默认会把Python侧的控制流完全展开为静态计算图,你写的1000步for循环会被展开为1000组重复的GRUCell、Dense运算节点,编译阶段需要处理数千个算子的优化、代码生成,因此耗时极长。如果要避免该问题,应该用jax.lax.scan算子实现序列循环,JAX可对scan做专门优化,不需要展开全量循环节点。
优化建议
- 替换手动for循环为
jax.lax.scan实现序列遍历,可大幅降低编译时间,同时得到优化后的序列运算性能。 - 直接使用Flax提供的
flax.linen.GRU层,和PyTorch的tnn.GRU功能对应,直接输入全序列即可,不需要手动处理循环逻辑。 - 修正PyTorch测速代码,排除测速逻辑错误的干扰。
内容的提问来源于stack exchange,提问作者Simon B
相关产品推荐
相关产品推荐

