From 475739c1670387b077a224768b09387ceb92176d Mon Sep 17 00:00:00 2001 From: Jingang Qu Date: Mon, 28 Sep 2026 22:40:27 +0200 Subject: [PATCH] Avoid ensemble selection copies - Select evenly spaced ensemble members as views. - Select shared DropConstantColumns groups before re-stacking members. - Select Choice members one option at a time. Signed-off-by: Jingang Qu --- sdm/ensemble.py | 19 +++++++++++------- sdm/processing/common/choice.py | 22 +++++++++----------- sdm/processing/numerical/constant.py | 30 ++++++++++++++-------------- test/test_ensemble.py | 15 ++++++++++++++ 4 files changed, 51 insertions(+), 35 deletions(-) diff --git a/sdm/ensemble.py b/sdm/ensemble.py index 70e629a99..3b4dc6962 100644 --- a/sdm/ensemble.py +++ b/sdm/ensemble.py @@ -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( diff --git a/sdm/processing/common/choice.py b/sdm/processing/common/choice.py index c4660235f..25da3bac8 100644 --- a/sdm/processing/common/choice.py +++ b/sdm/processing/common/choice.py @@ -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 @@ -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, @@ -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( @@ -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) @@ -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) @@ -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( diff --git a/sdm/processing/numerical/constant.py b/sdm/processing/numerical/constant.py index 4cbaaeb2f..4fe2873cf 100644 --- a/sdm/processing/numerical/constant.py +++ b/sdm/processing/numerical/constant.py @@ -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( @@ -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] diff --git a/test/test_ensemble.py b/test/test_ensemble.py index c3b39b6f6..7aa720844 100644 --- a/test/test_ensemble.py +++ b/test/test_ensemble.py @@ -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=(