Skip to content
Merged
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
19 changes: 12 additions & 7 deletions sdm/ensemble.py
Original file line number Diff line number Diff line change
Expand Up @@ -212,15 +212,20 @@ def _select_members(self, member_ids: Sequence[int]) -> Self:
group = self._groups[group_id]
selected_positions = tuple(positions)
group_ids[group_id] = new_group_id
start = selected_positions[0]
step = (
selected_positions[1] - start
if len(selected_positions) > 1
else 1
)
stop = start + step * len(selected_positions)
if selected_positions == tuple(range(group.size(0))):
groups.append(group)
elif len(selected_positions) == 1:
groups.append(
cast(
TableTensor,
group.narrow(0, selected_positions[0], 1),
)
)
elif step > 0 and selected_positions == tuple(
range(start, stop, step)
):
# Evenly spaced members are selected as a view, not a copy.
groups.append(group[start:stop:step])
else:
groups.append(
cast(
Expand Down
22 changes: 9 additions & 13 deletions sdm/processing/common/choice.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

from collections.abc import Iterator
from typing import Literal, cast

import torch
Expand Down Expand Up @@ -79,14 +80,13 @@ def _draw_option_ids(
def _tables_by_option(
self,
ensemble_table: EnsembleTable,
) -> dict[int, EnsembleTable]:
) -> Iterator[tuple[int, EnsembleTable]]:
member_ids_by_option: dict[int, list[int]] = {}
for member_id, option_id in enumerate(self._option_ids):
member_ids_by_option.setdefault(option_id, []).append(member_id)
return {
option_id: ensemble_table[member_ids]
for option_id, member_ids in sorted(member_ids_by_option.items())
}
# Select lazily, since selecting members may copy them.
for option_id, member_ids in sorted(member_ids_by_option.items()):
yield option_id, ensemble_table[member_ids]

def _gather_outputs(
self,
Expand Down Expand Up @@ -124,7 +124,7 @@ def _fit_ensemble(
ensemble_table,
generator=generator,
)
for option_id, table in self._tables_by_option(ensemble_table).items():
for option_id, table in self._tables_by_option(ensemble_table):
self.options[option_id].fit_ensemble(table, generator=generator)

def _fit_transform_ensemble(
Expand All @@ -142,9 +142,7 @@ def _fit_transform_ensemble(
table,
generator=generator,
)
for option_id, table in self._tables_by_option(
ensemble_table
).items()
for option_id, table in self._tables_by_option(ensemble_table)
}
return self._gather_outputs(ensemble_table, outputs)

Expand All @@ -155,9 +153,7 @@ def _transform_ensemble(
self._check_num_members(ensemble_table)
outputs = {
option_id: self.options[option_id].transform_ensemble(table)
for option_id, table in self._tables_by_option(
ensemble_table
).items()
for option_id, table in self._tables_by_option(ensemble_table)
}
return self._gather_outputs(ensemble_table, outputs)

Expand All @@ -167,7 +163,7 @@ def _inverse_transform_ensemble(
) -> EnsembleTable:
self._check_num_members(ensemble_table)
outputs = {}
for option_id, table in self._tables_by_option(ensemble_table).items():
for option_id, table in self._tables_by_option(ensemble_table):
processor = self.options[option_id]
if not isinstance(processor, EnsembleInvertibleMixin):
raise TypeError(
Expand Down
30 changes: 15 additions & 15 deletions sdm/processing/numerical/constant.py
Original file line number Diff line number Diff line change
Expand Up @@ -180,21 +180,6 @@ def _transform_ensemble(
"ensemble members before transform."
)

if sum(
group.size(0) for group in ensemble_table._iter_groups()
) == len(ensemble_table):
tables = [
self._select_columns(
ensemble_table[member_id],
kept_indices,
)
for member_id, kept_indices in enumerate(self._kept_indices)
]
return ensemble_table.replace_tables(
tables=tables,
member_table_ids=range(len(ensemble_table)),
)

member_ids_by_kept_indices: dict[tuple[int, ...], list[int]] = {}
for member_id, kept_indices in enumerate(self._kept_indices):
member_ids_by_kept_indices.setdefault(kept_indices, []).append(
Expand All @@ -210,6 +195,21 @@ def _transform_ensemble(
]
)

if sum(
group.size(0) for group in ensemble_table._iter_groups()
) == len(ensemble_table):
tables = [
self._select_columns(
ensemble_table[member_id],
kept_indices,
)
for member_id, kept_indices in enumerate(self._kept_indices)
]
return ensemble_table.replace_tables(
tables=tables,
member_table_ids=range(len(ensemble_table)),
)

outputs: dict[tuple[int, ...], EnsembleTable] = {}
for kept_indices, member_ids in member_ids_by_kept_indices.items():
selected = ensemble_table[member_ids]
Expand Down
15 changes: 15 additions & 0 deletions test/test_ensemble.py
Original file line number Diff line number Diff line change
Expand Up @@ -182,6 +182,21 @@ def test_getitem_preserves_groups_and_order() -> None:
assert output[1].equal(tables[2])


@pytest.mark.parametrize("member_ids", [(0, 2, 4), (4, 2, 0), (0, 3, 4)])
def test_getitem_strided_members(member_ids: tuple[int, ...]) -> None:
values = torch.arange(30).float().reshape(5, 3, 2)
table = EnsembleTable(
groups=(TableTensor.from_tensor(values),),
locations=tuple((0, index) for index in range(5)),
)

selected = table[member_ids]

assert len(selected) == len(member_ids)
for position, index in enumerate(member_ids):
torch.testing.assert_close(selected[position].numerical, values[index])


def test_concatenate_columns_preserves_member_order() -> None:
left = EnsembleTable.from_tables(
tables=(
Expand Down
Loading