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

PyTorch中重复第二矩阵的矩阵乘法及多进程方案问询

PyTorch矩阵乘法实现方案

一、批量运算无需复制B(最优方案)

你不需要复制B来匹配A的实例数,PyTorch的矩阵乘法支持广播机制,可以直接处理维度匹配的批量运算:

  • 当A形状为(10, 384),B形状为(384, 39000)时,直接执行矩阵乘法就能得到(10, 39000)的结果:
import torch

# 示例张量
A = torch.randn(10, 384)
B = torch.randn(384, 39000)

# 直接计算,自动适配批量
result = A @ B  # 或 torch.matmul(A, B)
print(result.shape)  # 输出 torch.Size([10, 39000])

这种方式完全不需要复制B,能极大节省内存,尤其适合实例数达10万的场景。

二、强制复制B的实现方法

如果确实需要将B扩展为(10, 384, 39000)的形状,有两种方式:

1. 无内存开销的视图扩展(推荐)

用unsqueeze增加维度,再用expand创建共享内存的视图:

# 将B从(384, 39000)扩展为(10, 384, 39000)
B_expanded = B.unsqueeze(0).expand(10, -1, -1)
# 将A调整为(10, 1, 384),用bmm做批量矩阵乘法
result = torch.bmm(A.unsqueeze(1), B_expanded).squeeze(1)
print(result.shape)  # 输出 torch.Size([10, 39000])

expand不会实际复制数据,只是修改张量的维度信息,内存占用和原B一致。

2. 实际复制数据的方式

用repeat强制复制B的内容,会占用额外内存:

B_repeated = B.unsqueeze(0).repeat(10, 1, 1)
result = torch.bmm(A.unsqueeze(1), B_repeated).squeeze(1)

这种方式仅在必须物理复制B的场景下使用,不推荐大实例量场景。

三、大实例量(如10万)下的高效处理方案

当实例数量极大时,优先选择分批次处理;若CPU计算需要利用多核心,再考虑多进程方案。

1. 分批次处理(优先推荐)

将A拆分为多个小批次,逐个计算后拼接结果,内存占用低且计算高效:

A = torch.randn(100000, 384)
B = torch.randn(384, 39000)
batch_size = 1000  # 根据内存情况调整批次大小

results = []
for batch_A in torch.split(A, batch_size):
    batch_result = batch_A @ B
    results.append(batch_result)

final_result = torch.cat(results, dim=0)
print(final_result.shape)  # 输出 torch.Size([100000, 39000])

如果用GPU计算,这种方式能充分利用GPU的并行计算能力,效率远高于多进程。

2. CPU多进程处理方案

如果是纯CPU计算,且机器有多个核心,可以用torch.multiprocessing实现并行计算:

import torch.multiprocessing as mp

def process_batch(batch_A, B, result_queue):
    # 计算当前批次的结果并放入队列
    batch_result = batch_A @ B
    result_queue.put(batch_result)

if __name__ == '__main__':
    A = torch.randn(100000, 384)
    B = torch.randn(384, 39000)
    batch_size = 10000
    num_processes = mp.cpu_count()  # 使用全部CPU核心

    # 拆分A为多个批次
    batches = list(torch.split(A, batch_size))
    # 创建进程队列和进程实例
    result_queues = [mp.Queue() for _ in range(num_processes)]
    processes = []

    for i in range(num_processes):
        if i < len(batches):
            p = mp.Process(target=process_batch, args=(batches[i], B, result_queues[i]))
            p.start()
            processes.append(p)

    # 收集所有批次结果
    final_result = []
    for q in result_queues:
        if not q.empty():
            final_result.append(q.get())
    final_result = torch.cat(final_result, dim=0)

    # 等待所有进程结束
    for p in processes:
        p.join()

注意:多进程存在进程间通信开销,需合理设置批次大小(不宜过小);GPU场景不推荐多进程,建议用分批次+GPU加速。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 03:52:42