import torch from torch_geometric.data import Data class MultimodalData(Data): def __init__(self): super().__init__() def add_metadata(self, metadata: dict): """Add metadata fields to the data object.""" for key, value in metadata.items(): setattr(self, key, value) def add_mol_embedding(self, encoder_name: str, embedding: torch.Tensor): """Global molecule-level embedding, shape (d,)""" embedding = embedding.flatten() if embedding.dim() > 1 else embedding embedding = embedding.unsqueeze(0) if embedding.dim() == 1 else embedding embedding = embedding.float() # Zero Imputation - this is mostly for when using target vectors as modalities embedding = torch.nan_to_num(embedding, nan=0.0) setattr(self, f"{encoder_name}_x", embedding) def add_atom_embedding(self, encoder_name: str, x: torch.Tensor): """Per-atom embeddings, x shape (n_atoms, d)""" setattr(self, f"{encoder_name}_x", x) def add_graph( self, encoder_name: str, x: torch.Tensor, edge_index: torch.Tensor, edge_attr: torch.Tensor = None, ): """Graph representation, x shape (n_atoms, d_a), edge_index shape (2, num_edges), x shape (num_edges, d_e)""" setattr(self, f"{encoder_name}_x", x) setattr(self, f"{encoder_name}_edge_index", edge_index) setattr(self, f"{encoder_name}_edge_attr", edge_attr) def __inc__(self, key, value, *args, **kwargs): if key.endswith("edge_index"): prefix = key.replace("edge_index", "x") x = getattr(self, prefix, None) if x is not None: return x.shape[0] # No node features for this modality on this sample -> treat as # an empty graph (0 nodes) instead of falling through to PyG's # default which raises when num_nodes cannot be inferred. return 0 return super().__inc__(key, value, *args, **kwargs) def __cat_dim__(self, key, value, *args, **kwargs): if key.endswith("edge_index"): return 1 return 0