MCPcopy Create free account
hub / github.com/computational-imaging/bacon / write_image_summary

Function write_image_summary

utils.py:158–195  ·  view source on GitHub ↗
(image_resolution, train_dataset, model, model_input, gt,
                        model_output, writer, total_steps, prefix='train_',
                        val_dataset=None)

Source from the content-addressed store, hash-verified

156
157
158def 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
198def write_psnr(pred_img, gt_img, writer, iter, prefix):

Callers

nothing calls this directly

Calls 1

write_psnrFunction · 0.85

Tested by

no test coverage detected