Download data.py from Flogrammer/Mol-JEPA: direct link, hf CLI and curl.
- Browser
- Download file 2.18 kB
-
https://huggingface.co/Flogrammer/Mol-JEPA/resolve/main/data.py
- Command line
-
hf download hf://Flogrammer/Mol-JEPA/data.py
-
curl -L -o data.py https://huggingface.co/Flogrammer/Mol-JEPA/resolve/main/data.py
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 | |