Perform 1 epoch of training with batch
(args, data_loaders, model, feat_model, hwf, half_res, device, world_setup_dict, **render_kwargs_test)
| 213 | return iter_loss, iter_psnr |
| 214 | |
| 215 | def eval_on_epoch(args, data_loaders, model, feat_model, hwf, half_res, device, world_setup_dict, **render_kwargs_test): |
| 216 | ''' Perform 1 epoch of training with batch ''' |
| 217 | model.eval() |
| 218 | batch_size = 1 |
| 219 | |
| 220 | train_dl, val_dl, test_dl = data_loaders |
| 221 | |
| 222 | total_loss = [] |
| 223 | total_psnr = [] |
| 224 | |
| 225 | #### Core optimization loop ##### |
| 226 | for data, pose, img_idx in val_dl: |
| 227 | # training one step with batch_size = args.batch_size |
| 228 | loss, psnr = eval_on_batch(args, data, model, feat_model, pose, img_idx, hwf, half_res, device, world_setup_dict, **render_kwargs_test) |
| 229 | total_loss.append(loss.item()) |
| 230 | total_psnr.append(psnr.item()) |
| 231 | total_loss_mean = np.mean(total_loss) |
| 232 | total_psnr_mean = np.mean(total_psnr) |
| 233 | return total_loss_mean, total_psnr_mean |
| 234 | |
| 235 | def train_on_feature_batch(args, data, model, feat_model, pose, img_idx, hwf, optimizer, device, world_setup_dict, **render_kwargs_test): |
| 236 | ''' Perform 1 step of training using scheme1 ''' |
no test coverage detected