ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

PyTorch Concat算子拼接Tensor报错全解析:从维度到显存

PyTorch Concat算子拼接Tensor报错全解析:从维度到显存 做深度学习的人谁没被Concat算子坑过几次呢明明看起来就是简简单单的“把张量拼在一起”可一旦跑起来报错信息五花八门从Sizes of tensors must match到Expected all tensors to be on the same device再到CUDA out of memory每一行都让人头皮发麻。这篇就想把Concat算子拼接多个Tensor的报错场景一次讲透包括每条报错背后的原理、最可能的触发原因、快速定位的手段还有我这些年实际排查中踩过的坑。无论你是刚入门的新手还是已经被模型训练折磨到麻木的老手这份内容都能帮你少走弯路。1. 为什么Concat算子拼接Tensor会频繁报错先明确一个事实Concat在深度学习代码里几乎是“出场率最高”的算子之一。数据预处理要把特征拼起来模型中间层要把多尺度特征拼起来后处理要拼预测结果强化学习里拼观测和历史状态……简直无处不在。但越常用的算子越容易因为“使用姿势不对”而炸。很多报错并不是Concat算子本身有bug而是框架对张量的要求极其严格任何一个前置环节留下了“尾巴”都会在拼接这一刻集中爆发。我把这类问题归纳成三大类第一类是形状系统问题。Tensor的shape不匹配或者拼接维度之外的维度尺寸对不上这是最常见的报错来源占掉了我遇到过的七成以上情况。第二类是属性一致性问题。包括dtype数据类型不一致、device设备不一致、layout内存布局不一致这些属性在单个张量内部看不出来但拼接时框架要求所有输入必须保持一致否则直接抛异常。第三类是资源与状态问题。比如显存不足、梯度图中断、张量非连续non-contiguous等。这类报错最迷惑人因为表面上看是Concat的问题实际是上游变量状态不对或显存分配失败。我见过很多同事、网友被这些报错折磨到怀疑人生甚至有人直接放弃尝试、改成for循环逐行append结果速度慢了几十倍。其实只要弄清楚框架在Concat时到底做了什么检查、每个报错对应什么原因排查起来就是“按图索骥”的事。2. Concat算子的核心机制与报错根因2.1 框架在Concat时到底做了什么以PyTorch为例torch.cat和torch.concattorch.concat是torch.cat的别名在底层执行时大致会经历这么几个阶段Shape推导框架根据所有输入张量的shape计算拼接后输出张量的shape。属性校验逐个检查输入张量的dtype、device、layout是否一致同时确认拼接维度是否在有效范围。内存分配依据推导出的输出shape和dtype申请一块连续内存或复用已有的内存块。Kernel执行启动一个数据搬运kernel将每个输入张量按偏移量拷贝到输出张量对应位置。这个过程中每一阶段都可能引发报错。比如shape推导阶段发现维度不匹配会立即抛出RuntimeError: Sizes of tensors must match...属性校验阶段发现dtype不统一会抛出RuntimeError: expected dtype ... but found dtype ...内存分配阶段失败则会抛出CUDA out of memory。理解这个流程的价值在于报错信息里的每一个关键词都有明确的指向性比如看到expect、found、dim、device、memory基本就知道是哪一环节出了问题。2.2 Tensor core与Concat的关系这里想顺便澄清一个经常被混淆的点很多人听到“GPU算子”就联想到Tensor Core但Concat本质上是数据搬运操作不涉及矩阵乘加运算所以它一般不跑在Tensor Core上而是以普通的elementwise copy kernel或专用融合kernel形式执行。那为什么网上会有“ConcatTensor Core”的讨论因为拼接后往往紧跟着矩阵乘法或卷积而Tensor Core对输入tensor的shape和内存对齐有要求比如ldm对齐、16×16分块等。有些优化库会把Concat Conv、Concat MatMul融合成一个算子目的就是减少中间张量的显存占用和拷贝开销。这种融合思路在一些推理引擎比如TensorRT、自研推理框架里非常常见。也就是说如果你在做算子开发、算子融合这时候才需要考虑Tensor Core的对齐规则。但如果你只是业务代码里调用torch.cat那大概率是前面三类常规问题先把基本功打牢再说。2.3 Concat、cat、stack和append的区别再补充一个基础但是高频混淆的知识点torch.cat/torch.concat沿已有维度拼接不增加新维度。torch.stack沿新维度堆叠所有输入张量shape必须完全一致。list.appendPython列表操作不会做张量校验但后续需要手动转换类型。很多人报错是因为搞混了cat和stack的使用场景或者用append拼了一堆不同shape的Tensor最后想一次性torch.cat结果维度全乱套。所以在正式开讲报错之前先记住一句话拼接前先print每个张量的shape确认拼接维度和非拼接维度的尺寸。3. 高频报错场景的完整复现与解决3.1 最常见的维度不匹配报错报错信息长这样RuntimeError: Sizes of tensors must match except in dimension 1. Expected size 3 but got size 2 for tensor number 1 in the list.产生这个报错的代码可能是import torch a torch.randn(3, 4) b torch.randn(2, 4) # 在第0维拼接要求第1维尺寸相同都是4合法 c torch.cat([a, b], dim0) print(c.shape) # torch.Size([5, 4]) # 在第1维拼接要求第0维尺寸相同a是3b是2报错 c torch.cat([a, b], dim1)你看代码本身没有写错错在选择的拼接维度不对。这里的规律很清晰torch.cat只要求“拼接维度对应的尺寸可以不同但其他所有维度的尺寸必须一致”。假设拼接维度是dim0那所有输入张量的第1维到最后一维必须完全相同如果拼接维度是dim1那所有输入张量的第0维、第2维及以后必须完全相同。实际排查中我个人的最快定位方法是在torch.cat上面加一行 print把所有输入张量的shape打出来print([t.shape for t in tensor_list])一眼就能看出来是哪个张量的哪个维度对不上。不要相信“大概没问题”的直觉数据在流动过程中改变shape是常有的事尤其经过squeeze、unsqueeze、view、permute之后维度顺序和含义很容易被偷换。3.2 dtype不一致的报错与陷阱报错信息通常长这样RuntimeError: expected dtype Float but found dtype Double触发代码很典型import torch a torch.randn(3, 4) # float32 b torch.randn(3, 4).double() # float64 c torch.cat([a, b], dim0)这种dtype不一致的报错在数据预处理流水线里极其常见。比如你用一个函数读取数据默认返回float64另一个分支的数据来自numpy默认是float64但模型权重和大部分张量都是float32。拼接时框架严格按“首次出现的dtype”来要求后续张量任何一个不匹配直接报错。还有一种隐蔽情况整数张量和浮点张量拼接。比如你有一个torch.tensor([1, 2, 3])默认int64和一个torch.tensor([1.0, 2.0, 3.0])默认float32拼接时会报错。这类问题在模型后处理阶段经常出现分类标签是long型预测概率是float型想拼成一个表格时就会中招。我的建议是统一约定张量的dtype。数据预处理最后输出为float32标签索引输出为long概率分数保持float32。如果实在不放心拼接前做一次显式转换b b.to(a.dtype)3.3 device不一致报错经典报错长这样RuntimeError: Expected all tensors to be on the same device, but found at least two devices, cuda:0 and cpu!触发代码import torch a torch.randn(3, 4).cuda() b torch.randn(3, 4) # 默认在CPU上 c torch.cat([a, b], dim0)这种报错多发生在混合了CPU数据和GPU数据的代码路径里。比如模型前向计算在GPU上但某个特征的预处理逻辑在CPU上完成最后拼接时忘了把CPU张量搬到GPU。还有一种情况在多卡环境下更隐蔽一张卡上的cuda:0和另一张卡上的cuda:1拼接报错内容变成Expected all tensors to be on the same device, but found at least two devices, cuda:0 and cuda:1!。这时候你需要确认为什么两张卡的数据会跑到同一个拼接点多半是分布式采样或者数据加载器的逻辑问题。解决方案没什么可炫技的全部统一到同一个设备上device a.device c torch.cat([a, b.to(device)], dim0)3.4 非连续张量(Non-contiguous)引发的诡异问题非连续张量导致的报错比较“诡异”因为它不总是报错。有时候正常拼接有时候就报RuntimeError: tried to construct a tensor from a lazily-initialized tensor...或者干脆在反向传播时报错。触发场景通常是import torch a torch.randn(3, 4).t() # 转置得到非连续张量shape是[4, 3] b torch.randn(4, 3) c torch.cat([a, b], dim1)torch.cat内部通常会处理非连续输入但处理过程涉及一次额外的contiguous拷贝这在显存紧张时可能造成峰值显存飙升。更麻烦的是如果你的模型用了自定义的autograd.Function或者经过torch.jit.trace之类工具非连续张量会触发一些底层断言报出极难排查的异常。我的经验是涉及拼接操作的张量尽量在上游就保证连续。如果实在无法避免转置、切片等操作就在拼接前调用.contiguous()。虽然多一次拷贝但至少稳定可控。3.5 CUDA out of memory拼接大张量时的隐形杀手这个报错不总是指向Concat但我有两次排查了很久最后发现确实是因为Concat产生的中间张量把显存撑爆了RuntimeError: CUDA out of memory. Tried to allocate 2.00 GiB (GPU 0; 23.70 GiB total capacity; 22.38 GiB already allocated; ...)为什么会这样因为torch.cat的输出张量是一个全新的连续张量所有输入数据都会被拷贝一份到新内存里。如果你拼接的输入本身已经占了不少显存输出张量又会占据同样量级的显存就产生了“双倍占用”的瞬时压力。一个典型的例子在目标检测的FPN层经常要把不同尺度的特征图分别上采样后拼接。每个特征图动辄几百MB拼接后的结果又是几百MB瞬间峰值显存就上去了。应对策略有几个尽量复用预分配的内存缓冲区避免频繁拼接大张量。用out参数指定输出张量减少多次拼接的中间分配。在显存充裕时做拼接在显存紧张时改用“直接写入目标位置”的算子或者分块处理。4. 从算子执行全流程理解Concat报错4.1 GPU上Concat算子的完整执行流程你在代码里写一个torch.cat([a, b], dim0)框架做了哪些事以GPU场景为例完整链路大致如下Python层的torch.cat调用进入ATen库经过dispatch机制分派到CUDA实现。框架计算输出shape校验输入属性shape规则、dtype、device。框架向缓存分配器caching allocator申请显存空间用于存放输出张量。框架把一个名为cat的CUDA kernel放到当前stream上。这个kernel会根据拼接维度和输入张量的内存偏移把输入数据逐块拷贝到输出空间。如果开启了自动混合精度或使用了融合执行计划框架可能选择融合kernel减少中途显存使用。如果中间任一步抛错你的训练进程就收到一个RuntimeError。而很多底层错误比如非法内存访问、kernel启动失败会被包装成CUDA error: an illegal memory access was encountered之类的信息这时候往往需要回溯到前面的Concat或者更上游的view、permute操作。4.2 为什么有时报错信息会“漂移”这是实践中最坑的一点报错提示的位置不一定准确但报错时间点一定是真实的。比如你在写数据加载代码用一个列表从不同渠道收集句子向量最终在训练前一步拼接。但实际上某个句子的向量在生成时就已经变成float64了问题源头在数据处理逻辑可报错会出现在训练循环里的Concat那一行。我排查过的很多case里真正的原因都在拼接之前的三四层调用栈之外。所以遇到Concat报错我的建议是不只盯着报错那一行要把报错点之前的张量属性全部检查一遍。用pdb或者打印大法把关键张量的shape、dtype、device、is_contiguous四个属性都打出来。这个动作看起来耗时但比瞎改代码高效得多。4.3 算子开发场景中的Concat陷阱如果你在做算子开发或者自定义算子融合比如AscendC算子、NPU算子或者CUDA自定义kernelConcat相关的坑又是另一层复杂度。首先是shape推导必须严格匹配。手写算子时拼接维度的偏移计算涉及多个张量的stride和存储偏移容易因维度顺序搞错导致数据错位这在debug时非常难查因为不报错只是结果错误。其次是内存对齐问题。GPU端的kernel常常需要按128字节或256字节对齐。拼接后输出张量的每一行起始地址如果不对齐就会触发未定义行为或性能骤降。很多自研算子库里的实现会在拼接前对输入做padding保证每段数据的起始偏移满足对齐要求。最后是融合算子中的维度推断。如果要把Concat Elementwise或Concat MatMul融合必须重新推导中间shape的变化。融合的目的是省掉中间张量的读写但前提是拼接后的数据在后续算子中的访问模式是连续的、可预测的。否则融合方案宁可放弃因为省下的显存远不如踩坑成本高。5. 常见报错速查表与排查技巧这里整理一张速查表方便你在碰到报错时直接对号入座。报错信息片段根本原因快速定位方法常规解法Sizes of tensors must match except in dimension X非拼接维度的尺寸不一致打印每组输入shape逐维对照调整数据shape或更换拼接维度expected dtype ... but found dtype ...输入张量dtype不一致打印每个张量的dtype统一dtype用.to(dtype)转换Expected all tensors to be on the same device输入张量device不一致打印每个张量的device统一搬到同一设备注意多卡场景CUDA out of memory显存峰值过高查看报错行附近是否有大张量拼接用out预分配、分块处理、减小batchtried to construct a tensor from a lazily-initialized tensor非连续张量或版本Bug检查前序permute、transpose、view拼接前.contiguous()dimension out of range拼接维度超出张量维度数检查tensor.ndim和dim值确认dim不超过tensor维度数间接报错发生在别的算子上游属性污染检查Concat之前的张量属性链在源头统一规格针对其中几个高频场景还有几条“独家”排查技巧如果你是新手就在每个张量生成之后立即检查shape不要等拼到一起再看。我常写一个超简单工具函数def show(*tensors, namesNone): for i, t in enumerate(tensors): name names[i] if names else ftensor_{i} print(f{name}: shape{tuple(t.shape)}, dtype{t.dtype}, fdevice{t.device}, contiguous{t.is_contiguous()}) show(a, b, names[a, b])如果报错信息里带有except in dimension先看数字是多少那个数字就是你的拼接维度。很多时候你实际想拼的维度是1但因为某个张量少了一维框架就认为你想拼0维于是报错维度也跟着变了。对于显存类报错别急着找Concat的麻烦。先用torch.cuda.memory_summary()看看内存分配的大头在哪。如果拼接后立即需要另一个大张量而且拼接结果本身不再修改可以考虑直接对目标张量切片赋值绕过torch.cat# 预分配输出张量 out torch.empty(a.shape[0] b.shape[0], a.shape[1], devicea.device) out[:a.shape[0]] a out[a.shape[0]:] b这种方式没有额外的临时分配峰值显存更低缺点是代码不够优雅但为了训练不崩值得。6. 我在实际处理Concat报错时的一些经验最后分享几条我自己的实际操作心得不一定成体系但都是真金白银踩出来的。第一不要迷信“用cat就是对的”。concat确实是深度学习里的基础操作但不同语义场景应该用不同的拼接方式。如果只是想把多个模型输出沿batch维度堆叠那torch.stack或者直接构造新张量可能更合适如果想把两个特征向量合并成一个更长的向量torch.cat是对的如果想把各类标签和预测结果拼成一个结构化表格用torch.stack再加permute反而更可控。认清语义比记住API更重要。第二防御性编程是救命稻草。我写数据流比较复杂的代码时习惯在关键拼接点前放断言assert all(t.shape a.shape for t in tensor_list[1:]), \ fshape mismatch: {[t.shape for t in tensor_list]} assert all(t.dtype tensor_list[0].dtype for t in tensor_list), \ fdtype mismatch: {[t.dtype for t in tensor_list]} assert all(t.device tensor_list[0].device for t in tensor_list), \ fdevice mismatch: {[t.device for t in tensor_list]}看上去像是“多余的代码”但真的能帮你把问题定位在源头而不是等到模型前向跑到一半才崩溃。第三多卡和混合精度环境下要格外小心。自动混合精度训练时PyTorch会自动为某些操作插入dtype转换但Concat不一定会自动做转换。多卡环境下不同DPU上生成的张量设备不同要拼接前必须统一device。我在分布式训练里还遇到过一个问题不同rank上输入数据的顺序不一致拼接后模型输出虽然能算但loss异常排查半天不是Concat的锅而是数据采样器的问题。所以报错只是线索不是答案。第四报错之后先确定范围再动手。我见过很多人一看到Concat报错就把代码改成stack、改成append、加上无数个.contiguous()结果越改越乱。正确姿势是先打印所有输入张量的shape、dtype、device、is_contiguous把四个属性确认完再决定动哪里。这四个属性都一致的情况下Concat几乎不会出问题如果不一致缺哪个补哪个别乱枪打鸟。第五大张量拼接前养成“预分配片段赋值”的习惯。训练和推理阶段还好一旦到了调优、部署环节显存就是硬约束。很多推理引擎不支持动态shape变化一个随意的Concat都可能让整个优化计划失效。如果拼接结果的大小在运行前就能确定尽量用预分配的方式既省显存又减少底层内存碎片。如果你看完这篇至少能分清“维度不匹配”和“dtype不一致”的区别遇到报错能第一时间打印属性而不是干瞪眼那这篇就值了。至于更复杂的自定义kernel、融合算子里的Concat实现问题那属于另一层面的技术深水区有机会我再单独写一篇展开聊。
返回列表