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: 0 additions & 1 deletion dgf/src/api/transform.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
1 change: 1 addition & 0 deletions dgf/src/transform/BUILD
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
33 changes: 1 addition & 32 deletions dgf/src/transform/extract.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
23 changes: 7 additions & 16 deletions dgf/src/transform/extract_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"]
Expand Down