From b159f6f8dc6bd8e40dde3785e51a0ba20abcbd47 Mon Sep 17 00:00:00 2001 From: Daniel Kohler <11864045+ddkohler@users.noreply.github.com> Date: Wed, 10 Jul 2024 09:57:23 -0500 Subject: [PATCH 1/3] inherit_attrs --- WrightTools/data/_data.py | 11 +++++++---- 1 file changed, 7 insertions(+), 4 deletions(-) diff --git a/WrightTools/data/_data.py b/WrightTools/data/_data.py index 57ed6652..c6b2753c 100644 --- a/WrightTools/data/_data.py +++ b/WrightTools/data/_data.py @@ -1891,7 +1891,7 @@ def smooth(self, factors, channel=None, verbose=True) -> "Data": print("smoothed data") def split( - self, expression, positions, *, units=None, parent=None, verbose=True + self, expression, positions, *, units=None, parent=None, inherit_attrs=False, verbose=True, ) -> wt_collection.Collection: """ Split the data object along a given expression, in units. @@ -1928,7 +1928,7 @@ def split( # axis ------------------------------------------------------------------------------------ old_expr = self.axis_expressions old_units = self.units - out = wt_collection.Collection(name="split", parent=parent) + out = wt_collection.Collection(name=f"{self.name}_split", parent=parent) if isinstance(expression, int): if units is None: units = self._axes[expression].units @@ -1962,8 +1962,11 @@ def split( omasks.append(None) cuts.append(None) for i in range(len(positions) - 1): - out.create_data("split%03i" % i) - + out.create_data(f"{self.name}_{i:0>3}") + + if inherit_attrs: + for d in out.values(): + {d.attrs[k] : self.attrs[k] for k in self.attrs.keys() if k not in d.attrs.keys()} for var in self.variables: for i, (imask, omask, cut) in enumerate(zip(masks, omasks, cuts)): if omask is None: From 4b0dd6d3a934603434a5293b63da198d778e38d5 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Wed, 10 Jul 2024 15:04:03 +0000 Subject: [PATCH 2/3] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- WrightTools/data/_data.py | 13 ++++++++++--- 1 file changed, 10 insertions(+), 3 deletions(-) diff --git a/WrightTools/data/_data.py b/WrightTools/data/_data.py index c6b2753c..f15d12e5 100644 --- a/WrightTools/data/_data.py +++ b/WrightTools/data/_data.py @@ -1891,7 +1891,14 @@ def smooth(self, factors, channel=None, verbose=True) -> "Data": print("smoothed data") def split( - self, expression, positions, *, units=None, parent=None, inherit_attrs=False, verbose=True, + self, + expression, + positions, + *, + units=None, + parent=None, + inherit_attrs=False, + verbose=True, ) -> wt_collection.Collection: """ Split the data object along a given expression, in units. @@ -1963,10 +1970,10 @@ def split( cuts.append(None) for i in range(len(positions) - 1): out.create_data(f"{self.name}_{i:0>3}") - + if inherit_attrs: for d in out.values(): - {d.attrs[k] : self.attrs[k] for k in self.attrs.keys() if k not in d.attrs.keys()} + {d.attrs[k]: self.attrs[k] for k in self.attrs.keys() if k not in d.attrs.keys()} for var in self.variables: for i, (imask, omask, cut) in enumerate(zip(masks, omasks, cuts)): if omask is None: From c8ca9400e2bca331eb81eea156e2b90efee40e5a Mon Sep 17 00:00:00 2001 From: Daniel Kohler <11864045+ddkohler@users.noreply.github.com> Date: Wed, 10 Jul 2024 10:50:46 -0500 Subject: [PATCH 3/3] fix names --- WrightTools/data/_data.py | 11 ++++++----- tests/data/split.py | 2 +- 2 files changed, 7 insertions(+), 6 deletions(-) diff --git a/WrightTools/data/_data.py b/WrightTools/data/_data.py index f15d12e5..5a307e53 100644 --- a/WrightTools/data/_data.py +++ b/WrightTools/data/_data.py @@ -1935,7 +1935,7 @@ def split( # axis ------------------------------------------------------------------------------------ old_expr = self.axis_expressions old_units = self.units - out = wt_collection.Collection(name=f"{self.name}_split", parent=parent) + out = wt_collection.Collection(name=f"{self.natural_name}_split", parent=parent) if isinstance(expression, int): if units is None: units = self._axes[expression].units @@ -1969,11 +1969,8 @@ def split( omasks.append(None) cuts.append(None) for i in range(len(positions) - 1): - out.create_data(f"{self.name}_{i:0>3}") + out.create_data(f"{self.natural_name}_{i:0>3}") - if inherit_attrs: - for d in out.values(): - {d.attrs[k]: self.attrs[k] for k in self.attrs.keys() if k not in d.attrs.keys()} for var in self.variables: for i, (imask, omask, cut) in enumerate(zip(masks, omasks, cuts)): if omask is None: @@ -2047,6 +2044,10 @@ def split( for ax, u in zip(self.axes, old_units): ax.convert(u) + if inherit_attrs: + for d in out.values(): + {d.attrs[k]: self.attrs[k] for k in self.attrs.keys() if k not in d.attrs.keys()} + return out def transform(self, *axes, verbose=True): diff --git a/tests/data/split.py b/tests/data/split.py index 5b2e8b61..fd4fbb44 100755 --- a/tests/data/split.py +++ b/tests/data/split.py @@ -122,7 +122,7 @@ def test_split_parent(): a = wt.data.from_PyCMDS(p) parent = wt.Collection() split = a.split(1, [1500], parent=parent) - assert "split" in parent + assert f"{a.natural_name}_split" in parent assert split.filepath == parent.filepath assert len(split) == 2 a.close()