(self, hparams)
| 40 | |
| 41 | class ImplicitVideoSystem(LightningModule): |
| 42 | def __init__(self, hparams): |
| 43 | super(ImplicitVideoSystem, self).__init__() |
| 44 | self.save_hyperparameters(hparams) |
| 45 | self.color_loss = loss_dict['mse'](coef=1) |
| 46 | if hparams.save_video: |
| 47 | self.video_visualizer = VideoVisualizer(fps=hparams.fps) |
| 48 | self.raw_video_visualizer = VideoVisualizer(fps=hparams.fps) |
| 49 | self.dual_video_visualizer = VideoVisualizer(fps=hparams.fps) |
| 50 | |
| 51 | self.models_to_train=[] |
| 52 | self.embedding_xyz = Embedding(2, 8) |
| 53 | self.embeddings = {'xyz': self.embedding_xyz} |
| 54 | self.models = {} |
| 55 | |
| 56 | # Construct normalized meshgrid. |
| 57 | h = self.hparams.img_wh[1] |
| 58 | w = self.hparams.img_wh[0] |
| 59 | self.h = h |
| 60 | self.w = w |
| 61 | |
| 62 | if self.hparams.mask_dir: |
| 63 | self.num_models = len(self.hparams.mask_dir) |
| 64 | else: |
| 65 | self.num_models = 1 |
| 66 | |
| 67 | # Decide the number of deformable mlp. |
| 68 | if hparams.encode_w: |
| 69 | # Multiple deformation MLP. |
| 70 | # Progressive Training for the Deformation (Annealed PE). |
| 71 | # No trainable parameters. |
| 72 | self.embeddings['xyz_w'] = [] |
| 73 | assert (isinstance(self.hparams.N_xyz_w, list)) |
| 74 | in_channels_xyz = [] |
| 75 | for i in range(self.num_models): |
| 76 | N_xyz_w = self.hparams.N_xyz_w[i] |
| 77 | in_channels_xyz += [2 + 2 * N_xyz_w * 2] |
| 78 | if hparams.annealed: |
| 79 | if hparams.deform_hash: |
| 80 | self.embedding_hash = AnnealedHash( |
| 81 | in_channels=2, |
| 82 | annealed_step=hparams.annealed_step, |
| 83 | annealed_begin_step=hparams.annealed_begin_step) |
| 84 | self.embeddings['aneal_hash'] = self.embedding_hash |
| 85 | else: |
| 86 | self.embedding_xyz_w = AnnealedEmbedding( |
| 87 | in_channels=2, |
| 88 | N_freqs=N_xyz_w, |
| 89 | annealed_step=hparams.annealed_step, |
| 90 | annealed_begin_step=hparams.annealed_begin_step) |
| 91 | self.embeddings['xyz_w'] += [self.embedding_xyz_w] |
| 92 | else: |
| 93 | self.embedding_xyz_w = Embedding(2, N_xyz_w) |
| 94 | self.embeddings['xyz_w'] += [self.embedding_xyz_w] |
| 95 | |
| 96 | for i in range(self.num_models): |
| 97 | embedding_w = torch.nn.Embedding(hparams.N_vocab_w, hparams.N_w) |
| 98 | torch.nn.init.uniform_(embedding_w.weight, -0.05, 0.05) |
| 99 | load_ckpt(embedding_w, hparams.weight_path, model_name=f'w_{i}') |
nothing calls this directly
no test coverage detected