MCPcopy Create free account
hub / github.com/CompVis/diff2flow / DatasetPreprocessor

Class DatasetPreprocessor

diff2flow/dataset/depth_preprocessing.py:115–187  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

113
114
115class 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:

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected