"""
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''
except Exception as e:
return f"
Visualization error: {str(e)}
" 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 )