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é
相关产品推荐
相关产品推荐

