| 78 | os.makedirs(os.path.join(self.logdir, 'images_val'), exist_ok=True) |
| 79 | |
| 80 | def prepare_batch_data(self, batch): |
| 81 | lrm_generator_input = {} |
| 82 | render_gt = {} |
| 83 | |
| 84 | # input images |
| 85 | images = batch['input_images'] |
| 86 | images = v2.functional.resize( |
| 87 | images, self.input_size, interpolation=3, antialias=True).clamp(0, 1) |
| 88 | |
| 89 | lrm_generator_input['images'] = images.to(self.device) |
| 90 | |
| 91 | # input cameras and render cameras |
| 92 | input_c2ws = batch['input_c2ws'] |
| 93 | input_Ks = batch['input_Ks'] |
| 94 | target_c2ws = batch['target_c2ws'] |
| 95 | |
| 96 | render_c2ws = torch.cat([input_c2ws, target_c2ws], dim=1) |
| 97 | render_w2cs = torch.linalg.inv(render_c2ws) |
| 98 | |
| 99 | input_extrinsics = input_c2ws.flatten(-2) |
| 100 | input_extrinsics = input_extrinsics[:, :, :12] |
| 101 | input_intrinsics = input_Ks.flatten(-2) |
| 102 | input_intrinsics = torch.stack([ |
| 103 | input_intrinsics[:, :, 0], input_intrinsics[:, :, 4], |
| 104 | input_intrinsics[:, :, 2], input_intrinsics[:, :, 5], |
| 105 | ], dim=-1) |
| 106 | cameras = torch.cat([input_extrinsics, input_intrinsics], dim=-1) |
| 107 | |
| 108 | # add noise to input_cameras |
| 109 | cameras = cameras + torch.rand_like(cameras) * 0.04 - 0.02 |
| 110 | |
| 111 | lrm_generator_input['cameras'] = cameras.to(self.device) |
| 112 | lrm_generator_input['render_cameras'] = render_w2cs.to(self.device) |
| 113 | |
| 114 | # target images |
| 115 | target_images = torch.cat([batch['input_images'], batch['target_images']], dim=1) |
| 116 | target_depths = torch.cat([batch['input_depths'], batch['target_depths']], dim=1) |
| 117 | target_alphas = torch.cat([batch['input_alphas'], batch['target_alphas']], dim=1) |
| 118 | target_normals = torch.cat([batch['input_normals'], batch['target_normals']], dim=1) |
| 119 | |
| 120 | render_size = self.render_size |
| 121 | target_images = v2.functional.resize( |
| 122 | target_images, render_size, interpolation=3, antialias=True).clamp(0, 1) |
| 123 | target_depths = v2.functional.resize( |
| 124 | target_depths, render_size, interpolation=0, antialias=True) |
| 125 | target_alphas = v2.functional.resize( |
| 126 | target_alphas, render_size, interpolation=0, antialias=True) |
| 127 | target_normals = v2.functional.resize( |
| 128 | target_normals, render_size, interpolation=3, antialias=True) |
| 129 | |
| 130 | lrm_generator_input['render_size'] = render_size |
| 131 | |
| 132 | render_gt['target_images'] = target_images.to(self.device) |
| 133 | render_gt['target_depths'] = target_depths.to(self.device) |
| 134 | render_gt['target_alphas'] = target_alphas.to(self.device) |
| 135 | render_gt['target_normals'] = target_normals.to(self.device) |
| 136 | |
| 137 | return lrm_generator_input, render_gt |