From 36cb261fda459b3f2792287ecffd83401a2bf0d6 Mon Sep 17 00:00:00 2001 From: Simon Meierhans Date: Tue, 25 Aug 2026 00:15:45 -0700 Subject: [PATCH] Remove per sample transforms from SampleGeneratorFromAnything. PiperOrigin-RevId: 970351662 --- dgf/src/learning/ten_lines/BUILD | 1 - dgf/src/learning/ten_lines/dataset.py | 91 +-------- dgf/src/learning/ten_lines/dataset_test.py | 204 +-------------------- 3 files changed, 7 insertions(+), 289 deletions(-) diff --git a/dgf/src/learning/ten_lines/BUILD b/dgf/src/learning/ten_lines/BUILD index 045d02f..f8d2a0d 100644 --- a/dgf/src/learning/ten_lines/BUILD +++ b/dgf/src/learning/ten_lines/BUILD @@ -349,7 +349,6 @@ py_test( "//dgf/src/data:schema", "//dgf/src/io:tf_graph_sample", "//dgf/src/sampling:config", - "//dgf/src/transform:timeseries", "//dgf/src/util:gen_test_graph", "//dgf/src/validate:in_memory_graph", # numpy dep, diff --git a/dgf/src/learning/ten_lines/dataset.py b/dgf/src/learning/ten_lines/dataset.py index 38a45d2..199f2d4 100644 --- a/dgf/src/learning/ten_lines/dataset.py +++ b/dgf/src/learning/ten_lines/dataset.py @@ -118,10 +118,6 @@ class SampleGeneratorFromAnything: sampler_returns_node_idxs_only: If `True`, the sampler returns only node indices without feature values. If `False` (default), the sampler returns feature values. - timeseries_pad_and_cap: Configuration for padding and capping timeseries - sequence features per sample before merging. - timedelta_extraction: Configuration for time delta extraction per sample - before merging. """ graph: Graph @@ -144,12 +140,6 @@ class SampleGeneratorFromAnything: default_factory=dict ) sampler_returns_node_idxs_only: bool = False - timeseries_pad_and_cap: Optional[ - timeseries_transform.PadAndCapTimeseriesConfig - ] = None - timedelta_extraction: Optional[ - timeseries_transform.TimestampFeatureExtractorConfig - ] = None num_seed_nodes: Optional[int] = dataclasses.field(init=False) batch_iterator: BatchSampleGeneratorIteratorFn = dataclasses.field(init=False) @@ -163,15 +153,6 @@ class SampleGeneratorFromAnything: _seed_timestamps_all: Optional[np.ndarray] = dataclasses.field( init=False, default=None ) - _pad_and_cap_transformer: Optional[ - timeseries_transform.PadAndCapTimeseries - ] = dataclasses.field(init=False, default=None) - _timestamp_extractor: Optional[ - timeseries_transform.TimestampFeatureExtractor - ] = dataclasses.field(init=False, default=None) - _has_per_sample_transforms: bool = dataclasses.field( - init=False, default=False - ) def __post_init__(self): if self.seed_node_idxs is not None: @@ -231,60 +212,14 @@ def __post_init__(self): len(self.seed_node_idxs) if self.seed_node_idxs is not None else None ) - # Initializing per-sample transforms and precomputing the output schema. - current_schema = self.schema - self._has_per_sample_transforms = ( - self.timeseries_pad_and_cap is not None - or self.timedelta_extraction is not None - ) - - if ( - self.temporal or self._has_per_sample_transforms - ) and self.format != GraphFormat.IN_MEMORY_GRAPH: + if self.temporal and self.format != GraphFormat.IN_MEMORY_GRAPH: raise ValueError( - "Temporal sampling (`temporal=True`) and per-sample transformations" - " (`timeseries_pad_and_cap`, `timedelta_extraction`) are" + "Temporal sampling (`temporal=True`) is" " only supported for GraphFormat.IN_MEMORY_GRAPH, but got format:" f" {self.format}" ) - if ( - self.timeseries_pad_and_cap is not None - and temporal_util.schema_has_timeseries_features(current_schema) - ): - self._pad_and_cap_transformer = timeseries_transform.PadAndCapTimeseries( - current_schema, self.timeseries_pad_and_cap - ) - current_schema = self._pad_and_cap_transformer.output_schema() - - if temporal_util.schema_has_dynamic_timeseries_features(current_schema): - raise ValueError( - "Dynamic shape timeseries features were detected in the schema;" - " please configure `timeseries_pad_and_cap` to pad/cap the sequences" - " first." - ) - - if self.timedelta_extraction is not None: - if self._ts_feature is None: - raise ValueError( - "Timedelta extraction is configured, but no creation timestamp" - f" feature was found in node set '{self._target_nodeset}'." - ) - self._timestamp_extractor = ( - timeseries_transform.TimestampFeatureExtractor( - current_schema, self.timedelta_extraction - ) - ) - current_schema = self._timestamp_extractor.output_schema() - - self._output_schema = current_schema - - if self.sampler_returns_node_idxs_only and self._has_per_sample_transforms: - raise ValueError( - "Per-sample transformations (`timeseries_pad_and_cap`," - " `timedelta_extraction`) cannot be used when" - " `sampler_returns_node_idxs_only=True`." - ) + self._output_schema = self.schema self.batch_iterator, self.single_iterator = self.iterator_builder() @@ -296,12 +231,6 @@ def set_sampler_returns_node_idxs_only( self, sampler_returns_node_idxs_only: bool ): """Changes whether the sampler returns only node indices.""" - if sampler_returns_node_idxs_only and self._has_per_sample_transforms: - raise ValueError( - "Per-sample transformations (`timeseries_pad_and_cap`," - " `timedelta_extraction`) cannot be used when" - " `sampler_returns_node_idxs_only=True`." - ) self.sampler_returns_node_idxs_only = sampler_returns_node_idxs_only if self.in_memory_sampler is not None: self.in_memory_sampler.set_return_options( @@ -380,20 +309,6 @@ def batch_generator(): else: graph_samples = self.in_memory_sampler.sample(node_idxs) - # Check if per-sample transforms are configured for efficiency. - if self._has_per_sample_transforms: - transformed_samples = [] - for i, sample in enumerate(graph_samples): - curr_sg = sample - if self._pad_and_cap_transformer is not None: - curr_sg = self._pad_and_cap_transformer(curr_sg) - if self._timestamp_extractor is not None: - assert seed_timestamps is not None - st = int(seed_timestamps[i]) - curr_sg = self._timestamp_extractor(curr_sg, seed_timestamp=st) - transformed_samples.append(curr_sg) - graph_samples = transformed_samples - try: yield merge_lib.merge_graphs( graph_samples, merge_schema, padding=self.padding diff --git a/dgf/src/learning/ten_lines/dataset_test.py b/dgf/src/learning/ten_lines/dataset_test.py index 452dd95..58923b6 100644 --- a/dgf/src/learning/ten_lines/dataset_test.py +++ b/dgf/src/learning/ten_lines/dataset_test.py @@ -19,7 +19,6 @@ from dgf.src.io import tf_graph_sample from dgf.src.learning.ten_lines import dataset from dgf.src.sampling import config as sampling_config_lib -from dgf.src.transform import timeseries as timeseries_transform from dgf.src.util import gen_test_graph from dgf.src.validate import in_memory_graph as in_memory_graph_validate_lib import numpy as np @@ -259,13 +258,9 @@ def _create_temporal_test_graph_and_schema(self): ) return graph, schema - def test_per_sample_transformations(self): + def test_temporal_sampling(self): graph, schema = self._create_temporal_test_graph_and_schema() - pad_and_cap_config = timeseries_transform.PadAndCapTimeseriesConfig( - sequence_length=5 - ) - generator = dataset.SampleGeneratorFromAnything( graph=graph, schema=schema, @@ -276,99 +271,25 @@ def test_per_sample_transformations(self): num_hops=1, hop_width=2, temporal_sampling=True, - max_timeseries_len=5, + max_timeseries_len=3, ), temporal=True, - timeseries_pad_and_cap=pad_and_cap_config, - timedelta_extraction=timeseries_transform.TimestampFeatureExtractorConfig(), drop_remainder=False, shuffle=False, ) - self.assertIn( + self.assertNotIn( "time_mask", generator.output_schema().node_sets["hardware"].features ) - self.assertIn( - "creation_time_seed_delta", - generator.output_schema().node_sets["alerts"].features, - ) - self.assertIn( - "time_seed_delta", - generator.output_schema().node_sets["hardware"].features, - ) - num_batches = 0 for sample, _ in generator.batch_iterator(): in_memory_graph_validate_lib.validate_graph( sample, generator.output_schema(), raise_on_warning=False ) - self.assertEqual( - sample.node_sets["hardware"].features["signal"].shape[1], 5 - ) - self.assertEqual( - sample.node_sets["hardware"].features["time_mask"].shape[1], 5 - ) - self.assertIn("time_seed_delta", sample.node_sets["hardware"].features) - self.assertIn( - "creation_time_seed_delta", sample.node_sets["alerts"].features - ) - self.assertFalse( - np.all(sample.node_sets["hardware"].features["time_mask"] == 0) - ) num_batches += 1 self.assertEqual(num_batches, 2) - def test_default_no_per_sample_transforms(self): - graph, schema = self._create_temporal_test_graph_and_schema() - schema.node_sets["hardware"].features["time"] = schema_lib.FeatureSchema( - format=schema_lib.FeatureFormat.INTEGER_64, - semantic=schema_lib.FeatureSemantic.TIMESTAMP, - is_timeseries=True, - is_creation_time=True, - shape=(3,), - ) - schema.node_sets["hardware"].features["signal"] = schema_lib.FeatureSchema( - format=schema_lib.FeatureFormat.FLOAT_32, - semantic=schema_lib.FeatureSemantic.NUMERICAL, - is_timeseries=True, - shape=(3,), - ) - graph.node_sets["hardware"].features["time"] = np.array( - [[50, 80, 120], [150, 250, 300]], dtype=np.int64 - ) - graph.node_sets["hardware"].features["signal"] = np.array( - [[1.5, 2.5, 3.5], [4.5, 5.5, 6.5]], dtype=np.float32 - ) - - generator = dataset.SampleGeneratorFromAnything( - graph=graph, - schema=schema, - batch_size=2, - seed_node_idxs=None, - sampling_config=sampling_config_lib.SimpleSamplingConfig( - seed_nodeset="alerts", - num_hops=1, - hop_width=2, - temporal_sampling=True, - max_timeseries_len=3, - ), - temporal=True, - drop_remainder=False, - shuffle=False, - ) - - self.assertNotIn( - "time_mask", generator.output_schema().node_sets["hardware"].features - ) - self.assertNotIn( - "creation_time_seed_delta", - generator.output_schema().node_sets["alerts"].features, - ) - self.assertNotIn( - "time_seed_delta", - generator.output_schema().node_sets["hardware"].features, - ) def test_temporal_sampling_requires_seed_timestamps(self): graph = gen_test_graph.generate_in_memory_graph(True, False) @@ -391,31 +312,7 @@ def test_temporal_sampling_requires_seed_timestamps(self): shuffle=False, ) - def test_sampler_returns_node_idxs_only_with_transforms_raises(self): - graph, schema = self._create_temporal_test_graph_and_schema() - with self.assertRaisesRegex( - ValueError, "cannot be used when `sampler_returns_node_idxs_only=True`" - ): - dataset.SampleGeneratorFromAnything( - graph=graph, - schema=schema, - batch_size=2, - seed_node_idxs=None, - sampling_config=sampling_config_lib.SimpleSamplingConfig( - seed_nodeset="alerts", - num_hops=1, - hop_width=2, - temporal_sampling=True, - max_timeseries_len=5, - ), - temporal=True, - timeseries_pad_and_cap=timeseries_transform.PadAndCapTimeseriesConfig(), - sampler_returns_node_idxs_only=True, - drop_remainder=False, - shuffle=False, - ) - - def test_non_in_memory_format_with_transforms_raises(self): + def test_non_in_memory_format_temporal_raises(self): _, schema = self._create_temporal_test_graph_and_schema() with self.assertRaisesRegex( ValueError, @@ -437,99 +334,6 @@ def test_non_in_memory_format_with_transforms_raises(self): shuffle=False, ) - def test_timedelta_extraction_with_temporal_false(self): - graph, schema = self._create_temporal_test_graph_and_schema() - generator = dataset.SampleGeneratorFromAnything( - graph=graph, - schema=schema, - batch_size=2, - seed_node_idxs=np.array([0, 1, 2, 3], dtype=np.int64), - sampling_config=sampling_config_lib.SimpleSamplingConfig( - seed_nodeset="alerts", - num_hops=1, - hop_width=2, - max_timeseries_len=5, - ), - timeseries_pad_and_cap=timeseries_transform.PadAndCapTimeseriesConfig( - sequence_length=5 - ), - timedelta_extraction=timeseries_transform.TimestampFeatureExtractorConfig(), - temporal=False, - drop_remainder=False, - shuffle=False, - ) - - self.assertIn( - "creation_time_seed_delta", - generator.output_schema().node_sets["alerts"].features, - ) - self.assertIn( - "time_seed_delta", - generator.output_schema().node_sets["hardware"].features, - ) - for sample, _ in generator.batch_iterator(): - self.assertIn("time_seed_delta", sample.node_sets["hardware"].features) - self.assertIn( - "creation_time_seed_delta", sample.node_sets["alerts"].features - ) - break - - def test_dynamic_set_sampler_returns_node_idxs_only_raises(self): - graph, schema = self._create_temporal_test_graph_and_schema() - generator = dataset.SampleGeneratorFromAnything( - graph=graph, - schema=schema, - batch_size=2, - seed_node_idxs=None, - sampling_config=sampling_config_lib.SimpleSamplingConfig( - seed_nodeset="alerts", - num_hops=1, - hop_width=2, - temporal_sampling=True, - max_timeseries_len=5, - ), - temporal=True, - timeseries_pad_and_cap=timeseries_transform.PadAndCapTimeseriesConfig(), - drop_remainder=False, - shuffle=False, - ) - with self.assertRaisesRegex( - ValueError, "cannot be used when `sampler_returns_node_idxs_only=True`" - ): - generator.set_sampler_returns_node_idxs_only(True) - - def test_timedelta_extraction_without_pad_and_cap_dynamic_ts_raises(self): - graph, schema = self._create_temporal_test_graph_and_schema() - # Add a dynamic shape timeseries feature - schema.node_sets["hardware"].features["ts_dynamic"] = ( - schema_lib.FeatureSchema( - format=schema_lib.FeatureFormat.FLOAT_32, - shape=(None,), - is_timeseries=True, - ) - ) - with self.assertRaisesRegex( - ValueError, - "Dynamic shape timeseries features were detected in the schema", - ): - dataset.SampleGeneratorFromAnything( - graph=graph, - schema=schema, - batch_size=2, - seed_node_idxs=None, - sampling_config=sampling_config_lib.SimpleSamplingConfig( - seed_nodeset="alerts", - num_hops=1, - hop_width=2, - temporal_sampling=True, - max_timeseries_len=5, - ), - temporal=True, - timedelta_extraction=timeseries_transform.TimestampFeatureExtractorConfig(), - drop_remainder=False, - shuffle=False, - ) - if __name__ == "__main__": absltest.main()