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
1 change: 1 addition & 0 deletions dgf/src/api/transform.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
136 changes: 135 additions & 1 deletion dgf/src/transform/normalize.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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:
Expand Down
187 changes: 187 additions & 0 deletions dgf/src/transform/normalize_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()