Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
36 changes: 26 additions & 10 deletions extra_data/components.py
Original file line number Diff line number Diff line change
Expand Up @@ -425,6 +425,17 @@ def __repr__(self):
return (f"<{det}: Data interface for detector {self.detector_name!r} "
f"- {rp} data with {len(self.source_to_modno)} modules>")

def _with_train_sel_data(self, data_sel):
# Used in select_trains & split_trains below
# Using a copy to bypass the source & train checks in __init__
res = copy(self)
res.data = data_sel
res.frame_counts = self.frame_counts[res.data.train_ids]
res.train_ids_perframe = np.repeat(
res.frame_counts.index.values, res.frame_counts.values.astype(np.intp)
)
return res

def select_trains(self, trains):
"""Select a subset of trains from this data as a new object.

Expand All @@ -438,14 +449,7 @@ def select_trains(self, trains):
sel1 = det.select_trains(by_id[142844490 : 142844495])
sel2 = det.select_trains(by_id[[142844490, 142844493, 142844494]])
"""
# Using a copy to bypass the source & train checks in __init__
res = copy(self)
res.data = self.data.select_trains(trains)
res.frame_counts = self.frame_counts[res.data.train_ids]
res.train_ids_perframe = np.repeat(
res.frame_counts.index.values, res.frame_counts.values.astype(np.intp)
)
return res
return self._with_train_sel_data(self.data.select_trains(trains))

def split_trains(self, parts=None, trains_per_part=None, frames_per_part=None):
"""Split this data into chunks with a fraction of the trains each.
Expand Down Expand Up @@ -474,9 +478,21 @@ def split_trains(self, parts=None, trains_per_part=None, frames_per_part=None):
raise ValueError(
"One of parts, trains_per_part, frames_per_part must be specified"
)
if frames_per_part is not None:
# If possible, convert frames_per_part to trains_per_part, which
# uses the more efficient code path.
try:
fpt = self.frames_per_train
except ValueError:
pass
else:
ntpp = max(1, frames_per_part // fpt)
trains_per_part = min(ntpp, trains_per_part or ntpp)
frames_per_part = None

if frames_per_part is None:
for s in split_trains(len(self.train_ids), parts, trains_per_part):
yield self.select_trains(s)
for part in self.data.split_trains(parts=parts, trains_per_part=trains_per_part):
yield self._with_train_sel_data(part)
else:
# frames_per_part was specified. We don't assume that the number
# of frames per train is constant, so we'll iterate over trains
Expand Down
5 changes: 3 additions & 2 deletions extra_data/tests/test_components.py
Original file line number Diff line number Diff line change
Expand Up @@ -479,11 +479,12 @@ def test_split_trains(mock_fxe_raw_run):

# trains_per_part cuts off before frames_per_part
parts = list(det.split_trains(trains_per_part=3, frames_per_part=1024))
assert [len(p.train_ids) for p in parts] == ([3] * 6) + [2]
assert {len(p.train_ids) for p in parts} == {2, 3}

# parts cuts off before frames_per_part
parts = list(det.split_trains(parts=6, frames_per_part=1024))
assert [len(p.train_ids) for p in parts] == ([3] * 6) + [2]
assert len(parts) == 6
assert {len(p.train_ids) for p in parts} == {3, 4}

# frames_per_part > all frames in selection
parts = list(det.split_trains(frames_per_part=3000))
Expand Down