atlas-1-demo / app.py
Reverb's picture
Update app.py
d3534d0 verified
Raw History Blame Contribute Delete
9.53 kB
"""
Atlas 1 Demo - Hugging Face Space Version
Interactive demo for molecular property prediction with 2D/3D visualization.
"""
import gradio as gr
import torch
import numpy as np
from rdkit import Chem
from rdkit.Chem import AllChem
from torch_geometric.data import Data
import sys
import os
sys.path.insert(0, os.path.join(os.path.dirname(__file__), '..'))
from src.models.geometric_gnn import GeometricGNN
# Try to import gradio_molecule3d, fallback to 2D if not available
try:
from gradio_molecule3d import Molecule3D
HAS_MOLECULE3D = True
except ImportError:
HAS_MOLECULE3D = False
print("gradio_molecule3d not found, will use 2D visualization")
# Atom type mapping (same as training)
ATOM_TYPES = {
'H': 0, 'C': 1, 'N': 2, 'O': 3, 'F': 4,
'S': 5, 'Cl': 6, 'Br': 7, 'I': 8, 'P': 9, 'Si': 10
}
def load_model(checkpoint_path='checkpoints/best.pt'):
"""Load trained model."""
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
model = GeometricGNN(
hidden_dim=256,
num_layers=6,
num_frequencies=16,
max_l=3,
cutoff=5.0,
).to(device)
checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False)
model.load_state_dict(checkpoint['model_state_dict'])
model.eval()
return model, device
def smiles_to_3d_coords(smiles):
"""Convert SMILES to 3D coordinates using RDKit."""
mol = Chem.MolFromSmiles(smiles)
if mol is None:
return None, None, None
# Add hydrogens
mol = Chem.AddHs(mol)
# Generate 3D coordinates
AllChem.EmbedMolecule(mol, randomSeed=42)
AllChem.MMFFOptimizeMolecule(mol)
# Extract coordinates and atom types
conf = mol.GetConformer()
coords = []
atom_types = []
for atom in mol.GetAtoms():
pos = conf.GetAtomPosition(atom.GetIdx())
coords.append([pos.x, pos.y, pos.z])
symbol = atom.GetSymbol()
atom_type = ATOM_TYPES.get(symbol, 0)
atom_types.append(atom_type)
return np.array(coords), np.array(atom_types), mol
def create_graph_data(coords, atom_types, cutoff=5.0):
"""Create PyG Data object from coordinates and atom types."""
# One-hot encode atom types
num_atoms = len(atom_types)
x = torch.zeros(num_atoms, 11)
x[torch.arange(num_atoms), atom_types] = 1.0
# Convert coords to tensor
pos = torch.FloatTensor(coords)
# Create edges based on distance cutoff
edge_index = []
for i in range(num_atoms):
for j in range(num_atoms):
if i != j:
dist = torch.norm(pos[i] - pos[j])
if dist < cutoff:
edge_index.append([i, j])
edge_index = torch.LongTensor(edge_index).t()
return Data(x=x, pos=pos, edge_index=edge_index)
def visualize_molecule_2d(mol):
"""Create 2D visualization of molecule."""
if mol is None:
return None
try:
from rdkit.Chem import Draw
import io
import base64
# Generate 2D image
img = Draw.MolToImage(mol, size=(500, 500))
buf = io.BytesIO()
img.save(buf, format='PNG')
img_str = base64.b64encode(buf.getvalue()).decode()
return f'<img src="data:image/png;base64,{img_str}" style="max-width: 100%; border-radius: 8px;">'
except Exception as e:
return f"<p style='color: red;'>Visualization error: {str(e)}</p>"
def get_molecule_3d_data(mol):
"""Get 3D molecular data for gradio_molecule3d (PDB format)."""
if mol is None or not HAS_MOLECULE3D:
return None
try:
import hashlib
# Convert to PDB format
pdb_block = Chem.MolToPDBBlock(mol)
# Create static directory if it doesn't exist
os.makedirs("molecule_cache", exist_ok=True)
# Create unique filename based on content hash
content_hash = hashlib.md5(pdb_block.encode()).hexdigest()[:8]
pdb_path = f"molecule_cache/mol_{content_hash}.pdb"
# Save PDB file
with open(pdb_path, 'w') as f:
f.write(pdb_block)
return pdb_path
except Exception as e:
print(f"3D data error: {e}")
return None
def predict_homo(smiles, model, device):
"""Predict HOMO energy for a molecule."""
# Convert SMILES to 3D structure
coords, atom_types, mol = smiles_to_3d_coords(smiles)
if coords is None:
return None, None, "❌ Invalid SMILES string!", ""
# Create graph data
data = create_graph_data(coords, atom_types).to(device)
# Predict
with torch.no_grad():
output = model(data)
pred_homo = output['pred_mean'].item()
# Denormalize (QM9 HOMO mean=-0.232, std=0.043)
pred_homo_denorm = pred_homo * 0.043 - 0.232
# Create 2D visualization
viz_2d = visualize_molecule_2d(mol)
# Create 3D data
viz_3d = get_molecule_3d_data(mol)
# Create results text
results = f"""
### ✅ Prediction Results
**Predicted HOMO Energy**: `{pred_homo_denorm:.4f} eV`
**Normalized Value**: `{pred_homo:.4f}`
**Molecule Info**:
- Atoms: {len(atom_types)}
- Edges: {data.edge_index.shape[1]}
**Model**: Atlas 1 (2.5M parameters)
"""
return viz_2d, viz_3d, results, ""
# Load model
print("Loading model...")
model, device = load_model()
print("Model loaded successfully!")
# Example molecules
EXAMPLES = [
["C", "Methane"],
["CC", "Ethane"],
["C1=CC=CC=C1", "Benzene"],
["CCO", "Ethanol"],
["CC(=O)O", "Acetic acid"],
["c1ccccc1O", "Phenol"],
["CN1C=NC2=C1C(=O)N(C(=O)N2C)C", "Caffeine"],
]
# Create Gradio interface
def demo_interface(smiles):
"""Main demo function."""
try:
viz_2d, viz_3d, results, error = predict_homo(smiles, model, device)
return viz_2d, viz_3d, results
except Exception as e:
return None, None, f"❌ Error: {str(e)}"
# Build interface - Fixed for Gradio 6.0
with gr.Blocks(title="Atlas 1: HOMO Predictor") as demo:
gr.Markdown("""
# 🧪 Atlas 1: Molecular Property Prediction
Predict the **Highest Occupied Molecular Orbital (HOMO)** energy using a SOTA Geometric GNN.
**Model Performance**:
- RMSE: 0.109 eV
- MAE: 0.061 eV
- Pearson r: 0.994
- Trained on QM9 dataset (130k molecules)
**Enter a SMILES string** to visualize the molecule and predict its HOMO energy!
""")
with gr.Row():
with gr.Column(scale=1):
smiles_input = gr.Textbox(
label="SMILES String",
placeholder="Enter molecule SMILES (e.g., C1=CC=CC=C1 for benzene)",
value="C1=CC=CC=C1"
)
predict_btn = gr.Button("🔮 Predict HOMO Energy", variant="primary", size="lg")
gr.Markdown("### 📚 Example Molecules")
gr.Examples(
examples=[[ex[0]] for ex in EXAMPLES],
inputs=smiles_input,
label="Click to try:",
)
with gr.Column(scale=1):
gr.Markdown("### 2D Structure")
viz_2d_output = gr.HTML(label="2D Molecular Structure")
if HAS_MOLECULE3D:
gr.Markdown("### 3D Interactive View")
viz_3d_output = Molecule3D(
label="3D Molecular Structure",
reps=[
{
"model": 0,
"chain": "",
"resname": "",
"style": "stick",
"color": "whiteCarbon",
"residue_range": "",
"around": 0,
"byres": False,
"visible": False
}
]
)
results_output = gr.Markdown(label="Prediction Results")
# Define outputs based on whether 3D is available
if HAS_MOLECULE3D:
outputs = [viz_2d_output, viz_3d_output, results_output]
else:
outputs = [viz_2d_output, gr.Textbox(visible=False), results_output]
gr.Markdown("💡 *Install `gradio_molecule3d` for interactive 3D visualization*")
predict_btn.click(
fn=lambda s: demo_interface(s),
inputs=smiles_input,
outputs=outputs
)
# Auto-run on load
demo.load(
fn=lambda s: demo_interface(s),
inputs=smiles_input,
outputs=outputs
)
gr.Markdown("""
---
### About Atlas 1
**Architecture**:
- Continuous-filter convolutions (CFConv)
- SE(3)-equivariant message passing
- 6 interaction blocks
- 256 hidden dimensions
- Spherical harmonics (L=3)
**Training**:
- Dataset: QM9 (130,831 molecules)
- Loss: MSE
- Optimizer: AdamW with warmup
- Training time: 1.4 hours (200 epochs)
**Performance**: Competitive (Top 50% of published models)
---
**Part of the Atlas Series**: From properties to structures
""")
if __name__ == "__main__":
# Fixed for HF Spaces - removed share parameter and moved theme to launch
demo.launch(
server_name="0.0.0.0",
server_port=7860,
ssr_mode=False # Disable SSR to avoid warnings
)