| 102 | return lrm_generator_input, render_gt |
| 103 | |
| 104 | def prepare_validation_batch_data(self, batch): |
| 105 | lrm_generator_input = {} |
| 106 | |
| 107 | # input images |
| 108 | images = batch['input_images'] |
| 109 | images = v2.functional.resize( |
| 110 | images, self.input_size, interpolation=3, antialias=True).clamp(0, 1) |
| 111 | |
| 112 | lrm_generator_input['images'] = images.to(self.device) |
| 113 | |
| 114 | input_c2ws = batch['input_c2ws'].flatten(-2) |
| 115 | input_Ks = batch['input_Ks'].flatten(-2) |
| 116 | |
| 117 | input_extrinsics = input_c2ws[:, :, :12] |
| 118 | input_intrinsics = torch.stack([ |
| 119 | input_Ks[:, :, 0], input_Ks[:, :, 4], |
| 120 | input_Ks[:, :, 2], input_Ks[:, :, 5], |
| 121 | ], dim=-1) |
| 122 | cameras = torch.cat([input_extrinsics, input_intrinsics], dim=-1) |
| 123 | |
| 124 | lrm_generator_input['cameras'] = cameras.to(self.device) |
| 125 | |
| 126 | render_c2ws = batch['render_c2ws'].flatten(-2) |
| 127 | render_Ks = batch['render_Ks'].flatten(-2) |
| 128 | render_cameras = torch.cat([render_c2ws, render_Ks], dim=-1) |
| 129 | |
| 130 | lrm_generator_input['render_cameras'] = render_cameras.to(self.device) |
| 131 | lrm_generator_input['render_size'] = 384 |
| 132 | lrm_generator_input['crop_params'] = None |
| 133 | |
| 134 | return lrm_generator_input |
| 135 | |
| 136 | def forward_lrm_generator( |
| 137 | self, |