最近在跑一个简单的CNN做图像分类,输入是64x64的RGB图,batch size设了32。网络结构就是conv+pool+fc,最后一层fc我算过应该输出128维,然后接一个Linear(128, 10)。但一跑训练就报错:size mismatch for fc2.weight: copying a param with shape torch.Size([64, 128]) from checkpoint, but the checkpoint has torch.Size([32, 128])。
我明明记得模型定义里是128输入到10输出,怎么加载权重时它说期望64?检查了state_dict的key和形状,发现是保存模型时把fc2写错了,但代码里没改过啊……是不是我哪里对nn.Sequential的索引理解有误?求指点,网上搜了一圈全是英文,看得头疼。
用PyTorch训练时报错shape不匹配,但自查逻辑没问题,大佬们帮看看?
全部回复
共 37 条这个报错信息其实已经说得很清楚了,fc2.weight在checkpoint里是[32,128],但你的模型定义里是[64,128],说明你保存的模型和现在加载的模型结构不一致。我猜你可能之前跑过别的batch size或者改过网络中间层的输出维度,然后直接load_state_dict了,没注意strict参数默认是True。建议打印一下模型的state_dict和checkpoint里的key对比看看,八成是中间某个fc层的输入输出弄混了,或者你加载的是旧版本的权重文件。
看到这个报错信息,我第一反应是你保存和加载模型时可能没统一用同一个模型类,或者中间改过网络结构。你检查下checkpoint里存的是不是之前某次实验的版本,那个fc2可能对应的是32维的中间层,而不是现在的128。另外,如果用了多GPU训练,保存的state_dict里会有module.前缀,加载时也要对应处理,不然也容易出这种诡异的不匹配。
看到这个报错信息我第一反应是检查下模型里是不是有两个叫fc2的层,有时候命名重复了PyTorch加载时会按名字匹配,容易串。你这里报错说期望64×128,但checkpoint里是32×128,这俩shape差的有点多,不像单纯笔误,我怀疑你保存checkpoint的时候用的模型结构和现在加载的结构不完全一致,比如之前训练时fc层输入维度是32,后来改代码改成128了,但加载时忘了同步。
另外你说自查逻辑没问题,但建议直接打印一下模型每一层的shape,用torchsummary或者自己遍历named_parameters看看,特别是fc2那层的weight shape,别光靠脑子算,有时候卷积输出展平后的维度跟你预期差很多,尤其是有自适应池化或padding调整的时候。
我之前也踩过类似的坑,就是改了几层网络后忘了重新初始化权重,直接load旧checkpoint,结果shape对不上。你有试过删掉strict=False加载吗?虽然不推荐,但至少能确认是不是只有fc2不匹配,如果其他层都正常,那问题就锁定在这一层了。
还有个小细节,你检查的代码里是不是漏了什么?比如在forward里对fc2的输入做了额外处理,或者用了nn.Flatten但位置不对。建议写个简单脚本,构造一个随机输入,跑一遍前向,看实际输出维度是多少,跟你的预期对比下,这样能快速定位是不是网络定义本身的问题。
如果实在找不到,可以把模型定义和保存checkpoint的代码贴出来一起看,光看报错信息确实有点难判断,但大概率是结构不一致或者加载逻辑写岔了。
大概率是checkpoint里存的是旧模型结构,你改过fc层但加载时没对应上,建议打印下state_dict的key看看。
这报错看着像是之前保存的模型权重和你现在定义的模型结构对不上,检查下是不是加载了旧checkpoint。
这报错明显是checkpoint里的fc2存错了,跟当前模型定义对不上,重新保存一下权重文件就行。
看到这个报错第一反应就是checkpoint和模型定义对不上,你确定加载的是同一个训练阶段保存的权重吗?我上次也这样,后来发现是保存时模型还在DataParallel包装下,key里多了module前缀,直接load就会错位。你检查下保存权重时是不是用了多卡训练,或者中途改过batch size,这个[32,128]的形状很像最后一层fc之前的特征图被flatten成了batch size的维度。建议打印一下model.state_dict()里fc2.weight的shape,和checkpoint里的对比下,多半是网络结构里哪层维度算漏了,尤其注意pool层之后要重新算一下flatten后的尺寸。
这个报错信息里已经写得很清楚了,是checkpoint里的fc2权重是32x128,但你当前模型期望的是64x128。大概率不是你模型定义的问题,而是你加载的checkpoint本身就不是这个模型存的,比如之前用别的batch size或者不同网络结构保存的。建议你打印一下checkpoint里所有key的shape,跟当前模型逐一比对,别只盯fc2。另外如果是自己刚存的权重,检查下保存时是不是用了model.state_dict(),而不是别的变量。
这个报错其实是模型定义和checkpoint里的状态字典对不上,不是你的逻辑问题。你定义的fc2是Linear(128,10),但保存的权重却是32x128的,说明之前保存的模型里fc2层的输入维度是32,大概率你改过网络结构或者batch size影响了那个层的定义。建议你加载权重时用strict=False,然后把模型打印出来对比一下state_dict的key和shape,重点看fc2那层前后的维度变化,别光看最后输出的10。另外你检查的那段代码是不是漏了什么,比如在某个地方把fc2重新赋值了?我上次也遇到类似问题,最后发现是模型里有个变量名被意外覆盖了。
八成是checkpoint里存的是旧模型或者中间变量,你load_state_dict时用strict=False试试,或者打印下模型和ckpt的key对比下。
看下checkpoint是不是之前用不同batch size或旧模型存的,加载时别用strict=True,或者直接打印下模型state_dict的key对比下。
大概率是checkpoint里存的是旧模型结构,你改代码后没重新保存权重,直接load旧档就串shape了。重新训两步再存个新权重试试。
看你这报错信息,我第一反应是checkpoint里存的根本不是你这个模型的权重,可能是之前跑别的实验时保存的。你那个“检查了”后面是不是漏了代码?建议直接打印一下model.fc2.weight.shape和ckpt里对应key的shape,如果确实不一致,大概率是加载路径写错了或者模型定义里fc2的输入维度被改过但没重新初始化。
另外,如果用的是多GPU训练,state_dict里的权重会带module.前缀,直接load就容易错位。可以先试试strict=False加载,看哪些key对不上,再针对性处理。
我遇到过类似情况,最后发现是保存模型时用了旧版本的代码,网络结构改了但没重新训练就直接load了旧权重。你最好确认下这个checkpoint是不是当前模型结构下保存的。
这个报错其实挺典型的,问题大概率不在你最后那个fc层,而在checkpoint本身。你说了模型定义是Linear(128, 10),但报错里显示checkpoint里fc2.weight的shape是[32, 128],说明你保存权重的时候,那个模型最后一层根本不是128->10,而是32->128。你可能是在保存checkpoint之前,把某个中间层的输出维度改成了32,或者用了不同版本的网络结构去训练,然后加载时用了新定义的模型。
我遇到过类似的情况,有时候是训练时用了DataParallel,保存的state_dict里带了module.前缀,但加载时没处理,导致层名对不上,虽然你这里报的是shape不匹配而不是key不匹配,但本质都是结构不一致。
建议你先打印一下checkpoint里所有层的shape,跟当前模型逐层对比,看看是只有fc2对不上,还是前面几层也有偏差。另外,如果这个模型是中途保存的,确认一下保存时的模型定义和现在是不是完全一样,哪怕改动了一个卷积层的padding,都可能让展平后的维度变化。
还有个比较隐蔽的点,如果你用的是预训练权重,但原始模型输入是32x32的图(比如CIFAR),你改成64x64后,前面卷积层的输出维度会变,但fc1的输入维度是固定的,所以如果你没改fc1,那报错可能就不在fc2,而是在fc1——你现在的报错是fc2,说明fc1倒是匹配的,那就要怀疑是不是保存时用了不同的网络。
其实最直接的办法,别加载checkpoint,先随机初始化跑一个batch看看能不能前向,如果能,那就纯粹是加载逻辑的问题,跟网络无关。你检查一下load_state_dict时有没有加strict=False,有时候这样能跳过不匹配的层,但会掩盖真正的问题。
大概率是checkpoint里存的是旧模型结构,你改过fc层但没重跑训练,直接加载老权重当然对不上。重新保存一下模型权重试试。
这个报错信息其实已经说得很明白了,checkpoint里存的是fc2的shape是[32,128],但你现在模型里期望的是[64,128],说明你加载的权重文件根本不是当前这个模型结构训练出来的。最可能的原因是你之前跑过别的实验,把那个模型的checkpoint覆盖了,或者你在定义模型时fc2的输入维度实际写的是64,但你自己记成128了。建议直接打印一下model.fc2.weight.shape,再对比一下checkpoint里对应的shape,基本就能定位问题。另外如果确认是旧权重,直接删掉重新训练就行,没必要纠结。
看到这个报错我感觉八成是checkpoint里存的不是模型权重,而是优化器状态或者之前某个中间变量的快照,shape正好是[32,128]就很可疑。你检查下保存的时候是不是用了model.state_dict(),如果是在某个forward hook里存的feature map那就对不上了。另外也可以打印一下checkpoint的keys,看看里面到底有哪些tensor,跟当前模型的state_dict对比下,比对着找问题快很多。我之前也遇到过类似情况,最后发现是保存时忘了加module前缀,但你这个shape差太多,更像是存错对象了。