| 137 | return lrm_generator_input, render_gt |
| 138 | |
| 139 | def prepare_validation_batch_data(self, batch): |
| 140 | lrm_generator_input = {} |
| 141 | |
| 142 | # input images |
| 143 | images = batch['input_images'] |
| 144 | images = v2.functional.resize( |
| 145 | images, self.input_size, interpolation=3, antialias=True).clamp(0, 1) |
| 146 | |
| 147 | lrm_generator_input['images'] = images.to(self.device) |
| 148 | |
| 149 | # input cameras |
| 150 | input_c2ws = batch['input_c2ws'].flatten(-2) |
| 151 | input_Ks = batch['input_Ks'].flatten(-2) |
| 152 | |
| 153 | input_extrinsics = input_c2ws[:, :, :12] |
| 154 | input_intrinsics = torch.stack([ |
| 155 | input_Ks[:, :, 0], input_Ks[:, :, 4], |
| 156 | input_Ks[:, :, 2], input_Ks[:, :, 5], |
| 157 | ], dim=-1) |
| 158 | cameras = torch.cat([input_extrinsics, input_intrinsics], dim=-1) |
| 159 | |
| 160 | lrm_generator_input['cameras'] = cameras.to(self.device) |
| 161 | |
| 162 | # render cameras |
| 163 | render_c2ws = batch['render_c2ws'] |
| 164 | render_w2cs = torch.linalg.inv(render_c2ws) |
| 165 | |
| 166 | lrm_generator_input['render_cameras'] = render_w2cs.to(self.device) |
| 167 | lrm_generator_input['render_size'] = 384 |
| 168 | |
| 169 | return lrm_generator_input |
| 170 | |
| 171 | def forward_lrm_generator(self, images, cameras, render_cameras, render_size=512): |
| 172 | planes = torch.utils.checkpoint.checkpoint( |