PyTorch模型训练代码:优化图像重建损失并添加前90%差异监督
import torch
def train(train_loader, model, criterion, optimizer, epoch, args, tb_writer, fp16_args):
total_losses, ssim_losses, mse_losses, batch_time, data_time = [AverageMeter() for _ in range(5)]
norm_ssim_losses = AverageMeter()
max_lr = max([param['lr'] for param in optimizer.param_groups])
print('epoch %d, processed %d samples, lr %.10f' % (epoch, epoch * len(train_loader.dataset), max_lr))
tb_writer.add_scalar('lr', max_lr, epoch)
model.train()
if args['norm_eval'] and args['model_type'].lower() != 'hrnet':
if args['norm_eval_encoder']:
model.encoder = freeze_bn(model.encoder)
else:
model = freeze_bn(model)
end = time.time()
for i_batch, (fname, img, fidt_map, kpoint) in enumerate(tqdm(train_loader)):
data_time.update(time.time() - end)
if args['fp16']:
with torch.autocast(device_type=fp16_args['device_type'], dtype=fp16_args['dtype'], enabled=True):
d6 = model(img.half().cuda())
mse_loss, ssim_loss = criterion(d6, fidt_map.half().cuda(), kpoint)
else:
if int(args['gpu_id']) >= 0:
img = img.cuda()
fidt_map = fidt_map.cuda()
if hasattr(torch.backends, 'mps') and torch.backends.mps.is_available():
img = img.to('mps')
fidt_map = fidt_map.to('mps')
d6 = model(img)
mse_loss, ssim_loss = criterion(d6, fidt_map, kpoint)
# 计算预测图d6与目标图fidt_map之间的差异
diff_map = torch.abs(d6 - fidt_map)
# 将差异向量展平为一维,并排序
diff_vector = diff_map.view(-1)
sorted_diff, _ = torch.sort(diff_vector)
# 计算差异分位点
num_pixels = len(diff_vector)
threshold_index = int(num_pixels * 0.9)
threshold = sorted_diff[threshold_index]
# 保留差异小于阈值的部分作为模型监督
supervised_diff = diff_vector[diff_vector <= threshold]
# 计算监督损失
supervised_loss = torch.mean(supervised_diff)
if args['loss_type'] == 'MSE':
loss = supervised_loss + mse_loss
else:
loss = supervised_loss + mse_loss + ssim_loss
if args['fp16']:
fp16_args['scaler'].scale(loss).backward()
fp16_args['scaler'].step(optimizer)
fp16_args['scaler'].update()
else:
loss.backward()
optimizer.step()
optimizer.zero_grad()
total_losses.update(loss.item())
mse_losses.update(mse_loss.item())
ssim_losses.update(ssim_loss.item())
norm_ssim_losses.update(ssim_loss.item() / criterion.ssim_coefficient)
batch_time.update(time.time() - end)
end = time.time()
print('Epoch: [{0}] '
'Time {batch_time.val:.3f} ({batch_time.avg:.3f}) '
'Data {data_time.val:.3f} ({data_time.avg:.3f}) '
'Loss {loss.val:.4f} ({loss.avg:.4f}) '.format(epoch, batch_time=batch_time,
data_time=data_time, loss=total_losses))
tb_writer.add_scalar('train/total_loss', total_losses.avg, epoch)
tb_writer.add_scalar('train/mse_loss', mse_losses.avg, epoch)
tb_writer.add_scalar('train/ssim_loss', ssim_losses.avg, epoch)
tb_writer.add_scalar('train/norm_ssim_loss', norm_ssim_losses.avg, epoch)
该代码通过以下步骤对训练过程进行了优化:
- 计算预测图
d6与目标图fidt_map之间的像素差异,使用torch.abs()函数计算差值的绝对值。 - 将差异图像展平为一维差异向量,并使用
torch.sort()函数对差异向量进行排序。 - 计算差异向量的分位点,这里选择了前90%的分位点。
- 保留差异小于阈值的部分作为模型的监督部分,使用布尔索引来选择差异小于等于阈值的元素。
- 计算监督损失,这里使用了差异向量的均值作为监督损失。
- 最终的总损失是监督损失和原始的MSE损失或SSIM损失的组合。
通过这个修改,模型只关注那些预测比较准确的部分,从而提高训练效率和最终的模型性能。
注意:
- 请确保
mse_loss和ssim_loss以及相应的损失系数(如果有)在你的代码中定义和计算。 - 你可能需要根据自己的具体情况对代码进行修改,例如调整差异阈值或添加其他损失项。
希望这个回答对你有所帮助。如果你还有其他问题,请随时提出。
原文地址: https://www.cveoy.top/t/topic/SgL 著作权归作者所有。请勿转载和采集!