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
7 changes: 3 additions & 4 deletions dgf/src/analyse/padding.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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(
[<some graph samples>],
merged_graph_samples, nodeset_offsets = dgf.transform.GraphMerger(
schema,
padding=padding,
)
)([<some graph samples>])
```

The padding size is:
Expand Down
2 changes: 1 addition & 1 deletion dgf/src/api/transform.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
2 changes: 1 addition & 1 deletion dgf/src/learning/ten_lines/BUILD
Original file line number Diff line number Diff line change
Expand Up @@ -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",
],
Expand Down
17 changes: 9 additions & 8 deletions dgf/src/learning/ten_lines/dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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:
Expand All @@ -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
Expand Down
41 changes: 18 additions & 23 deletions dgf/src/learning/ten_lines/link_prediction_dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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]
):
Expand All @@ -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() -> (
Expand All @@ -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() -> (
Expand All @@ -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 = (
Expand Down Expand Up @@ -817,25 +815,23 @@ 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(
batch_seed.pos_trg_node_idxs,
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 = (
Expand All @@ -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,
Expand Down
40 changes: 21 additions & 19 deletions dgf/src/learning/ten_lines/link_prediction_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -601,25 +601,26 @@ 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,
sub_src_samples: List[in_memory_graph.InMemoryGraph],
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)
Expand Down Expand Up @@ -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)
Expand Down
25 changes: 14 additions & 11 deletions dgf/src/learning/ten_lines/link_prediction_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
2 changes: 1 addition & 1 deletion dgf/src/learning/ten_lines/node_prediction_dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
Loading