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)

该代码通过以下步骤对训练过程进行了优化:

  1. 计算预测图d6与目标图fidt_map之间的像素差异,使用torch.abs()函数计算差值的绝对值。
  2. 将差异图像展平为一维差异向量,并使用torch.sort()函数对差异向量进行排序。
  3. 计算差异向量的分位点,这里选择了前90%的分位点。
  4. 保留差异小于阈值的部分作为模型的监督部分,使用布尔索引来选择差异小于等于阈值的元素。
  5. 计算监督损失,这里使用了差异向量的均值作为监督损失。
  6. 最终的总损失是监督损失和原始的MSE损失或SSIM损失的组合。

通过这个修改,模型只关注那些预测比较准确的部分,从而提高训练效率和最终的模型性能。

注意

  • 请确保mse_lossssim_loss以及相应的损失系数(如果有)在你的代码中定义和计算。
  • 你可能需要根据自己的具体情况对代码进行修改,例如调整差异阈值或添加其他损失项。

希望这个回答对你有所帮助。如果你还有其他问题,请随时提出。


原文地址: https://www.cveoy.top/t/topic/SgL 著作权归作者所有。请勿转载和采集!

免费AI点我,无需注册和登录