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
相关产品推荐
相关产品推荐

