(
self,
lrm_generator_config,
lrm_path=None,
input_size=256,
render_size=192,
)
| 13 | |
| 14 | class MVRecon(pl.LightningModule): |
| 15 | def __init__( |
| 16 | self, |
| 17 | lrm_generator_config, |
| 18 | lrm_path=None, |
| 19 | input_size=256, |
| 20 | render_size=192, |
| 21 | ): |
| 22 | super(MVRecon, self).__init__() |
| 23 | |
| 24 | self.input_size = input_size |
| 25 | self.render_size = render_size |
| 26 | |
| 27 | # init modules |
| 28 | self.lrm_generator = instantiate_from_config(lrm_generator_config) |
| 29 | if lrm_path is not None: |
| 30 | lrm_ckpt = torch.load(lrm_path) |
| 31 | self.lrm_generator.load_state_dict(lrm_ckpt['weights'], strict=False) |
| 32 | |
| 33 | self.lpips = LearnedPerceptualImagePatchSimilarity(net_type='vgg') |
| 34 | |
| 35 | self.validation_step_outputs = [] |
| 36 | |
| 37 | def on_fit_start(self): |
| 38 | if self.global_rank == 0: |
nothing calls this directly
no test coverage detected