matcha.datamodules.utils
Utility classes and collation functions for datamodules.
Attributes
Classes
Enum class to handle missing values in the data. |
|
A Dataset that can combine multiple `StackDataset`s. |
Functions
|
Collate function for PyTorch Geometric graphs. |
|
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.EnumEnum 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.DatasetA 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