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

JAX中如何重塑含不同形状数组的Pytree并合并前两维度?

合并JAX PyTree中数组的前两个维度并保留后续维度

你之前用(15, -1)的问题在于,-1会把数组剩余的所有维度直接展平成一维,自然不符合observations要保留250,23的需求。

解决思路很简单:针对每个数组动态构造目标形状——先把前两个维度的乘积作为新的第一个维度,再拼接原数组从第三个维度开始的所有形状。这样不管数组后面有多少个维度,都能完整保留。

直接用这个逻辑修改tree_map的lambda函数就行:

import jax
import jax.numpy as jnp

# 示例pytree(和你的结构一致)
my_pytree = {
    "observations": jnp.ones((5, 3, 250, 23)),
    "dones": jnp.zeros((5, 3, 250))
}

# 处理后的pytree
processed_pytree = jax.tree_map(
    lambda x: jnp.reshape(x, (x.shape[0] * x.shape[1],) + x.shape[2:]),
    my_pytree
)

# 检查结果形状
print(processed_pytree["observations"].shape)  # 输出 (15, 250, 23)
print(processed_pytree["dones"].shape)          # 输出 (15, 250)

这样处理后,不管pytree里的数组后续维度有几个,都能精准合并前两个维度,同时保留原有维度结构,完全适配你的需求。

内容的提问来源于stack exchange,提问作者Valentin Macé

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.12 12:35:25