You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.28 05:45:04