(image_resolution, train_dataset, model, model_input, gt,
model_output, writer, total_steps, prefix='train_',
val_dataset=None)
| 156 | |
| 157 | |
| 158 | def write_image_summary(image_resolution, train_dataset, model, model_input, gt, |
| 159 | model_output, writer, total_steps, prefix='train_', |
| 160 | val_dataset=None): |
| 161 | |
| 162 | gt_img = dataio.lin2img(gt['img'], image_resolution) |
| 163 | pred_img = dataio.lin2img(model_output['model_out']['output'], image_resolution) |
| 164 | |
| 165 | output_vs_gt = torch.cat((gt_img, pred_img), dim=-1) |
| 166 | writer.add_image(prefix + 'gt_vs_pred', make_grid(output_vs_gt, scale_each=False, normalize=True), |
| 167 | global_step=total_steps) |
| 168 | |
| 169 | write_psnr(pred_img, gt_img, writer, total_steps, prefix+'img_') |
| 170 | |
| 171 | # validation samples |
| 172 | if val_dataset is None: |
| 173 | return |
| 174 | |
| 175 | image_resolution = [2*r for r in image_resolution] |
| 176 | model_input, gt = val_dataset[0] |
| 177 | tmp = {} |
| 178 | for key, value in model_input.items(): |
| 179 | if isinstance(value, torch.Tensor): |
| 180 | tmp.update({key: value[None, ...].cuda()}) |
| 181 | else: |
| 182 | tmp.update({key: value}) |
| 183 | model_input = tmp |
| 184 | gt = {key: value[None, ...].cuda() for key, value in gt.items()} |
| 185 | |
| 186 | with torch.no_grad(): |
| 187 | model_output = model(model_input) |
| 188 | |
| 189 | gt_img = dataio.lin2img(gt['img'], image_resolution) |
| 190 | pred_img = dataio.lin2img(model_output['model_out']['output'], image_resolution) |
| 191 | |
| 192 | output_vs_gt = torch.cat((gt_img, pred_img), dim=-1) |
| 193 | writer.add_image(prefix + 'val_gt_vs_pred', make_grid(output_vs_gt, scale_each=False, normalize=True), |
| 194 | global_step=total_steps) |
| 195 | write_psnr(pred_img, gt_img, writer, total_steps, 'val_img_') |
| 196 | |
| 197 | |
| 198 | def write_psnr(pred_img, gt_img, writer, iter, prefix): |
nothing calls this directly
no test coverage detected