Mol-JEPA / data.py
Flogrammer's picture
Upload folder using huggingface_hub
0c1df4b verified
Raw History Blame Contribute Delete
2.18 kB
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