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

Numpy:如何实现无广播的数组按块乘法

如何让Numpy数组元素与对应维度的子数组相乘(非逐元素广播)

现有两个Numpy数组a和b,b的形状为(2, 2, 5),a的形状为(5,)。需要实现a的每个元素与b中对应的2x2数组相乘,而非像直接执行a * b那样,让a的每个元素与b中每个2x2数组的每个元素相乘。

以下代码展示了问题及期望结果:

import numpy as np

np.random.seed(1234)


class MyClass:
    def __mul__(self, other):
        print(f'{type(self).__name__} * {other}')
        return other

    def __rmul__(self, other):
        print(f'{other} * {type(self).__name__}')
        return other


a = np.full((5,), MyClass())
b = np.random.uniform(-1, 1, (2, 2, 5))

a * b
# MyClass * -0.6169610992422154
# MyClass * 0.24421754207966373
# ...
# MyClass * 0.5456532432247481
# MyClass * 0.7652823812722331

# 期望结果:
[ai * bi for ai, bi in zip(a, np.moveaxis(b, -1, 0))]
# MyClass * [[-0.6169611  -0.45481479]
#  [-0.28436546  0.12239237]]
# ...
# MyClass * [[ 0.55995162  0.75186527]
#  [-0.25949849  0.76528238]]

# 编辑补充:Guimoute提出的@运算符方案,效果类似a * b且引入加法,并非可行解
b @ a
# -0.6169610992422154 * MyClass
# 0.24421754207966373 * MyClass
# ...
# 0.5456532432247481 * MyClass
# 0.7652823812722331 * MyClass

c = np.split(np.moveaxis(b, -1, 0).reshape(-1, 2), 5)
# c现在是长度为5的Numpy数组列表

a * c
# 抛出异常
# ValueError: operands could not be broadcast together with shapes (5,) (5,2,2)

编辑说明:需要Numpy原生解决方案,不使用Python循环(如[ai * bi for ai, bi in zip(a, np.moveaxis(b, -1, 0))])。

补充验证示例:便于验证方案有效性,以下是验证函数示例:

import numpy as np

np.random.seed(1234)


class MyClass:
    def __mul__(self, other):
        return self

    def __rmul__(self, other):
        return self


def solution_involving_python_loop(a, b):
    return np.array([ai * bi for ai, bi in zip(a, np.moveaxis(b, -1, 0))])

def is_valid_solution(func):
    return func(a, b).ndim == 1


a = np.full((5,), MyClass())
b = np.random.uniform(-1, 1, (2, 2, 5))


print(is_valid_solution(solution_involving_python_loop))
# True

补充对比示例:区分标准Numpy乘法与期望的乘法类型:

import numpy as np

np.random.seed(1234)


class Symbol:
    def __init__(self, name):
        self.name = name

    def __repr__(self):
        return self.name

    def __mul__(self, other):
        return Mul(self, other)

    def __rmul__(self, other):
        return Mul(other, self)


class Mul:
    def __init__(self, a, b):
        self.a = a
        self.b = b

    def __repr__(self):
        return f'{self.a} * {self.b}'


a = np.array([Symbol(chr(i)) for i in range(ord('a'), ord('a') + 5)])
b = np.random.randint(0, 100, (2, 2, 5))

# 广播乘法结果
broadcast_result = a * b
print(broadcast_result)
# [[[a * 47 b * 83 c * 38 d * 53 e * 76]
#   [a * 24 b * 15 c * 49 d * 23 e * 26]]
#
#  [[a * 30 b * 43 c * 30 d * 26 e * 58]
#   [a * 92 b * 69 c * 80 d * 73 e * 47]]]

# 期望结果
desired_result = np.array([ai * bi for ai, bi in zip(a, np.moveaxis(b, -1, 0))])
for x in desired_result:
    print(x)
# a * [[47 24]
#  [30 92]]
# b * [[83 15]
#  [43 69]]
# c * [[38 49]
#  [30 80]]
# d * [[53 23]
#  [26 73]]
# e * [[76 26]
#  [58 47]]

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.24 20:09:22