Spaces:
Build error
Build error
Download app.py from Reverb/atlas-1-demo: direct link, hf CLI and curl.
- Browser
- Download file 9.53 kB
-
https://huggingface.co/spaces/Reverb/atlas-1-demo/resolve/main/app.py
- Command line
-
hf download hf://spaces/Reverb/atlas-1-demo/app.py
-
curl -L -o app.py https://huggingface.co/spaces/Reverb/atlas-1-demo/resolve/main/app.py
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 | |
| ) | |