From 0aca87d261c566c5e1699ddf80e531483b4b199c Mon Sep 17 00:00:00 2001 From: Simon Meierhans Date: Thu, 20 Aug 2026 21:36:36 -0700 Subject: [PATCH] Remove duplicate `filter_schema` function. PiperOrigin-RevId: 968246086 --- dgf/src/api/transform.py | 1 - dgf/src/transform/BUILD | 1 + dgf/src/transform/extract.py | 33 +------------------------------ dgf/src/transform/extract_test.py | 23 +++++++-------------- 4 files changed, 9 insertions(+), 49 deletions(-) diff --git a/dgf/src/api/transform.py b/dgf/src/api/transform.py index 90987e8..83fa026 100644 --- a/dgf/src/api/transform.py +++ b/dgf/src/api/transform.py @@ -29,7 +29,6 @@ from dgf.src.transform.normalize import SoftQuantileNormalizer from dgf.src.transform.normalize import SinusoidTimedeltaNormalizer -from dgf.src.transform.extract import filter_schema from dgf.src.transform.extract import filter_graph from dgf.src.transform.extract import drop_edge_features diff --git a/dgf/src/transform/BUILD b/dgf/src/transform/BUILD index 7d51812..d547c2f 100644 --- a/dgf/src/transform/BUILD +++ b/dgf/src/transform/BUILD @@ -185,6 +185,7 @@ py_test( srcs = ["extract_test.py"], deps = [ ":extract", + ":schema", # absl/testing:absltest dep, "//dgf/src/data:in_memory_graph", "//dgf/src/data:schema", diff --git a/dgf/src/transform/extract.py b/dgf/src/transform/extract.py index 066221d..f676912 100644 --- a/dgf/src/transform/extract.py +++ b/dgf/src/transform/extract.py @@ -15,43 +15,12 @@ """Operations on a schema.""" import copy -from typing import List, Tuple +from typing import Tuple from dgf.src.data import in_memory_graph as in_memory_graph_lib from dgf.src.data import schema as schema_lib from dgf.src.transform import schema as schema_transform_lib -def filter_schema( - src: schema_lib.GraphSchema, selected_features: List[str] -) -> schema_lib.GraphSchema: - """Creates a new schema with a subset of the features. - - The other parts of the schema are not modified. - - Args: - src: The source schema to extract. - selected_features: A list of feature names to include in the new schema. - - Returns: - The extracted schema. - """ - extracted_schema = copy.deepcopy(src) - - for node_name in extracted_schema.node_sets: - node_schema = extracted_schema.node_sets[node_name] - node_schema.features = { - k: v for k, v in node_schema.features.items() if k in selected_features - } - - for edge_name in extracted_schema.edge_sets: - edge_schema = extracted_schema.edge_sets[edge_name] - edge_schema.features = { - k: v for k, v in edge_schema.features.items() if k in selected_features - } - - return extracted_schema - - def filter_graph( graph: in_memory_graph_lib.InMemoryGraph, schema: schema_lib.GraphSchema, diff --git a/dgf/src/transform/extract_test.py b/dgf/src/transform/extract_test.py index b93a01e..ea24982 100644 --- a/dgf/src/transform/extract_test.py +++ b/dgf/src/transform/extract_test.py @@ -19,31 +19,22 @@ from dgf.src.data import in_memory_graph as in_memory_graph_lib from dgf.src.data import schema as schema_lib from dgf.src.transform import extract as extract_lib +from dgf.src.transform import schema as schema_transform_lib from dgf.src.util import gen_test_graph from dgf.src.util import test_util class FilterTest(absltest.TestCase): - def test_filter_schema(self): - schema = gen_test_graph.generate_schema() - selected_features = ["f1", "f3"] - extracted_schema = extract_lib.filter_schema(schema, selected_features) - expected_extracted_schema = copy.deepcopy(schema) - del expected_extracted_schema.node_sets["n1"].features["f2"] - del expected_extracted_schema.node_sets["n2"].features["f4"] - del expected_extracted_schema.node_sets["n2"].features["f5"] - del expected_extracted_schema.node_sets["n2"].features["f6"] - test_util.assert_are_equal( - self, - extracted_schema, - expected_extracted_schema, - ) - def test_filter_graph(self): graph = gen_test_graph.generate_in_memory_graph(variable_length=False) full_schema = gen_test_graph.generate_schema(variable_length=False) - extracted_schema = extract_lib.filter_schema(full_schema, ["f1", "f3"]) + extracted_schema = schema_transform_lib.filter_schema( + full_schema, + schema_lib.GraphSchemaFilter( + feature_fn=lambda name, _: name in ["f1", "f3"] + ), + ) extracted_graph = extract_lib.filter_graph(graph, extracted_schema) expected_extracted_graph = copy.deepcopy(graph) del expected_extracted_graph.node_sets["n1"].features["f2"]