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

基于广播实现含多级索引列的矩阵乘法通用方法求解

通用矩阵点乘函数实现(适配一维/多维列索引)

需求

实现一个通用函数,对两个矩阵执行点乘(.dot,非元素级乘法.mul),其中一个矩阵的列可能为一维普通索引或二维MultiIndex(多级索引),需整合现有分别适配两种场景的函数。

样本数据

import pandas as pd

# 矩阵A
A = pd.DataFrame({'b': [1, 0], 'c': [1, 0], 'e': [0, 1]}, index=['a','d'])
print(A)
# 输出:
#    b  c  e
# a  1  1  0
# d  0  0  1

# 带二维MultiIndex列的矩阵B2D
columns = pd.MultiIndex.from_product([['b', 'c', 'e'], ['metric1', 'metric2']])
B2D = pd.DataFrame(
    [
        [22, 24, 20, 31, 29, 20],
        [12, 14, 10, 21, 24, 91]
    ],
    columns=columns
)
print(B2D)
# 输出:
#         b               c               e        
#   metric1 metric2 metric1 metric2 metric1 metric2
# 0      22      24      20      31      29      20
# 1      12      14      10      21      24      91

# 一维列索引的矩阵B1D
B1D = B2D.xs('metric1', 1, 1)
print(B1D)
# 输出:
#    b   c   e
# 0  22  20  29
# 1  12  10  24

期望结果

# 传入二维MultiIndex列的矩阵
print(func(A, B2D))
# 输出:
#         a               d        
#   metric1 metric2 metric1 metric2
# 0      42      55      29      20
# 1      22      35      24      91

# 传入一维列索引的矩阵
print(func(A, B1D))
# 输出:
#     a   d
# 0  42  29
# 1  22  24

现有实现(分场景)

以下函数分别适配两种场景,但无法通用:

def func1D(a, b):
    return (a @ b.T).T

def func2D(a, b):
    return (b.stack(level=1) @ a.T).unstack()

通用解决方案

通过判断矩阵列索引的类型,动态选择对应的计算逻辑,整合为一个通用函数:

def func(a, b):
    # 判断b的列是否为MultiIndex
    if isinstance(b.columns, pd.MultiIndex):
        # 处理二维多级索引场景
        return (b.stack(level=1) @ a.T).unstack()
    else:
        # 处理一维普通索引场景
        return (a @ b.T).T

逻辑说明

  • 先检查输入矩阵b的列索引类型:
    • 若为MultiIndex,则执行stack将次级索引转为行维度,完成点乘后再unstack恢复多级列索引;
    • 若为普通一维索引,则直接通过转置配合点乘完成计算。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.03 08:25:30