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

如何使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.20 19:33:33