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

JAX两种scan实现性能差异原因及潜在影响咨询

JAX scan两种实现的速度差异原因及精度影响分析

我在项目中频繁使用JAX的scan函数,偶然发现按首维度扫描数组(示例中的normal_scan)比基于索引扫描(scan_with_index)速度更慢。想请教该现象的原因,以及这种索引扫描方式是否存在精度等负面影响?

import jax
import jax.numpy as jnp
import time

@jax.jit
def normal_scan(x0, arr):

    def body(state, input):
        x = state    
        a = input
        new_state = x + a**2 + jnp.sin(a)
        return (new_state, x)
    state = (x0)
    input = arr
    result = jax.lax.scan(body, state, input)
    return result

@jax.jit
def scan_with_index(x0, arr):
    N = len(arr)
    def body(state, input):
        x, ind = state
        new_state = (x+arr[ind]**2 + jnp.sin(arr[ind]), ind + 1)
        return (new_state, x)
    state = (x0, 0)

    result = jax.lax.scan(body, state, length=N)
    return result    

if __name__ == "__main__":
    key = jax.random.key(0) 
    N = 100                  

    arr = jax.random.normal(key, (N, 2))    
    x0 = jnp.array([1.0, 2.0])
    
    # warm up
    for i in range(2):
        result1 = normal_scan(x0, arr)
        result2 = scan_with_index(x0, arr)

    start_time = time.time()
    for i in range(100):
        result1 = normal_scan(x0, arr)
        result1[0][0].block_until_ready()
    end_time = time.time()    
    
    print(f"Execution time: {end_time - start_time:.4f} seconds")
    # around 0.08

    start_time = time.time()
    for i in range(100):
        result2 = scan_with_index(x0, arr)
        result2[0][0].block_until_ready()
    end_time = time.time()   

    print(f"Execution time: {end_time - start_time:.4f} seconds")
    # around 0.04

一、速度差异的原因

  • 输入数据传递开销不同:

    • normal_scan中,scan会把arr的每个元素作为input传给body函数,每次迭代都要提取对应维度的数据并完成传递。对于示例中(N,2)的数组,每次传递的是长度为2的向量,这种小数据的调度和传递开销占比很高。
    • scan_with_index里,scan只靠length=N控制迭代次数,arr作为外部数组直接通过索引访问。JAX的XLA编译器能更高效地优化这种模式,比如把整个arr加载到设备内存后直接寻址,省去了每次迭代传递输入数据的额外开销。
  • 编译器优化逻辑不同:
    基于输入数组的scan模式需要处理序列展开和映射,编译器会生成更通用的调度逻辑;而索引式的写法更接近传统循环,XLA更容易识别并生成紧凑高效的机器码,减少了调度层面的冗余开销。

  • 小尺寸输入放大差异:
    示例中N=100、每个输入元素是2维向量,这种场景下数据传递的开销占比远大于计算开销,所以两种实现的速度差会很明显。如果N大幅增大或者每个输入元素的尺寸变大,计算开销会掩盖传递开销,速度差异会缩小。

二、索引扫描的精度影响

这种索引式扫描完全不会带来额外的精度损失,理由如下:

  • 两种实现的计算逻辑完全等价:不管是通过input传元素,还是arr[ind]索引取元素,最终执行的都是x + a**2 + jnp.sin(a),运算步骤和数值处理完全一致。
  • JAX的浮点计算是确定性的:只要计算逻辑相同,不管用哪种方式获取输入,得到的结果完全一致(你可以用jnp.allclose(result1[1], result2[1])验证,结果肯定是True)。
  • 索引操作本身无损耗:JAX里的数组索引就是直接的内存寻址,不会对数值做任何修改或引入精度误差。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 04:03:18