Spaces:
Sleeping
Sleeping
| import os | |
| import gradio as gr | |
| import torch | |
| import matgl | |
| # 2026-06-05 migration: the previous model "M3GNet-MP-2021.2.8-PES" (DGL backend) | |
| # was RETIRED from MatGL's hosting (404 everywhere) and MatGL dropped the DGL | |
| # backend at 3.0.4 (now PyG-only). This Space now runs MatGL 4.x (PyG default) + | |
| # the current MatPES-trained M3GNet PES model. matgl is PINNED in the Dockerfile | |
| # to prevent the unpinned-dependency drift that broke this Space. | |
| MODEL_NAME = "M3GNet-PES-MatPES-PBE-2025.2" | |
| # Global model cache - load lazily to avoid startup timeout | |
| _model = None | |
| _model_error = None | |
| def get_model(): | |
| """Lazy load the M3GNet model on first use.""" | |
| global _model, _model_error | |
| if _model_error: | |
| return None, _model_error | |
| if _model is None: | |
| try: | |
| print(f"Loading {MODEL_NAME} ...") | |
| _model = matgl.load_model(MODEL_NAME) | |
| print("Model loaded successfully.") | |
| except Exception as e: | |
| _model_error = str(e) | |
| print(f"Failed to load model: {e}") | |
| import traceback | |
| traceback.print_exc() | |
| return None, _model_error | |
| return _model, None | |
| def predict_properties(cif_string: str) -> dict: | |
| """ | |
| Predicts formation energy from a CIF string using M3GNet. | |
| Uses ASE Calculator interface for energy calculation. | |
| """ | |
| if not cif_string or not cif_string.strip(): | |
| return {"status": "error", "message": "CIF string is empty"} | |
| pot, error = get_model() | |
| if error: | |
| return {"status": "error", "message": f"Model load failed: {error}"} | |
| try: | |
| from pymatgen.core import Structure | |
| from pymatgen.io.ase import AseAtomsAdaptor | |
| from matgl.ext.ase import PESCalculator | |
| # Parse CIF to pymatgen Structure | |
| struct = Structure.from_str(cif_string, fmt="cif") | |
| # Convert to ASE Atoms | |
| adaptor = AseAtomsAdaptor() | |
| atoms = adaptor.get_atoms(struct) | |
| # Create calculator and attach to atoms | |
| calc = PESCalculator(potential=pot) | |
| atoms.calc = calc | |
| # Get potential energy (total energy) | |
| e_total = atoms.get_potential_energy() | |
| e_per_atom = e_total / len(atoms) | |
| return { | |
| "status": "success", | |
| "formation_energy_per_atom": round(e_per_atom, 4), | |
| "unit": "eV/atom", | |
| "formula": struct.composition.reduced_formula, | |
| "num_atoms": len(struct), | |
| "model": MODEL_NAME, | |
| } | |
| except Exception as e: | |
| return {"status": "error", "message": str(e)} | |
| def health_check() -> dict: | |
| """Health check endpoint.""" | |
| pot, error = get_model() | |
| if error: | |
| return {"status": "error", "message": error} | |
| return {"status": "healthy", "model": MODEL_NAME} | |
| # Create Gradio interface | |
| with gr.Blocks() as demo: | |
| gr.Markdown("# MatGL M3GNet Property Predictor") | |
| gr.Markdown("Predict formation energy for crystal structures using M3GNet universal potential.") | |
| gr.Markdown("**Note:** First prediction may take longer as the model loads on-demand.") | |
| with gr.Tab("Predict"): | |
| cif_input = gr.Textbox( | |
| label="CIF Structure", | |
| placeholder="Paste CIF content here...", | |
| lines=10 | |
| ) | |
| predict_btn = gr.Button("Predict", variant="primary") | |
| output = gr.JSON(label="Result") | |
| predict_btn.click(predict_properties, inputs=cif_input, outputs=output) | |
| gr.Markdown("### Example CIF (NaCl)") | |
| gr.Markdown("""``` | |
| data_NaCl | |
| _symmetry_space_group_name_H-M 'F m -3 m' | |
| _cell_length_a 5.64 | |
| _cell_length_b 5.64 | |
| _cell_length_c 5.64 | |
| _cell_angle_alpha 90 | |
| _cell_angle_beta 90 | |
| _cell_angle_gamma 90 | |
| loop_ | |
| _atom_site_label | |
| _atom_site_type_symbol | |
| _atom_site_fract_x | |
| _atom_site_fract_y | |
| _atom_site_fract_z | |
| Na1 Na 0.0 0.0 0.0 | |
| Cl1 Cl 0.5 0.5 0.5 | |
| ```""") | |
| with gr.Tab("Health"): | |
| health_btn = gr.Button("Check Health") | |
| health_output = gr.JSON() | |
| health_btn.click(health_check, outputs=health_output) | |
| if __name__ == "__main__": | |
| demo.launch(server_name="0.0.0.0", server_port=7860) | |