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

如何在Python类中优雅定义jax.vmap向量化函数以避免重复定义?

如何在JAX类中优雅整合vmap向量化函数,避免重复定义

我希望将经jax.vmap向量化的函数作为类方法,但不确定该在类中何处定义它,核心目标是避免每次调用类方法时都重新定义该向量化函数。

以下是一个统计jnp.array中值出现次数的类的极简示例,包含非向量化和向量化版本:

import jax.numpy as jnp
import jax

class ValueCounter():

    def __init__(self): # 仅作示例完整性,未实际使用
        self.attribute_1 = None

    @staticmethod
    def _count_value_in_array( # 非向量化函数
        array: jnp.array, value: float
    ) -> jnp.array:
        """统计某个值在数组中的出现次数"""
        return jnp.count_nonzero(array == value)

    # 向量化函数定义处
    def count_values_in_array(self, array: jnp.array, value_array: jnp.array) -> jnp.array:
        """统计值数组中每个值在目标数组中的出现次数"""
        count_value_in_array_vec = jax.vmap(
            self._count_value_in_array, in_axes=(None, 0)
        ) # 每次调用方法时都会重新定义向量化函数,冗余
        return count_value_in_array_vec(array, value_array)

输入输出示例:

value_counter = ValueCounter()
value_counter.count_values_in_array(jnp.array([0, 1, 2, 2, 1, 1]), jnp.array([0, 1, 2]))

预期结果:

Array([1, 3, 2], dtype=int32)

但每次调用count_values_in_array时,向量化函数count_value_in_array_vec都会被重新定义,这显得多余。请问该如何更优雅地将向量化函数整合到类中?


解决方案1:在__init__中预初始化向量化函数

将vmap后的函数作为实例属性,在类初始化时仅创建一次,后续调用直接复用:

import jax.numpy as jnp
import jax

class ValueCounter():

    def __init__(self):
        self.attribute_1 = None
        # 初始化时创建向量化函数,仅执行一次
        self._count_value_in_array_vec = jax.vmap(
            self._count_value_in_array, in_axes=(None, 0)
        )

    @staticmethod
    def _count_value_in_array(array: jnp.array, value: float) -> jnp.array:
        """统计某个值在数组中的出现次数"""
        return jnp.count_nonzero(array == value)

    def count_values_in_array(self, array: jnp.array, value_array: jnp.array) -> jnp.array:
        """统计值数组中每个值在目标数组中的出现次数"""
        return self._count_value_in_array_vec(array, value_array)

解决方案2:用类属性直接定义向量化函数

直接在类级别定义vmap后的函数,类加载时完成初始化,所有实例共享同一个向量化函数:

import jax.numpy as jnp
import jax

class ValueCounter():

    def __init__(self):
        self.attribute_1 = None

    @staticmethod
    def _count_value_in_array(array: jnp.array, value: float) -> jnp.array:
        """统计某个值在数组中的出现次数"""
        return jnp.count_nonzero(array == value)

    # 类属性,仅在类定义时执行一次vmap转换
    _count_value_in_array_vec = jax.vmap(_count_value_in_array, in_axes=(None, 0))

    def count_values_in_array(self, array: jnp.array, value_array: jnp.array) -> jnp.array:
        return self._count_value_in_array_vec(array, value_array)

解决方案3:用静态方法封装向量化逻辑

如果需要对向量化函数添加额外说明或逻辑,可以将其封装为静态方法:

import jax.numpy as jnp
import jax

class ValueCounter():

    def __init__(self):
        self.attribute_1 = None

    @staticmethod
    def _count_value_in_array(array: jnp.array, value: float) -> jnp.array:
        """统计某个值在数组中的出现次数"""
        return jnp.count_nonzero(array == value)

    @staticmethod
    def _count_value_in_array_vec(array: jnp.array, value_array: jnp.array) -> jnp.array:
        """向量化版本:统计值数组中每个值在目标数组中的出现次数"""
        return jax.vmap(ValueCounter._count_value_in_array, in_axes=(None, 0))(array, value_array)

    def count_values_in_array(self, array: jnp.array, value_array: jnp.array) -> jnp.array:
        return self._count_value_in_array_vec(array, value_array)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 20:04:59