(
self, local_rank: int, world_rank, world_size: int, cfg: Config
)
| 434 | |
| 435 | class Runner: |
| 436 | def __init__( |
| 437 | self, local_rank: int, world_rank, world_size: int, cfg: Config |
| 438 | ) -> None: |
| 439 | os.makedirs(cfg.result_dir, exist_ok=True) |
| 440 | set_random_seed(42) |
| 441 | verify_random_seed(cfg.result_dir) |
| 442 | |
| 443 | self.cfg = cfg |
| 444 | self.world_rank = world_rank |
| 445 | self.local_rank = local_rank |
| 446 | self.world_size = world_size |
| 447 | self.device = f"cuda:{local_rank}" |
| 448 | |
| 449 | # Where to dump results. |
| 450 | os.makedirs(cfg.result_dir, exist_ok=True) |
| 451 | |
| 452 | # Setup output directories. |
| 453 | # self.ckpt_dir = f"{cfg.result_dir}/ckpts" |
| 454 | # os.makedirs(self.ckpt_dir, exist_ok=True) |
| 455 | self.stats_dir = f"{cfg.result_dir}/stats" |
| 456 | os.makedirs(self.stats_dir, exist_ok=True) |
| 457 | self.render_dir = f"{cfg.result_dir}/renders" |
| 458 | os.makedirs(self.render_dir, exist_ok=True) |
| 459 | |
| 460 | # Losses & Metrics. |
| 461 | self.ssim = StructuralSimilarityIndexMeasure(data_range=1.0).to(self.device) |
| 462 | self.psnr = PeakSignalNoiseRatio(data_range=1.0).to(self.device) |
| 463 | |
| 464 | if cfg.with_lpips: |
| 465 | if cfg.lpips_net == "alex": |
| 466 | self.lpips = LearnedPerceptualImagePatchSimilarity( |
| 467 | net_type="alex", normalize=True |
| 468 | ).to(self.device) |
| 469 | elif cfg.lpips_net == "vgg": |
| 470 | # The 3DGS official repo uses lpips vgg, which is equivalent with the following: |
| 471 | self.lpips = LearnedPerceptualImagePatchSimilarity( |
| 472 | net_type="vgg", normalize=False |
| 473 | ).to(self.device) |
| 474 | else: |
| 475 | raise ValueError(f"Unknown LPIPS network: {cfg.lpips_net}") |
| 476 | else: |
| 477 | self.lpips = None |
| 478 | |
| 479 | # frame num |
| 480 | self.frame_num = cfg.frame_num |
| 481 | |
| 482 | # load ply sequences |
| 483 | self.splats_list = self.load_ply_sequences(cfg.ply_dir, cfg.frame_num, cfg.ply_filename) |
| 484 | |
| 485 | # load dataset |
| 486 | self.trainset_list, self.valset_list = self.set_up_datasets(cfg.data_dir, cfg.frame_num, cfg) |
| 487 | |
| 488 | self.compression_cfg = cfg.compression_cfg.to_dict() |
| 489 | |
| 490 | if cfg.compression == "seq_hevc": |
| 491 | self.compression_method = SeqHevcCompression(**self.compression_cfg) |
| 492 | elif cfg.compression == "seq_yuv": |
| 493 | self.compression_method = SeqYUVCompression(**self.compression_cfg) |
nothing calls this directly
no test coverage detected