特征提取时是否必须使用torch.no_grad()?当前实现合理性问询
你的实现确实存在效率层面的问题,但不会直接破坏classifier的训练结果,具体分析如下:
eval模式不自动关闭梯度计算
DenseNet设置为eval()只是关闭了batch norm的运行均值更新、dropout的随机丢弃逻辑,但不会自动停止梯度的计算与存储。所以你当前的代码中,提取features_1和features_2时,PyTorch依然会为DenseNet的所有参数计算梯度并保留在内存中。对classifier训练结果无直接影响,但浪费资源
因为你的optimizer只绑定了classifier.parameters(),所以loss.backward()时,DenseNet的梯度不会被更新(optimizer只会更新它管理的参数),因此classifier的训练逻辑本身是正常的。但问题在于,DenseNet参数量极大,存储这些无用的梯度会占用大量显存,拖慢训练速度,甚至可能触发显存不足(OOM)的报错。必须添加torch.no_grad()优化
正确的做法是用torch.no_grad()包裹DenseNet的前向传播过程,彻底关闭梯度计算:with torch.no_grad(): features_1 = densenet(inputs_1) # extract features 1 features_2 = densenet(inputs_2) # extract features 2这样既不会影响classifier的梯度计算与参数更新,又能节省显存、提升训练效率。
代码小错误修正
你的代码里combined = combined(-1, 4416)是语法错误,应该改成combined = combined.view(-1, 4416)或者combined = combined.flatten(1)来完成张量重塑。
内容的提问来源于stack exchange,提问作者Ze0ruso

