(
self, ply_dir: str, frame_num: int, ply_filename: Optional[str] = None
)
| 497 | raise ValueError(f"Unknown compression method: {cfg.compression}") |
| 498 | |
| 499 | def load_ply_sequences( |
| 500 | self, ply_dir: str, frame_num: int, ply_filename: Optional[str] = None |
| 501 | ) -> List[torch.nn.ParameterDict]: |
| 502 | assert frame_num > 0, "frame_num must be greater than 0" |
| 503 | |
| 504 | splats_list = [] |
| 505 | if frame_num > 1: |
| 506 | self.ply_filename_list = sorted(glob.glob(os.path.join(ply_dir, "*.ply"))) |
| 507 | |
| 508 | for filename in tqdm(self.ply_filename_list[:frame_num], desc="Loading .ply file"): |
| 509 | splats = load_ply(filename) |
| 510 | splats_list.append(splats.to("cuda")) |
| 511 | else: |
| 512 | self.ply_filename_list = [ply_filename] |
| 513 | assert ply_filename is not None, "ply_filename must be provided if frame_num is 1" |
| 514 | splats = load_ply(ply_filename) |
| 515 | splats_list = [splats.to("cuda")] |
| 516 | |
| 517 | return splats_list |
| 518 | |
| 519 | def set_up_datasets( |
| 520 | self, data_dir: str, frame_num: int, cfg: Config |
no test coverage detected