diff --git a/dgf/src/analyse/padding.py b/dgf/src/analyse/padding.py index ae49ca8..1d3f950 100644 --- a/dgf/src/analyse/padding.py +++ b/dgf/src/analyse/padding.py @@ -13,7 +13,7 @@ # limitations under the License. import math -from typing import Iterator, List, Optional +from typing import Iterator, Optional from dgf.src.data import in_memory_graph as in_memory_graph_lib from dgf.src.data import padding as padding_lib from dgf.src.data import schema as schema_lib @@ -38,11 +38,10 @@ def padding_from_graph_generator( padding = padding_from_graph_generator(schema, graphs) # Later, the padding can be used to merge graphs - merged_graph_samples, nodeset_offsets = dgf.transform.merge_graphs( - [], + merged_graph_samples, nodeset_offsets = dgf.transform.GraphMerger( schema, padding=padding, - ) + )([]) ``` The padding size is: diff --git a/dgf/src/api/transform.py b/dgf/src/api/transform.py index 83fa026..5d16b98 100644 --- a/dgf/src/api/transform.py +++ b/dgf/src/api/transform.py @@ -17,7 +17,7 @@ # pylint: disable=unused-import,g-importing-member,g-import-not-at-top,g-bad-import-order,reimported,disable=attribute-error -from dgf.src.transform.merge import merge_graphs +from dgf.src.transform.merge import GraphMerger from dgf.src.transform.merge import remove_padding_sentinels from dgf.src.transform.normalize import GraphNormalizer diff --git a/dgf/src/learning/ten_lines/BUILD b/dgf/src/learning/ten_lines/BUILD index 8c361e3..9fd40ed 100644 --- a/dgf/src/learning/ten_lines/BUILD +++ b/dgf/src/learning/ten_lines/BUILD @@ -302,7 +302,7 @@ py_test( timeout = "long", srcs = ["node_prediction_test.py"], data = ["//test_data"], - shard_count = 4, + shard_count = 8, tags = [ #"requires-gpu-nvidia", ], diff --git a/dgf/src/learning/ten_lines/dataset.py b/dgf/src/learning/ten_lines/dataset.py index 622f7d8..12a80f3 100644 --- a/dgf/src/learning/ten_lines/dataset.py +++ b/dgf/src/learning/ten_lines/dataset.py @@ -286,9 +286,11 @@ def _generator_from_in_memory_graph( if self.sampler_returns_node_idxs_only else self.output_schema() ) - def batch_generator(): assert self.in_memory_sampler is not None + graph_merger = merge_lib.GraphMerger( + schema=merge_schema, padding=self.padding + ) for node_idxs in util.batch_indices_generator( self.seed_node_idxs # pyrefly: ignore[bad-argumet-type] if self.seed_node_idxs is not None @@ -309,9 +311,7 @@ def batch_generator(): graph_samples = self.in_memory_sampler.sample(node_idxs) try: - yield merge_lib.merge_graphs( - graph_samples, merge_schema, padding=self.padding - ) + yield graph_merger(graph_samples) except merge_lib.InsufficientPaddingError as e: if not self.skip_overflow_padding_error: raise e @@ -352,6 +352,9 @@ def batch_generator(): it = tf_graph_sample.read_tfgnn_graphs( self.graph, self.schema, container_type=container_type ) + graph_merger = merge_lib.GraphMerger( + schema=self.schema, padding=self.padding + ) while True: batch = [] try: @@ -360,15 +363,13 @@ def batch_generator(): except StopIteration: if batch and not self.drop_remainder: try: - yield merge_lib.merge_graphs( - batch, self.schema, padding=self.padding - ) + yield graph_merger(batch) except merge_lib.InsufficientPaddingError as e: if not self.skip_overflow_padding_error: raise e return try: - yield merge_lib.merge_graphs(batch, self.schema, padding=self.padding) + yield graph_merger(batch) except merge_lib.InsufficientPaddingError as e: if not self.skip_overflow_padding_error: raise e diff --git a/dgf/src/learning/ten_lines/link_prediction_dataset.py b/dgf/src/learning/ten_lines/link_prediction_dataset.py index 2a53fdc..b2aabd8 100644 --- a/dgf/src/learning/ten_lines/link_prediction_dataset.py +++ b/dgf/src/learning/ten_lines/link_prediction_dataset.py @@ -23,7 +23,7 @@ import copy import dataclasses import itertools -from typing import Dict, Iterator, Literal, Optional, Tuple, Union +from typing import Dict, Iterator, Literal, Optional from dgf.src.analyse import in_process_feature_statistics as in_process_feature_statistics_lib from dgf.src.analyse import padding as padding_lib from dgf.src.data import in_memory_graph as in_memory_graph_lib @@ -129,7 +129,7 @@ class GNNLinkDatasetPreparator: `dgf.analyse.feature_statistics_from_graphs` and `dgf.transform.AutoNormalizer`. - Padding graphs using `dgf.analyse.padding_from_graph_generator`. - - Merging of batches of graphs using `dgf.transform.merge_graphs`. + - Merging of batches of graphs using `dgf.transform.GraphMerger`. Attributes: graph: One of the graph format defined in data.Graph e.g. in-memory graph, @@ -609,6 +609,10 @@ def gen_raw_target_samples() -> Iterator[in_memory_graph_lib.InMemoryGraph]: # processed independently for padding, requiring separate padding # calculations. + graph_merger = merge_lib.GraphMerger( + schema=sampling_schema, padding=None, sentinel_offset=True + ) + def gen_normalized_merged_positive_source_samples() -> ( Iterator[in_memory_graph_lib.InMemoryGraph] ): @@ -624,9 +628,7 @@ def gen_normalized_merged_positive_source_samples() -> ( masked_edge_idxs=masked_edge_idxs, seed_timestamps=batch_seed.seed_timestamps, ) - merged_samples, _ = merge_lib.merge_graphs( - samples, sampling_schema, padding=None, sentinel_offset=True - ) + merged_samples, _ = graph_merger(samples) yield source_normalizer.normalize_numpy(merged_samples) def gen_normalized_merged_positive_target_samples() -> ( @@ -644,9 +646,7 @@ def gen_normalized_merged_positive_target_samples() -> ( masked_edge_idxs=masked_edge_idxs, seed_timestamps=batch_seed.seed_timestamps, ) - merged_samples, _ = merge_lib.merge_graphs( - samples, sampling_schema, padding=None, sentinel_offset=True - ) + merged_samples, _ = graph_merger(samples) yield target_normalizer.normalize_numpy(merged_samples) def gen_normalized_merged_negative_target_samples() -> ( @@ -669,9 +669,7 @@ def gen_normalized_merged_negative_target_samples() -> ( masked_edge_idxs=neg_masked_edge_idxs, seed_timestamps=neg_seed_timestamps, ) - merged_samples, _ = merge_lib.merge_graphs( - samples, sampling_schema, padding=None, sentinel_offset=True - ) + merged_samples, _ = graph_merger(samples) yield target_normalizer.normalize_numpy(merged_samples) gen_normalized_merged_positive_source_samples_iter = ( @@ -817,12 +815,11 @@ def _sample_and_merge( masked_edge_idxs=masked_edge_idxs, seed_timestamps=batch_seed.seed_timestamps, ) - pos_src_merged, pos_src_offsets = merge_lib.merge_graphs( - pos_src_samples, - merge_schema, + pos_src_merged, pos_src_offsets = merge_lib.GraphMerger( + schema=merge_schema, padding=live.positive_source_padding if padding else None, sentinel_offset=True, - ) + )(pos_src_samples) # Positive target pos_trg_samples = live.target_sampler.sample( @@ -830,12 +827,11 @@ def _sample_and_merge( masked_edge_idxs=masked_edge_idxs, seed_timestamps=batch_seed.seed_timestamps, ) - pos_trg_merged, pos_trg_offsets = merge_lib.merge_graphs( - pos_trg_samples, - merge_schema, + pos_trg_merged, pos_trg_offsets = merge_lib.GraphMerger( + schema=merge_schema, padding=live.positive_target_padding if padding else None, sentinel_offset=True, - ) + )(pos_trg_samples) # Negative target neg_seed_timestamps = ( @@ -853,12 +849,11 @@ def _sample_and_merge( masked_edge_idxs=neg_masked_edge_idxs, seed_timestamps=neg_seed_timestamps, ) - neg_trg_merged, neg_trg_offsets = merge_lib.merge_graphs( - neg_trg_samples, - merge_schema, + neg_trg_merged, neg_trg_offsets = merge_lib.GraphMerger( + schema=merge_schema, padding=live.negative_target_padding if padding else None, sentinel_offset=True, - ) + )(neg_trg_samples) return GNNLinkDatasetPreparatorSample( positive_source_graph=pos_src_merged, diff --git a/dgf/src/learning/ten_lines/link_prediction_model.py b/dgf/src/learning/ten_lines/link_prediction_model.py index 6ceb7d3..b5608c4 100644 --- a/dgf/src/learning/ten_lines/link_prediction_model.py +++ b/dgf/src/learning/ten_lines/link_prediction_model.py @@ -18,7 +18,7 @@ import dataclasses import os -from typing import Any, Callable, Dict, Iterator, List, Literal, Optional, Tuple, Union +from typing import Any, Callable, Dict, Iterator, List, Literal, Optional, Tuple import dataclasses_json from dgf.src.data import in_memory_graph @@ -601,6 +601,17 @@ def predict_batch( ), ) + source_graph_merger = merge_lib.GraphMerger( + schema=self._data.schema, + padding=self._data.positive_source_padding, + sentinel_offset=False, + ) + target_graph_merger = merge_lib.GraphMerger( + schema=self._data.schema, + padding=self._data.positive_target_padding, + sentinel_offset=False, + ) + def merge_and_predict( sub_src: np.ndarray, sub_trg: np.ndarray, @@ -608,18 +619,8 @@ def merge_and_predict( sub_trg_samples: List[in_memory_graph.InMemoryGraph], ): - source_merged, source_offsets = merge_lib.merge_graphs( - sub_src_samples, - self._data.schema, - padding=self._data.positive_source_padding, - sentinel_offset=False, - ) - target_merged, target_offsets = merge_lib.merge_graphs( - sub_trg_samples, - self._data.schema, - padding=self._data.positive_target_padding, - sentinel_offset=False, - ) + source_merged, source_offsets = source_graph_merger(sub_src_samples) + target_merged, target_offsets = target_graph_merger(sub_trg_samples) source_normalized = live.source_normalizer.normalize_numpy(source_merged) target_normalized = live.target_normalizer.normalize_numpy(target_merged) @@ -797,14 +798,15 @@ def predict_embedding( ), ) + graph_merger = merge_lib.GraphMerger( + schema=self._data.schema, + padding=padding, + sentinel_offset=False, + ) + def merge_and_predict_emb(sub_samples: List[in_memory_graph.InMemoryGraph]): - merged, offsets = merge_lib.merge_graphs( - sub_samples, - self._data.schema, - padding=padding, - sentinel_offset=False, - ) + merged, offsets = graph_merger(sub_samples) normalized = normalizer.normalize_numpy(merged) jax_graph = jax_lib.graph_to_jax_graph(normalized) diff --git a/dgf/src/learning/ten_lines/link_prediction_test.py b/dgf/src/learning/ten_lines/link_prediction_test.py index c7a78db..b0a4006 100644 --- a/dgf/src/learning/ten_lines/link_prediction_test.py +++ b/dgf/src/learning/ten_lines/link_prediction_test.py @@ -523,22 +523,25 @@ def test_predict_predict_batch(self): def test_predict_batch_insufficient_padding(self): """Tests that predict_batch handles InsufficientPaddingError by splitting.""" - original_merge_graphs = link_prediction_model.merge_lib.merge_graphs + original_graph_merger = link_prediction_model.merge_lib.GraphMerger call_count = 0 - def mock_merge_graphs(*args, **kwargs): - nonlocal call_count - call_count += 1 - if call_count == 1: - raise link_prediction_model.merge_lib.InsufficientPaddingError( - "Simulated insufficient padding" - ) - return original_merge_graphs(*args, **kwargs) + def mock_graph_merger(*args, **kwargs): + real_graph_merger = original_graph_merger(*args, **kwargs) + def graph_merger_wrapper(*call_args, **call_kwargs): + nonlocal call_count + call_count += 1 + if call_count == 1: + raise link_prediction_model.merge_lib.InsufficientPaddingError( + "Simulated insufficient padding" + ) + return real_graph_merger(*call_args, **call_kwargs) + return graph_merger_wrapper with unittest.mock.patch.object( link_prediction_model.merge_lib, - "merge_graphs", - side_effect=mock_merge_graphs, + "GraphMerger", + side_effect=mock_graph_merger, ): probs = self.model.predict( self.graph, diff --git a/dgf/src/learning/ten_lines/node_prediction_dataset.py b/dgf/src/learning/ten_lines/node_prediction_dataset.py index 5eb0ac6..13610f5 100644 --- a/dgf/src/learning/ten_lines/node_prediction_dataset.py +++ b/dgf/src/learning/ten_lines/node_prediction_dataset.py @@ -68,7 +68,7 @@ class GNNDatasetPreparator: `dgf.analyse.feature_statistics_from_graphs` and `dgf.transform.AutoNormalizer`. - Padding graphs using `dgf.analyse.padding_from_graph_generator`. - - Merging of batches of graphs using `dgf.transform.merge_graphs`. + - Merging of batches of graphs using `dgf.transform.GraphMerger`. This class is intended for basic GNN pipelines. For advanced GNN pipelines, users should apply those transformations manually. diff --git a/dgf/src/learning/ten_lines/node_prediction_model.py b/dgf/src/learning/ten_lines/node_prediction_model.py index 7d535e4..5265326 100644 --- a/dgf/src/learning/ten_lines/node_prediction_model.py +++ b/dgf/src/learning/ten_lines/node_prediction_model.py @@ -391,6 +391,12 @@ def predict_batch( ), ) + graph_merger = merge_lib.GraphMerger( + schema=schema, + padding=self._data.padding, + sentinel_offset=False, + ) + for batch_seed_node_idxs in batch_seed_node_idxs_generator: # TODO(gbm): The sampler should consume np array directly. @@ -410,23 +416,21 @@ def predict_batch( graph_samples = sampler.sample(batch_seed_node_idxs) yield from self._predict_sub_batch( - live, schema, batch_seed_node_idxs, graph_samples + live, + batch_seed_node_idxs, + graph_samples, + graph_merger=graph_merger, ) def _predict_sub_batch( self, live, - schema: schema_lib.GraphSchema, sub_seed_idxs: np.ndarray, sub_samples: List[in_memory_graph.InMemoryGraph], + graph_merger: merge_lib.GraphMerger, ) -> Iterator[BatchPrediction]: try: - merged_graph, merge_offsets = merge_lib.merge_graphs( - sub_samples, - schema, - padding=self._data.padding, - sentinel_offset=False, - ) + merged_graph, merge_offsets = graph_merger(sub_samples) except merge_lib.InsufficientPaddingError: # The graph is too large to fit in the padding. Let's split it in two # and try again. @@ -434,10 +438,16 @@ def _predict_sub_batch( raise mid = len(sub_samples) // 2 yield from self._predict_sub_batch( - live, schema, sub_seed_idxs[:mid], sub_samples[:mid] + live, + sub_seed_idxs[:mid], + sub_samples[:mid], + graph_merger=graph_merger, ) yield from self._predict_sub_batch( - live, schema, sub_seed_idxs[mid:], sub_samples[mid:] + live, + sub_seed_idxs[mid:], + sub_samples[mid:], + graph_merger=graph_merger, ) return @@ -473,12 +483,11 @@ def predict_on_graph_sample_batch( live = self._get_live() schema = schema_to_input_feature_schema(self._data.schema, self._data.task) - merged_graph, merge_offsets = merge_lib.merge_graphs( - graph_samples, - schema, + merged_graph, merge_offsets = merge_lib.GraphMerger( + schema=schema, padding=self._data.padding, sentinel_offset=False, - ) + )(graph_samples) normalized_merged = live.normalizer.normalize_numpy(merged_graph) normalized_merged_jax = jax_lib.graph_to_jax_graph(normalized_merged) seed_node_idxs = merge_offsets[self._data.task.target_nodeset] @@ -591,7 +600,11 @@ def evaluate_generator( def batch_prediction_generator(): live = self._get_live() - schema = self._data.schema + graph_merger = merge_lib.GraphMerger( + schema=self._data.schema, + padding=self._data.padding, + sentinel_offset=False, + ) iterator = graph_samples if verbose >= 2: iterator = tqdm.tqdm(iterator, desc="Evaluation", total=num_eval_steps) @@ -601,12 +614,18 @@ def batch_prediction_generator(): batch.append(sample) if len(batch) == batch_size: yield from self._predict_sub_batch( - live, schema, np.zeros(len(batch), dtype=np.int32), batch + live, + np.zeros(len(batch), dtype=np.int32), + batch, + graph_merger=graph_merger, ) batch = [] if batch: yield from self._predict_sub_batch( - live, schema, np.zeros(len(batch), dtype=np.int32), batch + live, + np.zeros(len(batch), dtype=np.int32), + batch, + graph_merger=graph_merger, ) if verbose >= 1: diff --git a/dgf/src/learning/ten_lines/node_prediction_test.py b/dgf/src/learning/ten_lines/node_prediction_test.py index ef545ef..40f4dcd 100644 --- a/dgf/src/learning/ten_lines/node_prediction_test.py +++ b/dgf/src/learning/ten_lines/node_prediction_test.py @@ -20,7 +20,6 @@ import tempfile from typing import Tuple import unittest -from unittest import mock from absl import logging from absl.testing import absltest from absl.testing import parameterized @@ -566,22 +565,25 @@ def in_mem_graphs(): def test_predict_batch_insufficient_padding(self): """Tests that predict_batch handles InsufficientPaddingError by splitting.""" - original_merge_graphs = node_prediction_model.merge_lib.merge_graphs + original_graph_merger = node_prediction_model.merge_lib.GraphMerger call_count = 0 - def mock_merge_graphs(*args, **kwargs): - nonlocal call_count - call_count += 1 - if call_count == 1: - raise node_prediction_model.merge_lib.InsufficientPaddingError( - "Simulated insufficient padding" - ) - return original_merge_graphs(*args, **kwargs) + def mock_graph_merger(*args, **kwargs): + real_graph_merger = original_graph_merger(*args, **kwargs) + def graph_merger_wrapper(*call_args, **call_kwargs): + nonlocal call_count + call_count += 1 + if call_count == 1: + raise node_prediction_model.merge_lib.InsufficientPaddingError( + "Simulated insufficient padding" + ) + return real_graph_merger(*call_args, **call_kwargs) + return graph_merger_wrapper with unittest.mock.patch.object( node_prediction_model.merge_lib, - "merge_graphs", - side_effect=mock_merge_graphs, + "GraphMerger", + side_effect=mock_graph_merger, ): # In test, batch_size is 5 (from RAPID_TRAINING_KWARGS). # We need to call predict with at least 2 examples to trigger splitting. diff --git a/dgf/src/transform/merge.py b/dgf/src/transform/merge.py index c66f28c..8cf3a47 100644 --- a/dgf/src/transform/merge.py +++ b/dgf/src/transform/merge.py @@ -340,26 +340,7 @@ def __call__( ) -def merge_graph( - graphs: List[in_memory_graph.InMemoryGraph], - schema: schema_lib.GraphSchema, - padding: Optional[padding_lib.Padding] = None, - sentinel_offset: bool = True, - schema_cache: Optional[temporal_util.TimeseriesSchemaCache] = None, -) -> Tuple[in_memory_graph.InMemoryGraph, Dict[str, np.ndarray]]: - """Merges multiple `InMemoryGraph` instances into a single graph. - - Temporary wrapper around the `GraphMerger` class. - """ - return GraphMerger( - schema=schema, - padding=padding, - sentinel_offset=sentinel_offset, - schema_cache=schema_cache, - )(graphs) - -merge_graphs = merge_graph def create_padding_item(value, num_padding_items): @@ -398,7 +379,7 @@ def pad_graph_tensorflow( ) -> tf_in_memory_graph.TFInMemoryGraph: """Pads a TFInMemoryGraph with sentinels according to the padding. - Similar to the padding implemented in merge_graphs with sentinel_offsets and + Similar to the padding implemented in GraphMerger with sentinel_offsets and padder, but work on TFInMemoryGraph instead of (Numpy)InMemoryGraphs. @@ -520,18 +501,19 @@ def remove_padding_sentinels( schema: schema_lib.GraphSchema, offsets: Dict[str, np.ndarray], ) -> in_memory_graph.InMemoryGraph: - """Removes the sentinel nodes and edges added by `merge_graphs`. + """Removes the sentinel nodes and edges added by `GraphMerger`. The `graph` and `offsets` arguments are the two return values from - `merge_graphs`. This method requires `merge_graphs` to have been called with + `GraphMerger`. This method requires `GraphMerger` to have been called with `sentinel_offset=True`. Usage example: ```python - padded_graph, offsets = merge_graphs([graph], schema, padding, - sentinel_offset=True) - unpadded_graph = remove_padding_sentinels(merged_graph, schema, offsets) + padded_graph, offsets = GraphMerger( + schema, padding, sentinel_offset=True + )([graph]) + unpadded_graph = remove_padding_sentinels(padded_graph, schema, offsets) assert unpadded_graph == graph ``` @@ -541,7 +523,7 @@ def remove_padding_sentinels( Args: graph: The padded graph. schema: Graph schema. - offsets: Node set offsets returned by `merge_graphs`. + offsets: Node set offsets returned by `GraphMerger`. Returns: A new `InMemoryGraph` with sentinel nodes and edges removed. diff --git a/dgf/src/transform/merge_test.py b/dgf/src/transform/merge_test.py index 9ad2a5c..cd44cd2 100644 --- a/dgf/src/transform/merge_test.py +++ b/dgf/src/transform/merge_test.py @@ -46,9 +46,9 @@ def test_batch_with_padding(self): "e2": padding_lib.EdgeSetPadding(num_edges=6), }, ) - merged_graph, offsets = merge_lib.merge_graphs( - graphs, schema, padding=padding - ) + merged_graph, offsets = merge_lib.GraphMerger( + schema=schema, padding=padding + )(graphs) expected_merged_graph = in_memory_graph_lib.InMemoryGraph( node_sets={ "n1": in_memory_graph_lib.InMemoryNodeSet( @@ -108,7 +108,9 @@ def test_batch_no_padding(self): gen_test_graph.generate_in_memory_graph(False, False), ] schema = gen_test_graph.generate_schema(False, False, variable_length=False) - merged_graph, offsets = merge_lib.merge_graphs(graphs, schema, padding=None) + merged_graph, offsets = merge_lib.GraphMerger( + schema=schema, padding=None + )(graphs) expected_merged_graph = in_memory_graph_lib.InMemoryGraph( node_sets={ "n1": in_memory_graph_lib.InMemoryNodeSet( @@ -182,7 +184,7 @@ def test_batch_padding_too_small(self): r"Required at least 5 nodes \(including the sentinel node\), but the" r" padder only defines 4.", ): - _ = merge_lib.merge_graphs(graphs, schema, padding=padding) + _ = merge_lib.GraphMerger(schema=schema, padding=padding)(graphs) def test_batch_with_padding_no_sentinel_offset(self): graphs = [ @@ -200,12 +202,12 @@ def test_batch_with_padding_no_sentinel_offset(self): "e2": padding_lib.EdgeSetPadding(num_edges=6), }, ) - merged_graph, offsets = merge_lib.merge_graphs( - graphs, schema, sentinel_offset=False, padding=padding - ) - expected_merged_graph, _ = merge_lib.merge_graphs( - graphs, schema, sentinel_offset=True, padding=padding - ) + merged_graph, offsets = merge_lib.GraphMerger( + schema=schema, sentinel_offset=False, padding=padding + )(graphs) + expected_merged_graph, _ = merge_lib.GraphMerger( + schema=schema, sentinel_offset=True, padding=padding + )(graphs) test_util.assert_are_equal(self, merged_graph, expected_merged_graph) test_util.assert_are_equal( self, offsets, {"n2": np.array([0, 2]), "n1": np.array([0, 2])} @@ -217,12 +219,12 @@ def test_batch_no_padding_no_sentinel_offset(self): gen_test_graph.generate_in_memory_graph(False, False), ] schema = gen_test_graph.generate_schema(False, False, variable_length=False) - merged_graph, offsets = merge_lib.merge_graphs( - graphs, schema, padding=None, sentinel_offset=False - ) - expected_merged_graph, _ = merge_lib.merge_graphs( - graphs, schema, padding=None, sentinel_offset=True - ) + merged_graph, offsets = merge_lib.GraphMerger( + schema=schema, padding=None, sentinel_offset=False + )(graphs) + expected_merged_graph, _ = merge_lib.GraphMerger( + schema=schema, padding=None, sentinel_offset=True + )(graphs) test_util.assert_are_equal(self, merged_graph, expected_merged_graph) test_util.assert_are_equal( self, offsets, {"n2": np.array([0, 2]), "n1": np.array([0, 2])} @@ -372,16 +374,16 @@ def test_remove_padding_sentinels_with_padding(self): "e2": padding_lib.EdgeSetPadding(num_edges=6), }, ) - merged_graph, offsets = merge_lib.merge_graphs( - graphs, schema, padding=padding - ) + merged_graph, offsets = merge_lib.GraphMerger( + schema=schema, padding=padding + )(graphs) unpadded_graph = merge_lib.remove_padding_sentinels( merged_graph, schema, offsets ) - expected_unpadded_graph, _ = merge_lib.merge_graphs( - graphs, schema, padding=None - ) + expected_unpadded_graph, _ = merge_lib.GraphMerger( + schema=schema, padding=None + )(graphs) test_util.assert_are_equal(self, unpadded_graph, expected_unpadded_graph) def test_batch_with_timeseries_padding(self): @@ -439,16 +441,15 @@ def test_batch_with_timeseries_padding(self): edge_sets={}, ) # Check that schema cache is no longer required and can be omitted. - merged_graph_no_cache, _ = merge_lib.merge_graphs( - [g1, g2], schema, padding=padding, schema_cache=None - ) + merged_graph_no_cache, _ = merge_lib.GraphMerger( + schema=schema, padding=padding, schema_cache=None + )([g1, g2]) schema_cache = temporal_util.extract_timeseries_schema_cache(schema) - merged_graph, _ = merge_lib.merge_graphs( - [g1, g2], - schema, + merged_graph, _ = merge_lib.GraphMerger( + schema=schema, padding=padding, schema_cache=schema_cache, - ) + )([g1, g2]) test_util.assert_are_equal(self, merged_graph_no_cache, merged_graph) # Total real nodes: 2 + 1 = 3. Sentinel nodes: 5 - 3 = 2. Total nodes: 5. self.assertEqual(merged_graph.node_sets["n1"].num_nodes, 5) @@ -535,12 +536,11 @@ def test_batch_with_edge_set_timeseries_padding(self): }, ) schema_cache = temporal_util.extract_timeseries_schema_cache(schema) - merged_graph, _ = merge_lib.merge_graphs( - [g1, g2], - schema, + merged_graph, _ = merge_lib.GraphMerger( + schema=schema, padding=padding, schema_cache=schema_cache, - ) + )([g1, g2]) # Total real edges: 1 + 1 = 2. Total edges with padding: 3. edge_set = merged_graph.edge_sets["e1"] self.assertEqual(edge_set.adjacency.shape, (2, 3)) @@ -612,12 +612,11 @@ def test_batch_with_timeseries_only_padding(self): edge_sets={}, ) schema_cache = temporal_util.extract_timeseries_schema_cache(schema) - merged_graph, offsets = merge_lib.merge_graphs( - [g1, g2], - schema, + merged_graph, offsets = merge_lib.GraphMerger( + schema=schema, padding=padding, schema_cache=schema_cache, - ) + )([g1, g2]) # Total real nodes: 2 + 1 = 3. No sentinel nodes added. self.assertEqual(merged_graph.node_sets["n1"].num_nodes, 3) ts_feat = merged_graph.node_sets["n1"].features["ts"] @@ -690,24 +689,17 @@ def test_graph_merger_class_and_output_schema(self): }, edge_sets={}, ) - merger = merge_lib.GraphMerger(schema=schema, padding=padding) - output_schema = merger.output_schema() + graph_merger = merge_lib.GraphMerger(schema=schema, padding=padding) + output_schema = graph_merger.output_schema() self.assertIn("ts", output_schema.node_sets["n1"].features) self.assertIn("ts_mask", output_schema.node_sets["n1"].features) self.assertEqual( output_schema.node_sets["n1"].features["ts"].shape, (3,) ) - merged_graph, offsets = merger([g1, g2]) + merged_graph, _ = graph_merger([g1, g2]) self.assertEqual(merged_graph.node_sets["n1"].num_nodes, 5) - # Test merge_graph function alias - merged_graph_alias, offsets_alias = merge_lib.merge_graph( - [g1, g2], schema=schema, padding=padding - ) - test_util.assert_are_equal(self, merged_graph, merged_graph_alias) - test_util.assert_are_equal(self, offsets, offsets_alias) - def test_unknown_padding_keys_raise_value_error(self): schema = schema_lib.GraphSchema( node_sets={"n1": schema_lib.NodeSchema(features={})}, diff --git a/examples/node_classification_model_advanced.py b/examples/node_classification_model_advanced.py index 8e01ee4..6d35b27 100644 --- a/examples/node_classification_model_advanced.py +++ b/examples/node_classification_model_advanced.py @@ -103,6 +103,11 @@ def batch_generator( also_return_merge_offsets=False, batch_size=32, ): + graph_merger = dgf.transform.GraphMerger( + schema=schema, + padding=padding, + sentinel_offset=False, + ) for seed_node_idxs in dgf.transform.batch_indices_generator( seed_node_idxs, batch_size=batch_size, @@ -114,12 +119,7 @@ def batch_generator( try: # Merge the graph samples into a single graph. - merged_samples, merge_offsets = dgf.transform.merge_graphs( - graphs=samples, - schema=schema, - padding=padding, - sentinel_offset=False, - ) + merged_samples, merge_offsets = graph_merger(samples) except dgf.exception.InsufficientPaddingError: # Skip if the number of nodes is too large for the padding. continue diff --git a/examples/node_classification_pyg.py b/examples/node_classification_pyg.py index 9afbaf1..899e86f 100644 --- a/examples/node_classification_pyg.py +++ b/examples/node_classification_pyg.py @@ -166,6 +166,11 @@ def main(argv: typing.Sequence[str]) -> None: ) def batch_generator(seed_node_idxs, batch_size): + graph_merger = dgf.transform.GraphMerger( + schema=schema, + padding=None, + sentinel_offset=False, + ) for indices in dgf.transform.batch_indices_generator( seed_node_idxs, batch_size=batch_size, @@ -175,12 +180,7 @@ def batch_generator(seed_node_idxs, batch_size): samples = sampler.sample(indices.tolist()) try: - merged_samples, merge_offsets = dgf.transform.merge_graphs( - graphs=samples, - schema=schema, - padding=None, - sentinel_offset=False, - ) + merged_samples, merge_offsets = graph_merger(samples) except dgf.exception.InsufficientPaddingError: continue