-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathsource_dataset_processor.py
More file actions
217 lines (197 loc) · 8.92 KB
/
Copy pathsource_dataset_processor.py
File metadata and controls
217 lines (197 loc) · 8.92 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
import logging
import os
from abc import ABC, abstractmethod
from typing import Any, final
from standard_e2e.caching.adapters import AbstractAdapter
from standard_e2e.caching.segment_context import SegmentContextAggregator
from standard_e2e.data_structures import (
FrameIndexData,
StandardFrameData,
TransformedFrameData,
)
from standard_e2e.enums import StandardFrameDataField
from standard_e2e.indexing import IndexDataGenerator
from standard_e2e.utils import _check_list_of_objects_or_none
class SourceDatasetProcessor(ABC):
"""Abstract base class for processing source datasets."""
def __init__(
self,
common_output_path: str,
split: str,
index_data_generator: IndexDataGenerator | None = None,
adapters: list[AbstractAdapter] | None = None,
context_aggregators: list[SegmentContextAggregator] | None = None,
):
_check_list_of_objects_or_none(adapters, AbstractAdapter)
_check_list_of_objects_or_none(context_aggregators, SegmentContextAggregator)
if not isinstance(index_data_generator, (IndexDataGenerator, type(None))):
raise TypeError(
"index_data_generator must be an instance of IndexDataGenerator"
f"or None, got {type(index_data_generator)}"
)
self._split = split
self._common_output_path = common_output_path
self._specific_output_path = self._prepare_output_directory()
self._inner_path = os.path.relpath(
self._specific_output_path, common_output_path
)
self._adapters = self._get_default_adapters() if adapters is None else adapters
self._context_aggregators = (
self._get_default_context_aggregators()
if context_aggregators is None
else context_aggregators
)
# Union of ``StandardFrameData`` attributes the registered adapter
# chain reads. Per-dataset ``_prepare_standardized_frame_data``
# implementations consult ``self.needs_attr(...)`` to skip
# building modalities no adapter consumes (lazy load).
self._consumed_attrs: set[StandardFrameDataField] = set()
for _adapter in self._adapters:
self._consumed_attrs |= _adapter.consumes_attrs
self._index_data_generator = (
index_data_generator if index_data_generator else IndexDataGenerator()
)
if self._split not in self.allowed_splits:
raise ValueError(
f"Invalid split: {self._split}. Must be one of {self.allowed_splits}."
)
logging.info("Initialized %s processor", self.dataset_name)
logging.info("Using adapters: %s", [a.name for a in self._adapters])
logging.info("Consumed SFD attrs: %s", sorted(self._consumed_attrs))
logging.info("Specific output path: %s", self._specific_output_path)
def needs_attr(self, attr: StandardFrameDataField) -> bool:
"""Whether at least one registered adapter reads this
``StandardFrameData`` field. Used by per-dataset processors to skip
expensive modality builds (cameras, lidar, hd_map, detections, …)
when no adapter would consume them. ``True`` when ``attr`` is in the
consumed-attrs union, plus a hard-coded special case: the
identifier / index fields are always treated as needed since they
are required for the cache + index regardless of adapter chain.
"""
always = {
StandardFrameDataField.DATASET_NAME,
StandardFrameDataField.SPLIT,
StandardFrameDataField.SEGMENT_ID,
StandardFrameDataField.FRAME_ID,
StandardFrameDataField.TIMESTAMP,
StandardFrameDataField.GLOBAL_POSITION,
}
if attr in always:
return True
return attr in self._consumed_attrs
def _get_default_adapters(self) -> list[AbstractAdapter]:
raise NotImplementedError("Subclasses must implement this method.")
def _get_default_context_aggregators(self) -> list[SegmentContextAggregator]:
return []
@final
def _prepare_output_directory(self) -> str:
"""Prepare the output directory for specific processed data."""
specific_output_path = os.path.join(
self._common_output_path, self.dataset_name, self.split
)
if not os.path.exists(specific_output_path):
os.makedirs(specific_output_path)
logging.info("Created output directory: %s", specific_output_path)
else:
logging.warning("Output directory already exists: %s", specific_output_path)
return specific_output_path
@property
@abstractmethod
def dataset_name(self) -> str:
"""Return the name of the dataset."""
raise NotImplementedError("Subclasses must implement this method.")
@property
def allowed_splits(self) -> list[str]:
"""Return the list of allowed splits for the dataset."""
raise NotImplementedError("Subclasses must implement this method.")
@property
def context_aggregators(self):
return self._context_aggregators
@property
def adapters(self) -> list[AbstractAdapter]:
"""The registered adapter chain (read-only view).
Exposed so the converter can serialize each adapter's
:attr:`~standard_e2e.caching.adapters.abstract_adapter.AbstractAdapter.spec`
into the per-(dataset, split) ``dataset_info.yaml``.
"""
return list(self._adapters)
@final
def process_frame(
self, raw_frame_data: Any
) -> tuple[TransformedFrameData, FrameIndexData]:
standard_frame_data = self._prepare_standardized_frame_data(raw_frame_data)
if not isinstance(standard_frame_data, StandardFrameData):
raise TypeError(
"_prepare_standardized_frame_data must return StandardFrameData, "
f"got {type(standard_frame_data)}"
)
transformed_modalities = {}
for adapter in self._adapters:
transformed_modalities.update(adapter.transform(standard_frame_data))
# Merge each adapter's per-frame metadata into aux_data so the .npz
# carries adapter-side configuration (e.g. the HD-map BEV channel list)
# that downstream consumers need to interpret modality outputs.
merged_aux_data: dict | None
if standard_frame_data.aux_data is None:
merged_aux_data = None
else:
merged_aux_data = dict(standard_frame_data.aux_data)
for adapter in self._adapters:
adapter_meta = adapter.metadata
if not adapter_meta:
continue
if merged_aux_data is None:
merged_aux_data = {}
merged_aux_data.update(adapter_meta)
transformed_frame_data = TransformedFrameData(
dataset_name=standard_frame_data.dataset_name,
segment_id=standard_frame_data.segment_id,
frame_id=standard_frame_data.frame_id,
timestamp=standard_frame_data.timestamp,
split=standard_frame_data.split,
global_position=standard_frame_data.global_position,
aux_data=merged_aux_data,
extra_index_data=standard_frame_data.extra_index_data,
_modality_data=transformed_modalities,
)
frame_index_data = self._index_data_generator.generate_index_data(
transformed_frame_data
)
return transformed_frame_data, frame_index_data
@abstractmethod
def _prepare_standardized_frame_data(
self, raw_frame_data: Any
) -> StandardFrameData:
"""Process a single frame of data."""
# Implement the logic to process a single frame
raise NotImplementedError("Subclasses must implement this method.")
@final
def process_frame_and_save_data(self, raw_frame_data: Any) -> FrameIndexData:
"""
Process a single frame of raw data, save the processed frame data to disk,
and return the corresponding FrameIndexData.
"""
frame_data: TransformedFrameData
frame_index_data: FrameIndexData
frame_data, frame_index_data = self.process_frame(raw_frame_data)
filename = frame_data.filename
if filename is None:
raise ValueError("Frame data must have a filename before saving.")
frame_data.to_npz(os.path.join(self._common_output_path, filename))
return frame_index_data
@property
def split(self) -> str:
"""Return the dataset split."""
return self._split
@property
def output_path(self) -> str:
"""Return the output path for the processed dataset."""
return self._common_output_path
@property
def inner_path(self) -> str:
"""Return the inner path relative to the common output path."""
return self._inner_path
@property
def specific_output_path(self) -> str:
"""Return the specific output path for the dataset."""
return self._specific_output_path