如何使用jax.lax.scan改写嵌套循环统计偶数的代码?
用jax.lax.scan改写统计偶数数量的代码
原代码的功能是统计二维数组中偶数的个数,下面用jax.lax.scan实现等价逻辑:
方法一:展平数组后单次扫描
先把二维数组展平成一维,再用scan逐个元素判断并累积计数:
import jax import jax.numpy as jnp numbers = jnp.array([[3.0, 14.0], [15.0, -7.0], [16.0, -11.0]]) # 定义扫描更新函数:输入当前计数和元素,返回更新后的计数 def update_count(carry, num): # 判断元素是否为偶数(浮点数取模2等于0) is_even = jnp.equal(num % 2, 0.0) # 将布尔值转为整数,累加到计数上 return carry + is_even.astype(jnp.int32), None # 初始计数为0,展平数组后执行扫描 total_evens, _ = jax.lax.scan(update_count, 0, numbers.flatten()) print(total_evens) # 输出:2
方法二:嵌套扫描处理二维结构
如果需要保留原二维数组的行遍历逻辑,可以用两层scan嵌套:
import jax import jax.numpy as jnp numbers = jnp.array([[3.0, 14.0], [15.0, -7.0], [16.0, -11.0]]) def count_row_evens(carry, num): is_even = jnp.equal(num % 2, 0.0) return carry + is_even.astype(jnp.int32), None def scan_rows(total_carry, row): # 对单行内的元素扫描计数,累加到总计数 row_evens, _ = jax.lax.scan(count_row_evens, 0, row) return total_carry + row_evens, None total_evens, _ = jax.lax.scan(scan_rows, 0, numbers) print(total_evens) # 输出:2
关键说明
jax.lax.scan的核心是累积式循环:第一个参数是更新函数(接收累积值carry和当前元素,返回新的累积值与可选中间结果),第二个参数是初始累积值,第三个参数是要遍历的序列。- 两种方法都能得到和原代码一致的结果,展平扫描更简洁,嵌套扫描适合需要按行处理的场景。
- 由于JAX的操作可追踪、可优化,
scan会把循环转换为JAX能高效处理的形式,配合jax.jit编译后性能会更优。
内容的提问来源于stack exchange,提问作者pepazdepa
相关产品推荐
相关产品推荐

