| 40 | os.makedirs(os.path.join(self.logdir, 'images_val'), exist_ok=True) |
| 41 | |
| 42 | def prepare_batch_data(self, batch): |
| 43 | lrm_generator_input = {} |
| 44 | render_gt = {} # for supervision |
| 45 | |
| 46 | # input images |
| 47 | images = batch['input_images'] |
| 48 | images = v2.functional.resize( |
| 49 | images, self.input_size, interpolation=3, antialias=True).clamp(0, 1) |
| 50 | |
| 51 | lrm_generator_input['images'] = images.to(self.device) |
| 52 | |
| 53 | # input cameras and render cameras |
| 54 | input_c2ws = batch['input_c2ws'].flatten(-2) |
| 55 | input_Ks = batch['input_Ks'].flatten(-2) |
| 56 | target_c2ws = batch['target_c2ws'].flatten(-2) |
| 57 | target_Ks = batch['target_Ks'].flatten(-2) |
| 58 | render_cameras_input = torch.cat([input_c2ws, input_Ks], dim=-1) |
| 59 | render_cameras_target = torch.cat([target_c2ws, target_Ks], dim=-1) |
| 60 | render_cameras = torch.cat([render_cameras_input, render_cameras_target], dim=1) |
| 61 | |
| 62 | input_extrinsics = input_c2ws[:, :, :12] |
| 63 | input_intrinsics = torch.stack([ |
| 64 | input_Ks[:, :, 0], input_Ks[:, :, 4], |
| 65 | input_Ks[:, :, 2], input_Ks[:, :, 5], |
| 66 | ], dim=-1) |
| 67 | cameras = torch.cat([input_extrinsics, input_intrinsics], dim=-1) |
| 68 | |
| 69 | # add noise to input cameras |
| 70 | cameras = cameras + torch.rand_like(cameras) * 0.04 - 0.02 |
| 71 | |
| 72 | lrm_generator_input['cameras'] = cameras.to(self.device) |
| 73 | lrm_generator_input['render_cameras'] = render_cameras.to(self.device) |
| 74 | |
| 75 | # target images |
| 76 | target_images = torch.cat([batch['input_images'], batch['target_images']], dim=1) |
| 77 | target_depths = torch.cat([batch['input_depths'], batch['target_depths']], dim=1) |
| 78 | target_alphas = torch.cat([batch['input_alphas'], batch['target_alphas']], dim=1) |
| 79 | |
| 80 | # random crop |
| 81 | render_size = np.random.randint(self.render_size, 513) |
| 82 | target_images = v2.functional.resize( |
| 83 | target_images, render_size, interpolation=3, antialias=True).clamp(0, 1) |
| 84 | target_depths = v2.functional.resize( |
| 85 | target_depths, render_size, interpolation=0, antialias=True) |
| 86 | target_alphas = v2.functional.resize( |
| 87 | target_alphas, render_size, interpolation=0, antialias=True) |
| 88 | |
| 89 | crop_params = v2.RandomCrop.get_params( |
| 90 | target_images, output_size=(self.render_size, self.render_size)) |
| 91 | target_images = v2.functional.crop(target_images, *crop_params) |
| 92 | target_depths = v2.functional.crop(target_depths, *crop_params)[:, :, 0:1] |
| 93 | target_alphas = v2.functional.crop(target_alphas, *crop_params)[:, :, 0:1] |
| 94 | |
| 95 | lrm_generator_input['render_size'] = render_size |
| 96 | lrm_generator_input['crop_params'] = crop_params |
| 97 | |
| 98 | render_gt['target_images'] = target_images.to(self.device) |
| 99 | render_gt['target_depths'] = target_depths.to(self.device) |