matcha.datamodules.utils

Utility classes and collation functions for datamodules.

Attributes

collate_fns

Classes

HandleMissing

Enum class to handle missing values in the data.

CombinedStackDataset

A Dataset that can combine multiple `StackDataset`s.

Functions

collate_fn_pyg_graph(→ torch_geometric.data.Data)

Collate function for PyTorch Geometric graphs.

concat_tensors_merge_fn(→ torch.Tensor)

Merge a list of tensors by concatenating them.

Module Contents

matcha.datamodules.utils.collate_fn_pyg_graph(graphs: list[torch_geometric.data.Data], collate_fn_map=None) torch_geometric.data.Data[source]

Collate function for PyTorch Geometric graphs.

Batches a list of PyG Data objects into a single batched Data object. Handles shortest path distances (spd) and virtual node initialization.

Parameters:
  • graphs – list of PyG Data objects

  • collate_fn_map – unused, for compatibility

Returns:

batched PyG Data object

matcha.datamodules.utils.collate_fns: dict[str, collections.abc.Callable]
class matcha.datamodules.utils.HandleMissing(*args, **kwds)[source]

Bases: enum.Enum

Enum class to handle missing values in the data.

RAISE = 'raise'
FILL = 'fill'
matcha.datamodules.utils.concat_tensors_merge_fn(values: list[torch.Tensor], dim=0) torch.Tensor[source]

Merge a list of tensors by concatenating them.

Parameters:
  • values – a list of tensors to concatenate

  • dim – the dimension along which to concatenate the tensors

Returns:

the concatenated tensor

class matcha.datamodules.utils.CombinedStackDataset(datasets: list[torch.utils.data.StackDataset], merge_fn: dict[str, collections.abc.Callable] = None)[source]

Bases: torch.utils.data.Dataset

A Dataset that can combine multiple `StackDataset`s.

This class is useful when you have multiple datasets that you want to combine into a single dataset.

If a key appears in multiple datasets, the values get merged according to merge_fn.

Parameters:

datasets – a list of datasets to combine

datasets
merge_fn