(self, data_path)
| 237 | return pred, loss_dict |
| 238 | |
| 239 | def train(self, data_path): |
| 240 | # load data |
| 241 | self._load_data(data_path, sdf_renorm=self.sdf_renorm) |
| 242 | self.net.reset_aabb(self.aabb) |
| 243 | |
| 244 | # set optimizer |
| 245 | self._set_optimizer(self.init_lr, self.min_lr_ratio) |
| 246 | |
| 247 | # set tensorboard writer |
| 248 | self.tb = SummaryWriter(os.path.join(self.log_dir, "tblog")) |
| 249 | |
| 250 | pbar = tqdm(range(self.n_iters)) |
| 251 | self.step = 0 |
| 252 | for i in pbar: |
| 253 | self.step = i |
| 254 | self.net.train() |
| 255 | |
| 256 | data = self._sample_batch(self.batch_size) |
| 257 | |
| 258 | rec, loss_dict = self._forward_batch(data) |
| 259 | |
| 260 | self.update_network(loss_dict) |
| 261 | |
| 262 | loss_values = {k: v.item() for k, v in loss_dict.items()} |
| 263 | self.tb.add_scalars("loss", loss_values, global_step=i) |
| 264 | pbar.set_postfix(loss_values) |
| 265 | |
| 266 | if i == 0 or (i + 1) % (self.n_iters // 5) == 0: |
| 267 | self._visualize_batch(i) |
| 268 | |
| 269 | if (i + 1) % (self.n_iters // 5) == 0: |
| 270 | eval_stat = self.evaluate() |
| 271 | for name in ["tsdf_l1", "tsdf_rel", "tsdf_acc"]: |
| 272 | stat_dict = {k: v for k, v in eval_stat.items() if name in k} |
| 273 | self.tb.add_scalars(name, stat_dict, global_step=i) |
| 274 | if self.data_type != "sdf": |
| 275 | self.tb.add_scalar("surf_tex_l1_error", eval_stat["surf_tex_l1_error"], global_step=i) |
| 276 | |
| 277 | eval_stat = self.evaluate() |
| 278 | with open(os.path.join(self.log_dir, "eval_stat.json"), "w") as f: |
| 279 | json.dump(eval_stat, f, indent=2) |
| 280 | self.save_ckpt("final") |
| 281 | |
| 282 | @torch.no_grad() |
| 283 | def _visualize_batch(self, step): |
no test coverage detected