diff --git a/dgf/src/api/transform.py b/dgf/src/api/transform.py index 5a0d7c8..33e34d3 100644 --- a/dgf/src/api/transform.py +++ b/dgf/src/api/transform.py @@ -29,6 +29,7 @@ from dgf.src.transform.normalize import SoftQuantileNormalizer from dgf.src.transform.normalize import SinusoidTimedeltaNormalizer from dgf.src.transform.normalize import TimedeltaNormalizer +from dgf.src.transform.normalize import SequentialNormalizer from dgf.src.transform.extract import filter_graph from dgf.src.transform.extract import drop_edge_features diff --git a/dgf/src/transform/normalize.py b/dgf/src/transform/normalize.py index 8e160f5..28fa896 100644 --- a/dgf/src/transform/normalize.py +++ b/dgf/src/transform/normalize.py @@ -17,9 +17,11 @@ from __future__ import annotations import abc +import collections.abc import copy import dataclasses -from typing import Any, Dict, List, Optional, Set, Tuple +import inspect +from typing import Any, Dict, List, Optional, Sequence, Set, Tuple import dataclasses_json from dgf.src.data import in_memory_graph @@ -662,6 +664,138 @@ def normalize_tensorflow( return {self.output_feature_name: deltas} +@normalizer_registry.register +@dataclasses_json.dataclass_json +@dataclasses.dataclass(kw_only=True) +class SequentialNormalizer(AbstractFeatureNormalizer): + """Applies multiple AbstractFeatureNormalizers in sequence. + + Executes normalizer stages sequentially and emits only the final stage output. + """ + + stages: List[AbstractFeatureNormalizer] = normalizer_registry.field_list() + type: str = dataclasses.field( + default="SequentialNormalizer", init=False + ) + _stage_kwargs: List[frozenset[str]] = dataclasses.field( + default_factory=list, + init=False, + metadata=dataclasses_json.config(exclude=dataclasses_json.Exclude.ALWAYS), + ) + + def __post_init__(self): + _validate_stages(self.input_feature, self.stages) + self._stage_kwargs: List[frozenset[str]] = [ + _accepted_kwargs(s) for s in self.stages + ] + + @classmethod + def create( + cls, + stages: Sequence[AbstractFeatureNormalizer], + ) -> "SequentialNormalizer": + if not stages: + raise ValueError("SequentialNormalizer requires at least one stage.") + return SequentialNormalizer( + input_feature=stages[0].input_feature, + stages=list(stages), + ) + + def output_schema(self) -> schema_lib.FeatureSetSchema: + return self.stages[-1].output_schema() + + def _normalize( + self, + stage_fn_getter: collections.abc.Callable[ + [AbstractFeatureNormalizer], Any + ], + value: Any, + **kwargs: Any, + ) -> Dict[str, Any]: + current_features: Dict[str, Any] = {self.input_feature: value} + stage_output: Dict[str, Any] = {} + for stage, accepted_kwargs in zip(self.stages, self._stage_kwargs): + assert stage.input_feature in current_features, ( + f"Stage '{stage.type}' expects input feature '{stage.input_feature}'," + f" but available features are {list(current_features.keys())}." + ) + stage_input = current_features[stage.input_feature] + stage_output = _filter_and_call( + stage_fn_getter(stage), stage_input, accepted_kwargs, kwargs + ) + current_features.update(stage_output) + return stage_output + + def normalize_numpy( + self, + value: np.ndarray, + **kwargs: Any, + ) -> Dict[str, np.ndarray]: + return self._normalize(lambda s: s.normalize_numpy, value, **kwargs) + + def normalize_tensorflow( + self, + value: tf.Tensor, + **kwargs: Any, + ) -> Dict[str, tf.Tensor]: + return self._normalize(lambda s: s.normalize_tensorflow, value, **kwargs) + + def tensorflow_resources(self) -> List[tf.Tensor]: + resources = [] + for stage in self.stages: + resources.extend(stage.tensorflow_resources()) + return resources + + +def _validate_stages( + input_feature: str, + stages: Sequence[AbstractFeatureNormalizer], +) -> None: + """Validates that stages form a valid, connected sequential pipeline.""" + if not stages: + raise ValueError("SequentialNormalizer requires at least one stage.") + if input_feature != stages[0].input_feature: + raise ValueError( + f"SequentialNormalizer input_feature '{input_feature}' does not" + f" match first stage input_feature '{stages[0].input_feature}'." + ) + available = {input_feature} + for stage in stages: + if stage.input_feature not in available: + raise ValueError( + f"Stage '{stage.type}' expects input feature" + f" '{stage.input_feature}', but available features up to this stage" + f" are {sorted(available)}." + ) + available.update(stage.output_schema().keys()) + + +def _accepted_kwargs(stage: AbstractFeatureNormalizer) -> frozenset[str]: + """Returns the set of accepted keyword argument names.""" + if isinstance(stage, SequentialNormalizer): + return frozenset().union(*stage._stage_kwargs) + np_params = list(inspect.signature(stage.normalize_numpy).parameters)[1:] + tf_params = list(inspect.signature(stage.normalize_tensorflow).parameters)[1:] + if set(np_params) != set(tf_params): + raise ValueError( + f"Stage '{stage.type}' has mismatched kwargs between normalize_numpy " + f"({np_params}) and normalize_tensorflow ({tf_params})." + ) + return frozenset(np_params) + + +def _filter_and_call( + method: Any, + value: Any, + accepted_kwargs: collections.abc.Set[str], + kwargs: Dict[str, Any], +) -> Any: + if not accepted_kwargs: + return method(value) + filtered = {k: kwargs[k] for k in accepted_kwargs if k in kwargs} + return method(value, **filtered) + + @dataclasses_json.dataclass_json @dataclasses.dataclass class AutoNormalizeConfig: diff --git a/dgf/src/transform/normalize_test.py b/dgf/src/transform/normalize_test.py index 53ed1c9..9bc26d3 100644 --- a/dgf/src/transform/normalize_test.py +++ b/dgf/src/transform/normalize_test.py @@ -1038,6 +1038,193 @@ def test_timedelta_normalizer_dynamic_shape_raises(self): normalize_lib.TimedeltaNormalizer.create("time", schema) +class SequentialNormalizerTest(parameterized.TestCase): + + def test_sequential_normalizer_numpy_and_tensorflow(self): + ts_schema = schema_lib.FeatureSchema( + format=schema_lib.FeatureFormat.INTEGER_64, + semantic=schema_lib.FeatureSemantic.TIMESTAMP, + shape=(2,), + is_timeseries=True, + ) + stage1 = normalize_lib.TimedeltaNormalizer.create("time", ts_schema) + delta_schema = stage1.output_schema()["time_seed_delta"] + stage2 = normalize_lib.SinusoidTimedeltaNormalizer.create( + "time_seed_delta", delta_schema, embedding_dim=4 + ) + sequential = normalize_lib.SequentialNormalizer.create([stage1, stage2]) + + # Output schema should only contain final stage output. + out_schema = sequential.output_schema() + self.assertLen(out_schema, 1) + self.assertIn("time_seed_delta_SINUSOID", out_schema) + self.assertNotIn("time_seed_delta", out_schema) + self.assertEqual(out_schema["time_seed_delta_SINUSOID"].shape, (2, 4)) + + raw_np = np.array([[100, 200], [300, 400]], dtype=np.int64) + seeds_np = np.array([500, 1000], dtype=np.int64) + + # NumPy normalization. + out_np = sequential.normalize_numpy(raw_np, seed_timestamps=seeds_np) + self.assertEqual(list(out_np.keys()), ["time_seed_delta_SINUSOID"]) + self.assertEqual(out_np["time_seed_delta_SINUSOID"].shape, (2, 2, 4)) + + # TensorFlow normalization. + raw_tf = tf.constant(raw_np) + seeds_tf = tf.constant(seeds_np) + out_tf = sequential.normalize_tensorflow(raw_tf, seed_timestamps=seeds_tf) + self.assertEqual(list(out_tf.keys()), ["time_seed_delta_SINUSOID"]) + np.testing.assert_allclose( + out_tf["time_seed_delta_SINUSOID"].numpy(), + out_np["time_seed_delta_SINUSOID"], + rtol=1e-5, + ) + + def test_sequential_normalizer_json_serialization(self): + ts_schema = schema_lib.FeatureSchema( + format=schema_lib.FeatureFormat.INTEGER_64, + semantic=schema_lib.FeatureSemantic.TIMESTAMP, + shape=(), + ) + stage1 = normalize_lib.TimedeltaNormalizer.create("time", ts_schema) + delta_schema = stage1.output_schema()["time_seed_delta"] + stage2 = normalize_lib.SinusoidTimedeltaNormalizer.create( + "time_seed_delta", delta_schema, embedding_dim=4 + ) + sequential = normalize_lib.SequentialNormalizer.create([stage1, stage2]) + + json_str = sequential.to_json() # pyrefly: ignore[missing-attribute] + # pyrefly: ignore[missing-attribute] + reconstructed = normalize_lib.SequentialNormalizer.from_json(json_str) + self.assertLen(reconstructed.stages, 2) + self.assertEqual(reconstructed.input_feature, "time") + self.assertIn("time_seed_delta_SINUSOID", reconstructed.output_schema()) + # Verify execution on reconstructed normalizer. + out = reconstructed.normalize_numpy( + np.array([100], dtype=np.int64), + seed_timestamps=np.array([500], dtype=np.int64), + ) + self.assertIn("time_seed_delta_SINUSOID", out) + + def test_validate_stages(self): + ts_schema = schema_lib.FeatureSchema( + format=schema_lib.FeatureFormat.INTEGER_64, + semantic=schema_lib.FeatureSemantic.TIMESTAMP, + shape=(), + ) + stage1 = normalize_lib.TimedeltaNormalizer.create("time", ts_schema) + delta_schema = stage1.output_schema()["time_seed_delta"] + stage2 = normalize_lib.SinusoidTimedeltaNormalizer.create( + "time_seed_delta", delta_schema, embedding_dim=4 + ) + + # Valid stages. + normalize_lib._validate_stages("time", [stage1, stage2]) + + # Empty stages raises. + with self.assertRaisesRegex(ValueError, "at least one stage"): + normalize_lib._validate_stages("time", []) + + # Mismatched input feature raises. + with self.assertRaisesRegex(ValueError, "does not match first stage"): + normalize_lib._validate_stages("other_feature", [stage1, stage2]) + + # Missing intermediate feature raises. + unconnected_stage = normalize_lib.SinusoidTimedeltaNormalizer.create( + "non_existent_feature", delta_schema, embedding_dim=4 + ) + with self.assertRaisesRegex(ValueError, "expects input feature"): + normalize_lib._validate_stages("time", [stage1, unconnected_stage]) + + def test_sequential_normalizer_empty_stages_raises(self): + with self.assertRaisesRegex(ValueError, "at least one stage"): + normalize_lib.SequentialNormalizer.create([]) + + def test_sequential_normalizer_mismatched_input_feature_raises(self): + ts_schema = schema_lib.FeatureSchema( + format=schema_lib.FeatureFormat.INTEGER_64, + semantic=schema_lib.FeatureSemantic.TIMESTAMP, + shape=(), + ) + stage = normalize_lib.TimedeltaNormalizer.create("time", ts_schema) + with self.assertRaisesRegex(ValueError, "does not match first stage"): + normalize_lib.SequentialNormalizer( + input_feature="mismatched_feature", + stages=[stage], + ) + + def test_sequential_normalizer_missing_intermediate_feature_raises(self): + ts_schema = schema_lib.FeatureSchema( + format=schema_lib.FeatureFormat.INTEGER_64, + semantic=schema_lib.FeatureSemantic.TIMESTAMP, + shape=(), + ) + stage1 = normalize_lib.TimedeltaNormalizer.create("time", ts_schema) + delta_schema = schema_lib.FeatureSchema( + format=schema_lib.FeatureFormat.INTEGER_64, + semantic=schema_lib.FeatureSemantic.TIMEDELTA, + shape=(), + ) + # Stage 2 expects "non_existent_feature" which stage 1 does not output. + stage2 = normalize_lib.SinusoidTimedeltaNormalizer.create( + "non_existent_feature", delta_schema, embedding_dim=4 + ) + with self.assertRaisesRegex(ValueError, "expects input feature"): + normalize_lib.SequentialNormalizer.create([stage1, stage2]) + + def test_sequential_normalizer_nested(self): + ts_schema = schema_lib.FeatureSchema( + format=schema_lib.FeatureFormat.INTEGER_64, + semantic=schema_lib.FeatureSemantic.TIMESTAMP, + shape=(), + ) + stage1 = normalize_lib.TimedeltaNormalizer.create("time", ts_schema) + delta_schema = stage1.output_schema()["time_seed_delta"] + stage2 = normalize_lib.SinusoidTimedeltaNormalizer.create( + "time_seed_delta", delta_schema, embedding_dim=4 + ) + inner = normalize_lib.SequentialNormalizer.create([stage1]) + outer = normalize_lib.SequentialNormalizer.create([inner, stage2]) + out = outer.normalize_numpy( + np.array([100], dtype=np.int64), + seed_timestamps=np.array([500], dtype=np.int64), + ) + self.assertIn("time_seed_delta_SINUSOID", out) + + def test_sequential_normalizer_tensorflow_resources(self): + ts_schema = schema_lib.FeatureSchema( + format=schema_lib.FeatureFormat.INTEGER_64, + semantic=schema_lib.FeatureSemantic.TIMESTAMP, + shape=(), + ) + stage1 = normalize_lib.TimedeltaNormalizer.create("time", ts_schema) + delta_schema = stage1.output_schema()["time_seed_delta"] + stage2 = normalize_lib.SinusoidTimedeltaNormalizer.create( + "time_seed_delta", delta_schema, embedding_dim=4 + ) + sequential = normalize_lib.SequentialNormalizer.create([stage1, stage2]) + self.assertEmpty(sequential.tensorflow_resources()) + + def test_sequential_normalizer_extra_kwargs_ignored(self): + ts_schema = schema_lib.FeatureSchema( + format=schema_lib.FeatureFormat.INTEGER_64, + semantic=schema_lib.FeatureSemantic.TIMESTAMP, + shape=(), + ) + stage1 = normalize_lib.TimedeltaNormalizer.create("time", ts_schema) + delta_schema = stage1.output_schema()["time_seed_delta"] + stage2 = normalize_lib.SinusoidTimedeltaNormalizer.create( + "time_seed_delta", delta_schema, embedding_dim=4 + ) + sequential = normalize_lib.SequentialNormalizer.create([stage1, stage2]) + out = sequential.normalize_numpy( + np.array([100], dtype=np.int64), + seed_timestamps=np.array([500], dtype=np.int64), + unrelated_kwarg="ignored", + ) + self.assertIn("time_seed_delta_SINUSOID", out) + + if __name__ == "__main__": absltest.main()