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

如何在JAX中监控vmapped函数的执行进度(类似tqdm的功能)

如何在JAX中监控vmapped函数的执行进度(类似tqdm的功能)

这个问题确实挺常见的——JAX的jit和vmap是把整个操作编译成一个整体批量执行的,不像普通Python循环那样能随时插入打印或进度条更新,不过我们有两种实用的办法来解决这个问题:

方法一:分块处理(性能友好型)

既然直接在vmapped的jit函数里没法插进度监控,我们可以把大数组拆成小块,在Python层面循环处理每一块,每处理完一块就更新一次进度。这种方法几乎不会损失JAX的编译优化优势,还能正常用tqdm显示进度。

示例代码:

import jax
import jax.numpy as jnp
from tqdm import tqdm

def f(x):
    # 模拟你的昂贵计算操作
    return x ** 2

xarr = jnp.arange(1000)
batch_size = 100  # 根据你的计算资源调整块大小
num_batches = len(xarr) // batch_size

# 提前编译好处理单块的函数,避免循环中重复编译
batched_process = jax.jit(jax.vmap(f))

output_blocks = []
# 用tqdm包裹循环,实时显示进度
for i in tqdm(range(num_batches)):
    current_batch = xarr[i * batch_size : (i+1) * batch_size]
    output_blocks.append(batched_process(current_batch))

# 处理数组长度不能被块大小整除的剩余部分
if len(xarr) % batch_size != 0:
    remaining_batch = xarr[num_batches * batch_size :]
    output_blocks.append(batched_process(remaining_batch))

# 合并所有块的结果
final_output = jnp.concatenate(output_blocks)

这种方法的好处是:进度条更新清晰,编译开销极低(只编译一次块处理函数),几乎不影响原本的计算性能。唯一的小缺点是需要手动处理分块和结果合并,代码量稍微多一点。

方法二:使用JAX主机回调(细粒度监控)

如果你需要更细粒度的逐元素进度更新,可以用JAX的jax.experimental.host_callback.call功能,在vmap的每个元素处理时插入主机端的回调函数来更新进度。不过要注意,这种方法会带来一定的性能开销,因为每个元素处理都要和主机交互,适合计算本身非常昂贵、这点开销可以接受的场景。

示例代码:

import jax
import jax.numpy as jnp
from tqdm import tqdm

def f(x):
    # 模拟你的昂贵计算操作
    return x ** 2

xarr = jnp.arange(1000)
# 初始化进度条
pbar = tqdm(total=len(xarr))

def update_progress(input_x):
    # 每次处理一个元素就更新进度条
    pbar.update(1)
    return input_x  # 返回输入,不干扰原计算逻辑

# 给原函数加上进度回调
def f_with_progress(x):
    # 插入主机回调,指定返回的形状和类型和输入一致
    x = jax.experimental.host_callback.call(
        update_progress,
        x,
        result_shape=jax.ShapeDtypeStruct(x.shape, x.dtype)
    )
    return f(x)

# 正常执行vmapped+jitted的函数
final_output = jax.jit(jax.vmap(f_with_progress))(xarr)
# 记得关闭进度条
pbar.close()

需要注意的是,JAX的vmap是向量化并行执行的,元素的处理顺序可能和数组顺序不完全一致,所以进度条的更新可能不是严格线性的,但整体能准确反映完成的比例。

小总结

  • 如果追求性能优先,选分块处理的方法,既不怎么影响计算速度,又能清晰监控进度;
  • 如果需要逐元素的细粒度监控,并且可以接受一点性能损失,就用主机回调的方法。

内容来源于stack exchange

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.07 11:09:34