| 113 | |
| 114 | |
| 115 | class DatasetPreprocessor: |
| 116 | def __init__( |
| 117 | self, |
| 118 | size=None, |
| 119 | depth_key="depth", |
| 120 | out_channels=3, |
| 121 | return_valid_mask=True, |
| 122 | exclude_keys_for_resize=None, |
| 123 | keep_raw_depth=False, |
| 124 | ): |
| 125 | self.size = tuple(size) if size is not None else None |
| 126 | self.out_channels = out_channels |
| 127 | self.return_valid_mask = return_valid_mask |
| 128 | self.exclude_keys_for_resize = exclude_keys_for_resize or [] |
| 129 | self.saved_processed_sample = None |
| 130 | self.keep_raw_depth = keep_raw_depth |
| 131 | self.depth_key = depth_key |
| 132 | |
| 133 | def preprocess_sample(self, sample): |
| 134 | """ get dataset name """ |
| 135 | if "dataset" not in sample: |
| 136 | sample["dataset"] = "hypersim" # default to hypersim |
| 137 | dataset_name = sample.get("dataset") |
| 138 | assert isinstance(self.size, tuple) or isinstance(self.size, list) or self.size is None, "Invalid size" |
| 139 | try: |
| 140 | dataset_name = dataset_name.decode() # convert bytes to string |
| 141 | except AttributeError: |
| 142 | pass |
| 143 | sample["dataset"] = dataset_name |
| 144 | |
| 145 | """ exceptions """ |
| 146 | # if dataset_name in ["depth_anything"]: |
| 147 | # ... |
| 148 | |
| 149 | """ Preprocess depth map """ |
| 150 | depth = sample[self.depth_key] |
| 151 | depth, valid_mask = preprocess_depth(depth, dataset_name, self.out_channels, self.keep_raw_depth) |
| 152 | |
| 153 | if self.return_valid_mask: |
| 154 | if "valid_mask" in sample: |
| 155 | valid_mask_sample = sample["valid_mask"] |
| 156 | # merge the valid masks |
| 157 | valid_mask = valid_mask * valid_mask_sample |
| 158 | sample["valid_mask"] = valid_mask |
| 159 | sample[self.depth_key] = depth |
| 160 | |
| 161 | """ convert to tensor and resize """ |
| 162 | if self.size is not None: |
| 163 | for key in sample: |
| 164 | if isinstance(sample[key], np.ndarray): |
| 165 | sample[key] = torch.tensor(sample[key], dtype=torch.float32) |
| 166 | if key in self.exclude_keys_for_resize: |
| 167 | continue |
| 168 | if isinstance(sample[key], torch.Tensor): |
| 169 | sample[key] = resize(sample[key], size=self.size) |
| 170 | |
| 171 | # filter the resized valid mask with 1 |
| 172 | if self.return_valid_mask: |
nothing calls this directly
no outgoing calls
no test coverage detected