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

如何在Jax中高效映射函数处理形状不一致的多参数元组/列表

处理JAX中形状不一致的可迭代对象批量计算

你的场景中,xs和ys是元素形状不同的列表,jax.vmap确实不适用(它要求输入具有统一的形状结构,用于对指定维度做批量映射)。要高效替代手动for循环,可以使用JAX内置的jax.tree_util.tree_map,它专门用于处理结构一致但元素形状可不同的树形/序列结构,自动遍历对应元素并应用函数。

实现代码

from jax import numpy as jnp
from jax import tree_util

xs = [jnp.zeros((1, 3)), jnp.zeros((3, 2, 3))]
ys = [jnp.ones((1, 3)), jnp.ones((3, 2, 3))]

def f(x, y):
    return jnp.sum(x - y)

# 用tree_map替代for循环
res = tree_util.tree_map(f, xs, ys)

说明

  • tree_map会自动遍历xs和ys的对应位置元素,对每一对(x, y)调用f函数,最终返回的res结构与输入的列表结构完全一致,结果和手动for循环完全相同。
  • 这种方式无需手动编写循环逻辑,且依托JAX的内部优化,执行效率与手动循环相当,同时更符合JAX的函数式编程风格。
  • 如果需要将结果转为数组(这里所有结果都是标量,转数组无问题),可以后续调用jnp.array(res)。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.29 19:20:02