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

Colab+Jax+GPU:为何单元执行耗时60秒但%%timeit仅显示70ms?

JAX GPU计算Mandelbrot集时编译耗时过长导致总运行时间远超预期

问题现象

  • 单次计算耗时符合预期:%%timeit显示GPU函数计算耗时70-80ms,但整个%%timeit单元(默认连续执行7次)实际耗时约60秒;直接执行函数生成600万像素图像,单元完成同样耗时约60秒。
  • 额外开销随计算量缩放:当函数包含1000次迭代时,%%timeit显示71ms但实际总耗时60秒;仅20次迭代时,%%timeit显示10ms但实际总耗时约10秒。

问题根源

你在run_jax_kernel中使用了Python原生for循环,JAX的JIT编译器会将这个循环完全展开成1000个独立的计算步骤,生成的计算图极其庞大,导致编译时间远超实际计算时间。%%timeit显示的是编译完成后的单次运行耗时,但整个单元的总时间包含了漫长的JIT编译过程,这就是总耗时远超预期的核心原因。

解决方案

使用JAX提供的jax.lax.fori_loop代替Python原生循环,它会被编译成高效的GPU循环结构,而非展开所有步骤,大幅缩短编译时间。

修改后的代码

import math
import numpy as np
import matplotlib.pyplot as plt
import jax
from jax import lax

assert len(jax.devices("gpu")) == 1

def run_jax_kernel(c, fractal):
    def loop_body(i, carry):
        z, fractal = carry
        z = z**2 + c
        diverged = jax.numpy.absolute(z) > 2
        diverging_now = diverged & (fractal == 1000)
        fractal = jax.numpy.where(diverging_now, i, fractal)
        return (z, fractal)
    
    z = c
    _, fractal = lax.fori_loop(0, 1000, loop_body, (z, fractal))
    return fractal

run_jax_gpu_kernel = jax.jit(run_jax_kernel, backend="gpu")

def run_jax_gpu(height, width):
    mx = -0.69291874321833995150613818345974774914923989808007473759199
    my = 0.36963080032727980808623018005116209090839988898368679237704
    zw = 4 / 1e3

    y, x = jax.numpy.ogrid[(my-zw/2):(my+zw/2):height*1j, (mx-zw/2):(mx+zw/2):width*1j]
    c = x + y*1j
    fractal = jax.numpy.full(c.shape, 1000, dtype=np.int32)
    return np.asarray(run_jax_gpu_kernel(c, fractal).block_until_ready())

效果验证

修改后,首次编译时间会大幅缩短(从几十秒降至几秒内),后续调用均使用缓存的编译结果,%%timeit的总运行时间会接近单次耗时乘以执行次数,生成图像的速度也会显著提升。

内容的提问来源于stack exchange,提问作者crazygringo

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.21 08:19:51