Views
No views yet
1{
2 "name": "baseline",
3 "hidden_channels": 64,
4 "num_layers": 2,
5 "num_timesteps": 2,
6 "dropout": 0.2,
7 "learning_rate": 0.001,
8 "weight_decay": 1e-05,
9 "batch_size": 32,
10 "epochs": 50,
11 "patience": 10
12}pip install torch torch-geometric rdkit-pypi1import torch
2from torch_geometric.nn import AttentiveFP
3from rdkit import Chem
4from torch_geometric.data import Data
5
6# Load the model
7model_dict = torch.load('pytorch_model.bin', map_location='cpu')
8state_dict = model_dict['model_state_dict']
9hyperparams = model_dict['hyperparameters']
10
11# Create model with correct architecture
12model = AttentiveFP(
13 in_channels=10, # Enhanced atom features
14 hidden_channels=hyperparams["hidden_channels"],
15 out_channels=1,
16 edge_dim=6, # Enhanced bond features
17 num_layers=hyperparams["num_layers"],
18 num_timesteps=hyperparams["num_timesteps"],
19 dropout=hyperparams["dropout"],
20)
21
22model.load_state_dict(state_dict)
23model.eval()1def smiles_to_data(smiles):
2 """Convert SMILES string to PyG Data object"""
3 mol = Chem.MolFromSmiles(smiles)
4 if mol is None:
5 return None
6
7 # Enhanced atom features (10 dimensions)
8 atom_features = []
9 for atom in mol.GetAtoms():
10 features = [
11 atom.GetAtomicNum(),
12 atom.GetTotalDegree(),
13 atom.GetFormalCharge(),
14 atom.GetTotalNumHs(),
15 atom.GetNumRadicalElectrons(),
16 int(atom.GetIsAromatic()),
17 int(atom.IsInRing()),
18 # Hybridization as one-hot (3 dimensions)
19 int(atom.GetHybridization() == Chem.rdchem.HybridizationType.SP),
20 int(atom.GetHybridization() == Chem.rdchem.HybridizationType.SP2),
21 int(atom.GetHybridization() == Chem.rdchem.HybridizationType.SP3)
22 ]
23 atom_features.append(features)
24
25 x = torch.tensor(atom_features, dtype=torch.float)
26
27 # Enhanced bond features (6 dimensions)
28 edges_list = []
29 edge_features = []
30 for bond in mol.GetBonds():
31 i = bond.GetBeginAtomIdx()
32 j = bond.GetEndAtomIdx()
33 edges_list.extend([[i, j], [j, i]])
34
35 features = [
36 # Bond type as one-hot (4 dimensions)
37 int(bond.GetBondType() == Chem.rdchem.BondType.SINGLE),
38 int(bond.GetBondType() == Chem.rdchem.BondType.DOUBLE),
39 int(bond.GetBondType() == Chem.rdchem.BondType.TRIPLE),
40 int(bond.GetBondType() == Chem.rdchem.BondType.AROMATIC),
41 # Additional features (2 dimensions)
42 int(bond.GetIsConjugated()),
43 int(bond.IsInRing())
44 ]
45 edge_features.extend([features, features])
46
47 if not edges_list:
48 return None
49
50 edge_index = torch.tensor(edges_list, dtype=torch.long).t()
51 edge_attr = torch.tensor(edge_features, dtype=torch.float)
52
53 return Data(x=x, edge_index=edge_index, edge_attr=edge_attr)
54
55def predict(model, smiles):
56 """Make prediction for a SMILES string"""
57 data = smiles_to_data(smiles)
58 if data is None:
59 return None
60
61 batch = torch.zeros(data.num_nodes, dtype=torch.long)
62 with torch.no_grad():
63 output = model(data.x, data.edge_index, data.edge_attr, batch)
64 return output.item()
65
66# Example usage
67smiles = "CC(=O)OC1=CC=CC=C1C(=O)O" # Aspirin
68prediction = predict(model, smiles)
69print(f"Prediction for {smiles}: {prediction}")1@misc{pyrosageames,
2 title={Pyrosage AMES AttentiveFP Model},
3 author={Pyrosage Team},
4 year={2024},
5 publisher={Hugging Face},
6 url={https://huggingface.co/alarv/pyrosage-ames-attentivefp}
7}