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

JAX中PyTree加权求和优化:规避循环提升性能

优化PyTree加权求和的性能方案

问题背景

我有一组结构完全一致的嵌套PyTree,需要基于给定的权重列表对这些PyTree做加权求和,输出结果的结构要和单个子PyTree保持一致。目前已有可行代码,但包含循环和多步骤操作,希望优化性能。

原实现代码

import jax
import jax.numpy as jnp

list_1 = [
    [jnp.asarray([[1, 2], [3, 4]]), jnp.asarray([2, 3])],
    [jnp.asarray([[1, 2], [3, 4]]), jnp.asarray([2, 3])],
]

list_2 = [
    [jnp.asarray([[2, 3], [3, 4]]), jnp.asarray([5, 3])],
    [jnp.asarray([[2, 3], [3, 4]]), jnp.asarray([5, 3])],
]

list_3 = [
    [jnp.asarray([[7, 1], [4, 4]]), jnp.asarray([6, 2])],
    [jnp.asarray([[6, 4], [3, 7]]), jnp.asarray([7, 3])],
]

weights = [1, 2, 3] 
pytree = [list_1, list_2, list_3]

weighted_pytree = [jax.tree_map(lambda tree: weight * tree, tree) for weight, tree in zip(weights, pytree)]
reduced = jax.tree_util.tree_map(lambda *args: sum(args), *weighted_pytree)

性能优化方案

原代码的问题在于会先创建所有加权后的中间PyTree,额外占用内存且增加计算步骤。下面的优化方案直接在单次tree_map操作中完成加权求和,避免中间结构的生成,同时利用JAX的向量化计算能力提升效率:

import jax
import jax.numpy as jnp

# 数据定义与原代码一致
list_1 = [
    [jnp.asarray([[1, 2], [3, 4]]), jnp.asarray([2, 3])],
    [jnp.asarray([[1, 2], [3, 4]]), jnp.asarray([2, 3])],
]

list_2 = [
    [jnp.asarray([[2, 3], [3, 4]]), jnp.asarray([5, 3])],
    [jnp.asarray([[2, 3], [3, 4]]), jnp.asarray([5, 3])],
]

list_3 = [
    [jnp.asarray([[7, 1], [4, 4]]), jnp.asarray([6, 2])],
    [jnp.asarray([[6, 4], [3, 7]]), jnp.asarray([7, 3])],
]

# 将权重转为JAX数组以支持向量化计算
weights = jnp.array([1, 2, 3]) 
pytree = [list_1, list_2, list_3]

# 单步完成加权求和
reduced = jax.tree_map(lambda *leaves: jnp.dot(weights, jnp.stack(leaves)), *pytree)

优化说明

  • 直接对所有PyTree对应位置的叶子节点进行操作:用jnp.stack把同一位置的所有叶子堆叠成数组,再通过jnp.dot和权重做内积,一步得到该位置的加权和。
  • 避免了中间weighted_pytree结构的生成,节省内存开销的同时减少了计算调度的额外成本。
  • 利用JAX对dot和stack操作的优化,相比原代码的sum操作能获得更好的向量化性能。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.05 12:22:11