Views
No views yet

pip install torch torch-geometric transformers safetensors rdkit molfeat numpyAutoModel API. Because it ships
custom modeling code, load it with trust_remote_code=True:1from transformers import AutoModel
2
3model = AutoModel.from_pretrained("Flogrammer/Mol-JEPA", trust_remote_code=True)
4model.eval()ModelOutput,
so the fields are accessible both by name and by index.1smiles_list = ["Cn1cnc2n(C)c(=O)n(C)c(=O)c12", "CC(=O)Oc1ccccc1C(=O)O"]
2out = model(smiles_list)
3
4print("Predicted embeddings shape:", out.predictions.shape) # batch size, modalities, embedding dimension
5>>> Predicted embeddings shape: torch.Size([2, 12, 512])
6
7print("CLS token shape:", out.cls.shape) # batch size, embedding dimension
8>>> CLS token shape: torch.Size([2, 512])
9
10print("Latent embeddings shape:", out.embeddings.shape) # batch size, modalities, latent dimension
11>>> Latent embeddings shape: torch.Size([2, 13, 512])
12
13# Index access (predictions, cls, embeddings) is also supported:
14predictions, cls, embeddings = out[0], out[1], out[2]1from tabicl import TabICLRegressor
2
3# Take features from above
4cls_features = out.cls
5
6# Do some splitting
7train_idx, test_index = ...
8
9model = TabICLRegressor(random_state=123)
10model.fit(cls_features[train_idx], y[train_idx])
11test_pred = model.predict(cls_features[test_idx])1class TransformerProbe(nn.Module):
2 def __init__(self, n_tokens, token_dim, n_layers=2, n_heads=4, dropout=0.1):
3 super().__init__()
4 self.n_tokens = n_tokens
5 self.token_dim = token_dim
6 encoder_layer = nn.TransformerEncoderLayer(
7 d_model=token_dim,
8 nhead=n_heads,
9 dim_feedforward=token_dim * 4,
10 dropout=dropout,
11 batch_first=True,
12 norm_first=True,
13 )
14 self.transformer = nn.TransformerEncoder(
15 encoder_layer, num_layers=n_layers, enable_nested_tensor=False
16 )
17 self.head = nn.Linear(token_dim, 1)
18
19 def forward(self, x):
20 # x: (B, n_tokens * token_dim) -> (B, n_tokens, token_dim)
21 x = x.view(-1, self.n_tokens, self.token_dim)
22 x = self.transformer(x)
23 # Mean pool over tokens, then predict
24 x = x.mean(dim=1)
25 return self.head(x)
26
27
28# Extract the number of modalities
29n_tokens = out.predictions.shape[-1]
30
31# Build model
32hidden_dim = 512
33model = TransformerProbe(n_tokens=n_tokens, token_dim=hidden_dim)
34
35# Optimize
36epochs = 10
37for _ in range(epochs):
38 ...