""" 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 )