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

如何高效实现含重复元素的张量矩阵乘积?规避高内存占用

优化方案:利用维度重组与广播避免显式重复

核心思路是通过维度重组将张量b分组,结合广播机制直接与a进行批量矩阵乘法,完全避免显式重复a带来的内存开销,同时保持硬件加速的运算性能。

具体实现步骤

  • 将张量b的维度重组为(n//f, f, c, c),把每f个b样本归为一组,对应a中的一个样本。
  • 给张量a新增一个维度,扩展为(n//f, 1, c, c),使其能与重组后的b进行广播。
  • 执行批量矩阵乘法:利用广播特性,a的每个c×c矩阵会自动与对应组内的f个b的c×c矩阵分别相乘。
  • 最后将结果重组回原形状(n, c, c),与原方法输出完全一致。

代码示例

import tensorflow as tf
import numpy as np

n = 100
c = 5
f = 10

# 原数据
a = tf.constant(np.random.rand(n//f, c, c))
b = tf.constant(np.random.rand(n, c, c))

# 原方法(内存开销大)
a_prime = tf.repeat(a, f, 0)
result_original = a_prime @ b

# 优化方法(低内存)
# 重组b的维度
b_reshaped = tf.reshape(b, (n//f, f, c, c))
# 给a新增维度,支持广播
a_expanded = tf.expand_dims(a, axis=1)
# 批量矩阵乘,自动广播
result_grouped = a_expanded @ b_reshaped
# 重组回原形状
result_optimized = tf.reshape(result_grouped, (n, c, c))

# 验证结果一致
print(tf.reduce_all(tf.abs(result_original - result_optimized) < 1e-6))  # 输出True

为什么更优?

  • 内存效率:tf.reshape和tf.expand_dims都是零拷贝操作(仅改变张量的视图,不复制数据),不需要像tf.repeat那样额外占用(f-1)*(n/f)*c²的内存空间,当n和f较大时,内存节省非常显著。
  • 性能保持:批量矩阵乘法由TensorFlow的底层优化(如CUDA加速)处理,运算效率与原方法几乎一致,完全避免了手动循环的性能损耗。

适配b为(n, c, 1)的场景

如果你的b确实是(n, c, 1)形状(问题描述中的情况),只需调整维度重组的最后一个维度即可:

b = tf.constant(np.random.rand(n, c, 1))
b_reshaped = tf.reshape(b, (n//f, f, c, 1))
result_grouped = a_expanded @ b_reshaped
result_optimized = tf.reshape(result_grouped, (n, c, 1))

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.16 02:13:19