MCPcopy Create free account
hub / github.com/ActiveVisionLab/DFNet / train_on_batch_with_random_view_synthesis

Function train_on_batch_with_random_view_synthesis

script/run_feature.py:166–230  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

164 return train_loss
165
166def 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

Callers 1

train_featureFunction · 0.85

Calls 3

freeze_bn_layer_trainFunction · 0.90
PoseLossFunction · 0.50

Tested by

no test coverage detected