PyTorch小批量训练中能否关闭训练集的requires_grad以节省内存?
问题解答
1. 能否关闭positive_set_train和negative_set_train的requires_grad?
完全可以,而且这才是正确的操作。这些张量属于训练输入数据,我们训练时只需要更新模型的参数,根本不需要对输入数据计算梯度。开启requires_grad=True会让PyTorch为这些大张量额外分配梯度存储空间,纯粹是内存浪费,关闭后能有效节省内存。
2. 关于data的requires_grad理解是否正确?
你的理解存在偏差:
data是从原始正负样本张量切片拼接而来的,它的requires_grad属性会继承原始张量的设置。如果原始的positive_set_train和negative_set_train关闭了requires_grad,那么data的requires_grad也会是False。- 但这完全不影响训练!模型训练只需要模型的可训练参数(即
net里的权重、偏置等)带有requires_grad=True(模型初始化时默认就是这个状态)。输入数据的requires_grad不需要为True,因为我们不需要对输入数据求导,loss.backward()只会计算模型参数的梯度,和输入数据的梯度状态无关。
额外代码小问题提醒
你的get_bootstrap_batch函数里的shuffle逻辑有个小bug:
if shuffle: p = torch.randperm(positives.size(0)) # 这里只生成了正样本数量的随机索引 data = data[p, :, :] # 但data是正样本+负样本,长度是positives.shape[0]+BOOTSTRAP_SIZE target = target[p]
这样只会打乱前positives.size(0)个元素(正样本),负样本位置完全没变,不符合整体shuffle的需求。应该改成对整个data的长度生成随机索引:
if shuffle: p = torch.randperm(data.size(0)) data = data[p, :, :] target = target[p]
内容的提问来源于stack exchange,提问作者Gábor Erdős
相关产品推荐
相关产品推荐

