we implement random view synthesis for generating more views to help training posenet
(args, targets, rgbs, poses, virtue_view, poses_perturb, feat_model, dset_size, FeatureLoss, optimizer, hwf, img_idxs, render_kwargs_test)
| 164 | return train_loss |
| 165 | |
| 166 | def train_on_batch_with_random_view_synthesis(args, targets, rgbs, poses, virtue_view, poses_perturb, feat_model, dset_size, FeatureLoss, optimizer, hwf, img_idxs, render_kwargs_test): |
| 167 | ''' we implement random view synthesis for generating more views to help training posenet ''' |
| 168 | feat_model.train() |
| 169 | |
| 170 | H, W, focal = hwf |
| 171 | H, W = int(H), int(W) |
| 172 | |
| 173 | if args.freezeBN: |
| 174 | feat_model = freeze_bn_layer_train(feat_model) |
| 175 | |
| 176 | train_loss_epoch = [] |
| 177 | |
| 178 | # random generate batch_size of idx |
| 179 | select_inds = np.random.choice(dset_size, size=[dset_size], replace=False) # (N_rand,) |
| 180 | |
| 181 | batch_size=args.featurenet_batch_size # manual setting, use smaller batch size like featurenet_batch_size = 4 if OOM |
| 182 | if dset_size % batch_size == 0: |
| 183 | N_iters = dset_size//batch_size |
| 184 | else: |
| 185 | N_iters = dset_size//batch_size + 1 |
| 186 | |
| 187 | i_batch = 0 |
| 188 | for i in range(0, N_iters): |
| 189 | if i_batch + batch_size > dset_size: |
| 190 | i_batch = 0 |
| 191 | break |
| 192 | i_inds = select_inds[i_batch:i_batch+batch_size] |
| 193 | i_batch = i_batch + batch_size |
| 194 | |
| 195 | # convert input shape to [B, 3, H, W] |
| 196 | target_in = targets[i_inds].clone().permute(0,3,1,2).to(device) |
| 197 | rgb_in = rgbs[i_inds].clone().permute(0,3,1,2).to(device) |
| 198 | pose = poses[i_inds].clone().reshape(batch_size, 12).to(device) |
| 199 | rgb_perturb = virtue_view[i_inds].clone().permute(0,3,1,2).to(device) |
| 200 | pose_perturb = poses_perturb[i_inds].clone().reshape(batch_size, 12).to(device) |
| 201 | |
| 202 | # inference feature model for GT and nerf image |
| 203 | pose = torch.cat([pose, pose]) # double gt pose tensor |
| 204 | features, predict_pose = feat_model(torch.cat([target_in, rgb_in]), return_feature=True, upsampleH=H, upsampleW=W) # features: (1, [2, B, C, H, W]) |
| 205 | |
| 206 | # get features_target and features_rgb |
| 207 | if args.DFNet: |
| 208 | features_target = features[0] # [3, B, C, H, W] |
| 209 | features_rgb = features[1] |
| 210 | |
| 211 | loss_pose = PoseLoss(args, predict_pose, pose, device) # target |
| 212 | |
| 213 | if args.tripletloss: |
| 214 | loss_f = triplet_loss_hard_negative_mining_plus(features_rgb, features_target, margin=args.triplet_margin) |
| 215 | else: |
| 216 | loss_f = FeatureLoss(features_rgb, features_target) # feature Maybe change to s2d-ce loss |
| 217 | |
| 218 | # inference model for RVS image |
| 219 | _, virtue_pose = feat_model(rgb_perturb.to(device), False) |
| 220 | |
| 221 | # add relative pose loss here. TODO: This FeatureLoss is nn.MSE. Should be fixed later |
| 222 | loss_pose_perturb = PoseLoss(args, virtue_pose, pose_perturb, device) |
| 223 | loss = args.combine_loss_w[0]*loss_pose + args.combine_loss_w[1]*loss_f + args.combine_loss_w[2]*loss_pose_perturb |
no test coverage detected