自定义损失函数训练中loss.backward()耗时远超前向传播的问题求助
自定义损失函数训练中loss.backward()耗时远超前向传播的问题求助
我最近在训练模型时遇到了个棘手的问题:用自定义损失函数训练时,loss.backward()这一步居然要花2秒,但如果去掉反向传播步骤,整个训练迭代过程只需要0.2秒,差距大得离谱。我已经尝试过优化计算图,但几乎没什么效果,而且所有张量都已经移到CUDA上运行了,实在搞不懂问题出在哪。
涉及的网络结构
class half2(nn.Module): def __init__(self): super().__init__() self.conv1 = nn.Conv2d(3,20,(11,11)) self.relu = nn.ReLU() self.maxpool = nn.MaxPool2d((2,2)) self.conv2 = nn.Conv2d(20,40,(5,5)) self.conv3 = nn.Conv2d(40,80,(3,3)) self.fc1 = nn.Linear(80,160) self.fc2 = nn.Linear(160,80) def forward(self,img): res = self.relu(self.conv1(img)) #res = self.maxpool(res) res = self.relu(self.conv2(res)) res = self.relu(self.conv3(res)) res = res.view((res.shape[1],res.shape[2],res.shape[0])) res = self.relu(self.fc1(res)) res = torch.relu(self.fc2(res)) return torch.flatten(res) def train_forward(self,img): res = self.relu(self.conv1(img)) #res = self.maxpool(res) res = self.relu(self.conv2(res)) res = self.relu(self.conv3(res)) res = res.view((res.shape[1],res.shape[2],res.shape[0])) res = self.relu(self.fc1(res)) res = torch.relu(self.fc2(res)) return torch.flatten(res)
我的自定义损失函数
class ContrastiveLoss(torch.nn.Module): def __init__(self, margin=1.0): super(ContrastiveLoss, self).__init__() self.margin = margin def forward(self, output1,output2, Y): D = tfunc.pairwise_distance(output1, output2) left = Y*0.5 left = left * torch.pow(D,2) right = (1-Y)*0.5 choose = torch.clamp(self.margin-D, min=0.0) right = right * torch.pow(choose,2) return left + right
训练循环代码
running_loss = 0 cnt = 0 same = torch.tensor(0,dtype=torch.float32).to(device) diff = torch.tensor(1,dtype=torch.float32).to(device) for epoch in range(5): for dr1 in range(1,3): for ig1 in range(1,9): for dr2 in range(1,3): for ig2 in range(1,9): if dr1 == dr2 and ig1 == ig2: continue optimizer.zero_grad() if dr1 == dr2: target_dist = same else: target_dist = diff img1 = image_to_input(f"./face_dataset/s{dr1}/{ig1}.pgm").to(device) img2 = image_to_input(f"./face_dataset/s{dr2}/{ig2}.pgm").to(device) output1 = net3.train_forward(img1) output2 = net3.train_forward(img2) loss = criterion(output1,output2,target_dist) loss.backward() optimizer.step() running_loss += loss.item() cnt += 1 if cnt != 0: print("epoch",epoch,"set",dr1,"loss",running_loss / cnt) running_loss = 0.0 cnt = 0
备注:内容来源于stack exchange,提问作者yariba3000
相关产品推荐
相关产品推荐

