| 75 | self.viewmat.requires_grad = False |
| 76 | |
| 77 | def train( |
| 78 | self, |
| 79 | iterations: int = 1000, |
| 80 | lr: float = 0.01, |
| 81 | save_imgs: bool = False, |
| 82 | model_type: Literal["3dgs", "2dgs"] = "3dgs", |
| 83 | ): |
| 84 | optimizer = optim.Adam( |
| 85 | [self.rgbs, self.means, self.scales, self.opacities, self.quats], lr |
| 86 | ) |
| 87 | mse_loss = torch.nn.MSELoss() |
| 88 | frames = [] |
| 89 | times = [0] * 2 # rasterization, backward |
| 90 | K = torch.tensor( |
| 91 | [ |
| 92 | [self.focal, 0, self.W / 2], |
| 93 | [0, self.focal, self.H / 2], |
| 94 | [0, 0, 1], |
| 95 | ], |
| 96 | device=self.device, |
| 97 | ) |
| 98 | |
| 99 | if model_type == "3dgs": |
| 100 | rasterize_fnc = rasterization |
| 101 | elif model_type == "2dgs": |
| 102 | rasterize_fnc = rasterization_2dgs |
| 103 | |
| 104 | for iter in range(iterations): |
| 105 | start = time.time() |
| 106 | |
| 107 | renders = rasterize_fnc( |
| 108 | self.means, |
| 109 | self.quats / self.quats.norm(dim=-1, keepdim=True), |
| 110 | self.scales, |
| 111 | torch.sigmoid(self.opacities), |
| 112 | torch.sigmoid(self.rgbs), |
| 113 | self.viewmat[None], |
| 114 | K[None], |
| 115 | self.W, |
| 116 | self.H, |
| 117 | packed=False, |
| 118 | )[0] |
| 119 | out_img = renders[0] |
| 120 | torch.cuda.synchronize() |
| 121 | times[0] += time.time() - start |
| 122 | loss = mse_loss(out_img, self.gt_image) |
| 123 | optimizer.zero_grad() |
| 124 | start = time.time() |
| 125 | loss.backward() |
| 126 | torch.cuda.synchronize() |
| 127 | times[1] += time.time() - start |
| 128 | optimizer.step() |
| 129 | print(f"Iteration {iter + 1}/{iterations}, Loss: {loss.item()}") |
| 130 | |
| 131 | if save_imgs and iter % 5 == 0: |
| 132 | frames.append((out_img.detach().cpu().numpy() * 255).astype(np.uint8)) |
| 133 | if save_imgs: |
| 134 | # save them as a gif with PIL |