Spaces:
Sleeping
Sleeping
Download app.py from glasgow-lab/PFNet: direct link, hf CLI and curl.
- Browser
- Download file 46.9 kB
-
https://huggingface.co/spaces/glasgow-lab/PFNet/resolve/main/app.py
- Command line
-
hf download hf://spaces/glasgow-lab/PFNet/app.py
-
curl -L -o app.py https://huggingface.co/spaces/glasgow-lab/PFNet/resolve/main/app.py
46.9 kB
| import os | |
| import tempfile | |
| import glob | |
| import json | |
| import zipfile | |
| import urllib.request | |
| import traceback | |
| import gradio as gr | |
| import numpy as np | |
| import pandas as pd | |
| import matplotlib.pyplot as plt | |
| import matplotlib.colors as mcolors | |
| import plotly.express as px | |
| import plotly.graph_objects as go | |
| import torch | |
| import py3Dmol | |
| import html | |
| from pfnet.run_inference import predict | |
| from pfnet.plot import BFactorPlot | |
| from pigeon_feather.analysis import get_res_avg_logP, get_res_avg_logP_std, get_res_avg_log_kex | |
| from pigeon_feather.hxio import load_HXMS_file | |
| import MDAnalysis | |
| from mpl_toolkits.axes_grid1.inset_locator import inset_axes | |
| def process_pf_value(value): | |
| """ | |
| Safely converts a protection factor value to float, handling non-numeric types. | |
| """ | |
| try: | |
| return float(value) | |
| except (ValueError, TypeError): | |
| return np.nan | |
| def get_display_state_name(state_key, state_keys): | |
| """ | |
| Helper function to determine whether to keep or remove index from state name. | |
| If there are duplicate original state names, keep the index; otherwise remove it. | |
| """ | |
| original_state_names = ["_".join(key.split("_")[:-1]) for key in state_keys] | |
| has_duplicates = len(original_state_names) != len(set(original_state_names)) | |
| if has_duplicates: | |
| return state_key | |
| else: | |
| return "_".join(state_key.split("_")[:-1]) | |
| def resolve_uploaded_file_path(file_obj): | |
| """Normalize Gradio file input into a filesystem path.""" | |
| if file_obj is None: | |
| return None | |
| if isinstance(file_obj, str): | |
| return file_obj | |
| file_name = getattr(file_obj, "name", None) | |
| if isinstance(file_name, str): | |
| return file_name | |
| if isinstance(file_obj, dict): | |
| for key in ("name", "path"): | |
| value = file_obj.get(key) | |
| if isinstance(value, str): | |
| return value | |
| return None | |
| def get_default_structure_color_range(file1, file2): | |
| """Pick structure coloring defaults from the number of HXMS inputs.""" | |
| has_file1 = resolve_uploaded_file_path(file1) is not None | |
| has_file2 = resolve_uploaded_file_path(file2) is not None | |
| if has_file1 and has_file2: | |
| return -25, 25 | |
| return 0, 50 | |
| def predict_gradio( | |
| file1, | |
| file2, | |
| pdb_id_input, | |
| pdb_file_input, | |
| centroid_model, | |
| uptake_plots, | |
| vmin, | |
| vmax, | |
| ): | |
| print(file1, file2,) | |
| input_files = [ | |
| resolve_uploaded_file_path(file_obj) | |
| for file_obj in (file1, file2) | |
| ] | |
| input_files = [path for path in input_files if path] | |
| if len(input_files) == 0: | |
| raise gr.Error("Please upload at least one HXMS file.") | |
| if vmin is None or vmax is None: | |
| vmin, vmax = get_default_structure_color_range(file1, file2) | |
| if vmin >= vmax: | |
| raise gr.Error("Structure color min must be smaller than structure color max.") | |
| hdxms_data_list = [ load_HXMS_file(input_file)[0] for input_file in input_files ] | |
| global state_names | |
| state_names = [state.state_name for data in hdxms_data_list for state in data.states] | |
| if not centroid_model: | |
| _all_tps = [tp for data in hdxms_data_list for state in data.states for pep in state.peptides for tp in pep.timepoints if tp.deut_time != np.inf and tp.deut_time != 0.0] | |
| _envelope_tps = [tp for tp in _all_tps if tp.isotope_envelope is not None and len(tp.isotope_envelope) > 3] | |
| if len(_envelope_tps) == 0: | |
| raise gr.Error( | |
| "No envelope data found in the input file(s). " | |
| "You may want to use centroid mode or check the input." | |
| ) | |
| global output_dir | |
| output_dir = tempfile.mkdtemp() | |
| os.makedirs(f"{output_dir}/pfnet_output", exist_ok=True) | |
| os.makedirs(f"{output_dir}/pfnet_plots", exist_ok=True) | |
| output_dicts = {} | |
| for idx, input_file in enumerate(input_files): | |
| output_dict_i = predict( | |
| input=input_file, | |
| output_json=f"{output_dir}/pfnet_output/results_{state_names[idx]}_{idx}.json", | |
| centroid_model=centroid_model, | |
| uptake_plots=uptake_plots, | |
| plots_dir=f"{output_dir}/pfnet_plots", | |
| benchmark=False | |
| ) | |
| output_dicts[f"{state_names[idx]}_{idx}"] = output_dict_i | |
| if uptake_plots: | |
| for plot_file in glob.glob(f"{output_dir}/pfnet_plots/PFNet_uptake_?.pdf"): | |
| os.rename(plot_file, plot_file.replace("PFNet_uptake_", f"PFNet_uptake_{state_names[idx]}_")) | |
| for envelope_type in ["centroid", "envelope"]: | |
| plot_file = f"{output_dir}/pfnet_plots/ae_histogram_{envelope_type}.png" | |
| if os.path.exists(plot_file): | |
| os.rename(plot_file, plot_file[:-4] + f"_{state_names[idx]}.png") | |
| analysis_objs = { f"{state_names[idx]}_{idx}": output_dicts[f"{state_names[idx]}_{idx}"]["analysis_objs"][f"analysis_pfnet"] | |
| for idx in range(len(state_names)) if f"{state_names[idx]}_{idx}" in output_dicts } | |
| log_kex_plot = get_log_kex_plot(analysis_objs, output_dir) | |
| csv_results_path_list = get_csv_results(analysis_objs, output_dicts, output_dir) | |
| summary, summary_path = get_summary(output_dicts, hdxms_data_list) | |
| dg_pdb_path = None | |
| if pdb_id_input or pdb_file_input: | |
| pdb_path = load_pdb_data(pdb_id_input, pdb_file_input) | |
| dg_pdb_path = make_BFactorPlot(pdb_path, analysis_objs, output_dir) | |
| molstar_view_html = get_structure_html(dg_pdb_path, vmin=vmin, vmax=vmax) | |
| else: | |
| molstar_view_html = None | |
| json_files = glob.glob(f"{output_dir}/pfnet_output/results_*.json") | |
| heatmaps = get_heatmap(hdxms_data_list, output_dir) | |
| ae_histogram = sorted(glob.glob(f"{output_dir}/pfnet_plots/ae_histogram_*.png")) | |
| if ae_histogram: | |
| ae_histogram_labeled = [] | |
| for img_path in ae_histogram: | |
| basename = os.path.basename(img_path) | |
| if "centroid" in basename: | |
| envelope_type = "Centroid" | |
| elif "envelope" in basename: | |
| envelope_type = "Envelope" | |
| else: | |
| envelope_type = "Unknown" | |
| for state in state_names: | |
| if state in basename: | |
| title = f"{envelope_type} - {state}" | |
| break | |
| else: | |
| title = envelope_type | |
| ae_histogram_labeled.append((img_path, title)) | |
| ae_info = "" | |
| else: | |
| ae_histogram_labeled = None | |
| ae_info = "No AE histograms generated. Enable 'Generate uptake plots' in Settings to create absolute error histograms." | |
| uptake_gallery = glob.glob(f"{output_dir}/pfnet_plots/PFNet_uptake_*.pdf") | |
| if uptake_gallery: | |
| uptake_info = "" | |
| else: | |
| uptake_info = "No uptake plots generated. Enable 'Generate uptake plots' in Settings to create uptake curves." | |
| zip_path = f"{output_dir}/pfnet_results.zip" | |
| with zipfile.ZipFile(zip_path, "w", zipfile.ZIP_DEFLATED) as zipf: | |
| files_to_zip = json_files + csv_results_path_list + [log_kex_plot] + ae_histogram | |
| if (pdb_id_input or pdb_file_input) and dg_pdb_path: | |
| files_to_zip.append(dg_pdb_path) | |
| for file in files_to_zip: | |
| if file and os.path.exists(file): | |
| zipf.write(file, os.path.basename(file)) | |
| for plot_file in uptake_gallery: | |
| if os.path.exists(plot_file): | |
| zipf.write(plot_file, os.path.basename(plot_file)) | |
| zipf.write(summary_path, os.path.basename(summary_path)) | |
| uptake_visible = len(uptake_gallery) > 0 | |
| uptake_info_visible = not uptake_visible | |
| ae_message = None | |
| if not ae_histogram: | |
| ae_message = "No AE histograms generated. Enable 'Generate uptake plots' in Settings to create absolute error histograms." | |
| else: | |
| ae_histogram_labeled = [] | |
| for img_path in ae_histogram: | |
| basename = os.path.basename(img_path) | |
| if "centroid" in basename: | |
| envelope_type = "Centroid" | |
| elif "envelope" in basename: | |
| envelope_type = "Envelope" | |
| else: | |
| envelope_type = "Unknown" | |
| for state in state_names: | |
| if state in basename: | |
| title = f"{envelope_type} - {state}" | |
| break | |
| else: | |
| title = envelope_type | |
| ae_histogram_labeled.append((img_path, title)) | |
| return ( | |
| summary, | |
| log_kex_plot, | |
| molstar_view_html, | |
| uptake_gallery, | |
| uptake_info, | |
| ae_histogram_labeled, | |
| ae_info, | |
| heatmaps, | |
| zip_path | |
| ) | |
| def get_color_map(dg_pdb_path, vmin=None, vmax=None): | |
| white = (1.0, 1.0, 1.0) | |
| if len(state_names) == 1: | |
| orange_hex = '#F6851F' | |
| orange_rgb = mcolors.to_rgb(orange_hex) | |
| CMAP = mcolors.LinearSegmentedColormap.from_list('my_cm', [white, orange_rgb]) | |
| color_list = [mcolors.to_hex(CMAP(i/15)) for i in range(16)] | |
| mda_u = MDAnalysis.Universe(dg_pdb_path) | |
| if vmin is None and vmax is None: | |
| vmin, vmax = 0.0, np.ceil(np.nanmax(mda_u.residues.atoms.tempfactors)/5)*5 - 5 | |
| else: | |
| purple_hex = '#B162A7' | |
| green_hex = '#66BB45' | |
| purple_rgb = mcolors.to_rgb(purple_hex) | |
| green_rgb = mcolors.to_rgb(green_hex) | |
| CMAP = mcolors.LinearSegmentedColormap.from_list('my_cm', [green_rgb, white, purple_rgb]) | |
| color_list = [mcolors.to_hex(CMAP(i/15)) for i in range(16)] | |
| mda_u = MDAnalysis.Universe(dg_pdb_path) | |
| if vmin is None and vmax is None: | |
| vmax = np.ceil(np.nanmax(abs(mda_u.residues.atoms.tempfactors))/5)*5 - 5 | |
| vmin = -vmax | |
| protein = mda_u.select_atoms("protein") | |
| from collections import defaultdict | |
| color_map = defaultdict(dict) | |
| for res in protein.residues: | |
| bfactor = res.atoms.tempfactors[0] | |
| if res.resname == "PRO" or np.isnan(bfactor): | |
| color_hex = "#969696" | |
| else: | |
| norm_val = np.clip((bfactor - vmin) / (vmax - vmin), 0, 1) | |
| color_hex = mcolors.to_hex(CMAP(norm_val)) | |
| key = f"{res.atoms[0].chainID}_{res.resname}_{res.resid}" | |
| color_map[key] = int(color_hex.lstrip('#'), 16) | |
| print("Sample color map keys:", list(color_map.keys())[:10]) | |
| return color_map, color_list, vmin, vmax | |
| def get_structure_html(dg_pdb_path, title="", vmin=None, vmax=None): | |
| try: | |
| # Validate input file | |
| if not dg_pdb_path or not os.path.exists(dg_pdb_path): | |
| return "<div style='color: red; padding: 20px; text-align: center;'>Error: PDB file not found or invalid path.</div>" | |
| # Read PDB content with error handling | |
| try: | |
| with open(dg_pdb_path, 'r') as f: | |
| pdb_content = f.read() | |
| if pdb_content == "": | |
| return "<div style='color: #666; padding: 20px; text-align: center;'>No prediction data available for visualization.</div>" | |
| except IOError as e: | |
| return f"<div style='color: red; padding: 20px; text-align: center;'>Error reading PDB file: {str(e)}</div>" | |
| # Generate color mapping with error handling | |
| try: | |
| color_map, color_list, vmin, vmax = get_color_map(dg_pdb_path, vmin, vmax) | |
| except Exception as e: | |
| return f"<div style='color: red; padding: 20px; text-align: center;'>Error generating color mapping: {str(e)}</div>" | |
| # Prepare JSON data safely | |
| try: | |
| pdb_content_json = json.dumps(pdb_content) | |
| color_map_json = json.dumps(color_map) | |
| gradient_css = get_colorbar_css(color_list) | |
| except (TypeError, ValueError) as e: | |
| return f"<div style='color: red; padding: 20px; text-align: center;'>Error preparing visualization data: {str(e)}</div>" | |
| # Read Molstar CSS and JS files with error handling | |
| try: | |
| with open('static/molstar.css', 'r', encoding='utf-8') as f: | |
| molstar_css = f.read() | |
| except IOError: | |
| return "<div style='color: red; padding: 20px; text-align: center;'>Error: Molstar CSS file not found. Please ensure 'static/molstar.css' exists.</div>" | |
| try: | |
| with open('static/molstar.js', 'r', encoding='utf-8') as f: | |
| molstar_js = f.read() | |
| except IOError: | |
| return "<div style='color: red; padding: 20px; text-align: center;'>Error: Molstar JS file not found. Please ensure 'static/molstar.js' exists.</div>" | |
| # Generate legend title | |
| try: | |
| if len(state_names) == 1: | |
| legend_title = f"ΔG<sub>op,</sub> (kJ/mol)" | |
| else: | |
| state_names_str = [str(name) for name in state_names] | |
| display_state_name_0 = get_display_state_name(state_names_str[0], state_names_str) | |
| display_state_name_1 = get_display_state_name(state_names_str[1], state_names_str) | |
| legend_title = f"ΔΔG<sub>op,{display_state_name_0}-{display_state_name_1}</sub> (kJ/mol)" | |
| except (NameError, IndexError): | |
| legend_title = "ΔG (kJ/mol)" | |
| if len(state_names) == 1: | |
| hover_label = f"ΔG<sub>op" | |
| else: | |
| hover_label = f"ΔΔG<sub>op" | |
| hover_label_json = json.dumps(hover_label) | |
| # Generate HTML with error handling | |
| try: | |
| full_html = f"""<!DOCTYPE html> | |
| <html> | |
| <head> | |
| <meta charset="utf-8" /> | |
| <title>Mol* Viewer with Per-Residue Coloring</title> | |
| <style>{molstar_css}</style> | |
| <style> | |
| body {{ margin: 0; font-family: sans-serif; }} | |
| #viewer {{ width: 100vw; height: 100vh; }} | |
| .legend {{ position: absolute; left: 20px; bottom: 20px; z-index: 1000; background: rgba(255,255,255,0.95); padding: 12px 16px; border-radius: 8px; box-shadow: 0 2px 8px rgba(0,0,0,0.2); font-size: 12px; backdrop-filter: blur(4px); border: 1px solid rgba(0,0,0,0.1); min-width: 200px; }} | |
| .colorbar {{ width: 100%; height: 16px; border-radius: 4px; background: linear-gradient(to right, {gradient_css}); border: 1px solid rgba(0,0,0,0.1); margin: 8px 0; }} | |
| </style> | |
| </head> | |
| <body> | |
| <div id="viewer"></div> | |
| <div class="legend"> | |
| <div style="font-weight:600;margin-bottom:4px;">{legend_title}</div> | |
| <div style="width:180px;height:14px;border-radius:4px; | |
| background:linear-gradient(to right, {gradient_css});"></div> | |
| <div style="display:flex;justify-content:space-between;font-size:11px;margin-top:2px;"> | |
| <span>{vmin}</span><span>{vmax}</span> | |
| </div> | |
| <div style="margin-top:4px;font-size:11px;color:#555;">Gray: Proline/NaN</div> | |
| </div> | |
| <script>{molstar_js}</script> | |
| <script> | |
| const pdbStr = {pdb_content_json}; | |
| const colorMap = {color_map_json}; | |
| const hoverLabel = {hover_label_json}; | |
| const defaultColor = 0xcccccc; | |
| async function init() {{ | |
| try {{ | |
| const viewer = await molstar.Viewer.create('viewer', {{ | |
| layoutIsExpanded: false, | |
| layoutShowControls: false, | |
| layoutShowRemoteState: false, | |
| layoutShowSequence: true, | |
| layoutShowLog: false, | |
| layoutShowLeftPanel: true | |
| }}); | |
| const plugin = viewer.plugin; | |
| window.hoverLabelText = hoverLabel; | |
| window.occupancyLabelText = 'confidence'; | |
| plugin.representation.structure.themes.colorThemeRegistry.add({{ | |
| name: 'dg_color', | |
| label: 'ΔGop', | |
| category: 'Residue Property', | |
| factory: (ctx, props) => {{ | |
| return {{ | |
| color: (location) => {{ | |
| try {{ | |
| const unit = location.unit; | |
| const model = unit.model; | |
| const residueIndex = unit.residueIndex[location.element]; | |
| const chainIndex = unit.chainIndex[location.element]; | |
| const chainid = model.atomicHierarchy.chains.auth_asym_id.value(chainIndex); | |
| const firstAtomIndex = model.atomicHierarchy.residueAtomSegments.offsets[residueIndex]; | |
| const resname = model.atomicHierarchy.atoms.auth_comp_id.value(firstAtomIndex); | |
| const resid = model.atomicHierarchy.residues.auth_seq_id.value(residueIndex); | |
| const key = `${{chainid}}_${{resname}}_${{resid}}`; | |
| console.log("Looking for key:", key); | |
| if (colorMap[key] !== undefined) {{ | |
| const colorValue = colorMap[key]; | |
| if (isNaN(colorValue) || colorValue === null) {{ | |
| return 0xcccccc; | |
| }} | |
| return colorValue; | |
| }} | |
| }} catch (e) {{ | |
| console.log("Coloring error:", e); | |
| }} | |
| return defaultColor; | |
| }}, | |
| granularity: 'group', | |
| props: props, | |
| description: 'Color by ΔG' | |
| }}; | |
| }}, | |
| getParams: () => ({{}}), | |
| defaultValues: () => ({{}}), | |
| isApplicable: (ctx) => true | |
| }}); | |
| await viewer.loadStructureFromData(pdbStr, 'pdb'); | |
| const structureCell = plugin.managers.structure.hierarchy.current.structures[0]; | |
| const structure = structureCell?.cell?.obj?.data; | |
| if (structureCell && structure) {{ | |
| const model = structure.model; | |
| const deltaGopData = new Map(); | |
| const residues = new Set(); | |
| for (let i = 0; i < model.atomicHierarchy.atoms._rowCount; i++) {{ | |
| const elementIndex = i; | |
| const residueIndex = model.atomicHierarchy.residueAtomSegments.index[elementIndex]; | |
| const chainId = model.atomicHierarchy.chains.auth_asym_id.value(model.atomicHierarchy.chainAtomSegments.index[elementIndex]); | |
| const auth_seq_id = model.atomicHierarchy.residues.auth_seq_id.value(residueIndex); | |
| const atomName = model.atomicHierarchy.atoms.label_atom_id.value(elementIndex); | |
| if (atomName === 'CA') {{ | |
| const key = `${{chainId}}_${{auth_seq_id}}`; | |
| if (!residues.has(key)) {{ | |
| residues.add(key); | |
| const bfactor = model.atomicConformation.B_iso_or_equiv.value(elementIndex); | |
| const deltaGopValue = bfactor; | |
| console.log(key, deltaGopValue); | |
| deltaGopData.set(key, deltaGopValue); | |
| }} | |
| }} | |
| }} | |
| model._staticPropertyData['dGop'] = {{ | |
| value: deltaGopData, | |
| props: {{}} | |
| }}; | |
| }} | |
| if (structureCell && structure) {{ | |
| const components = structureCell.components; | |
| await plugin.managers.structure.component.updateRepresentationsTheme(components, {{ | |
| color: 'dg_color' | |
| }}); | |
| }} | |
| }} catch (error) {{ | |
| console.error('Molstar initialization failed:', error); | |
| document.getElementById('viewer').innerHTML = '<div style="padding: 20px; color: red; text-align: center;">Failed to initialize molecular viewer: ' + error.message + '</div>'; | |
| }} | |
| }} | |
| init(); | |
| </script> | |
| </body> | |
| </html>""" | |
| except Exception as e: | |
| return f"<div style='color: red; padding: 20px; text-align: center;'>Error generating HTML visualization: {str(e)}</div>" | |
| # Generate final iframe with error handling | |
| try: | |
| return f"<iframe style='width: 100%; height: 720px; border: none;' srcdoc='{html.escape(full_html)}'></iframe>" | |
| except Exception as e: | |
| return f"<div style='color: red; padding: 20px; text-align: center;'>Error creating iframe: {str(e)}</div>" | |
| except Exception as e: | |
| # Catch any unexpected errors in the entire function | |
| return f"<div style='color: red; padding: 20px; text-align: center;'>Unexpected error in structure visualization: {str(e)}<br><br>Traceback: {html.escape(traceback.format_exc())}</div>" | |
| def get_summary(output_dicts, hdxms_data_list): | |
| SUMMARY = "" | |
| for idx, state_name in enumerate(state_names): | |
| results = output_dicts[f"{state_name}_{idx}"]["results"] | |
| stats_text = get_all_statics_info([hdxms_data_list[idx]]) | |
| SUMMARY += stats_text | |
| SUMMARY += "\n"*2 | |
| PFNET_HEADER = "=" * 60 + "\n" + " " * 20 + f"PFNet Results Summary for {state_name} \n" + "=" * 60 + "\n" | |
| SUMMARY += PFNET_HEADER | |
| single_num = sum(np.array(results['resolution_limits'])[:,0] == np.array(results['resolution_limits'])[:,1]) | |
| non_covered_residues = (torch.tensor(results['resolution_grouping']) == torch.tensor([1, 0, 0, 0])).all(axis=1) | |
| non_covered_residues_num = non_covered_residues.sum().item() | |
| covered_residues_num = len(results['protein_sequence']) - non_covered_residues_num | |
| seq_mask = results["seq_mask"] | |
| highest_log_kex = torch.max(results["pfnet_pred_log_kex"][~seq_mask]) | |
| lowest_log_kex = torch.min(results["pfnet_pred_log_kex"][~seq_mask]) | |
| std_log_kex = torch.std(results["pfnet_pred_log_kex"][~seq_mask]) | |
| high_confidence_residues = results["pfnet_pred_log_kex_confidence"] > 0.8 | |
| high_confidence_residues_num = high_confidence_residues.sum().item() | |
| text = f"Highest log(kex) (PFNet): {highest_log_kex:.2f}\nLowest log(kex) (PFNet): {lowest_log_kex:.2f}\n" | |
| text += f"Std log(kex) (PFNet): {std_log_kex:.2f}\n" | |
| SUMMARY += text + "\n" | |
| SUMMARY += f"Number of single resolved residues: {single_num}\n" | |
| SUMMARY += f"Number of non covered residues: {non_covered_residues_num}\n" | |
| SUMMARY += f"Number of high confidence residues: {high_confidence_residues_num} ({high_confidence_residues_num/covered_residues_num*100:.2f}%)\n" | |
| SUMMARY += f"AE_mean_pfnet: {results['AE_mean_pfnet']:.2f}\n" | |
| SUMMARY += f"AE_median_pfnet: {results['AE_median_pfnet']:.2f}\n" | |
| SUMMARY += f"centroid model: {results['centroid_model']}\n" | |
| SUMMARY += "=" * 60 + "\n"*2 | |
| summary_path = f"{output_dir}/pfnet_plots/PFNet_summary.txt" | |
| with open(summary_path, 'w') as f: | |
| f.write(SUMMARY) | |
| return SUMMARY, summary_path | |
| def get_all_statics_info(hdxms_datas): | |
| from pigeon_feather.tools import calculate_coverages | |
| # Ensure input is a list for consistent processing | |
| if not isinstance(hdxms_datas, list): | |
| hdxms_datas = [hdxms_datas] | |
| state_names = list(set([state.state_name for data in hdxms_datas for state in data.states])) | |
| protein_sequence = hdxms_datas[0].protein_sequence | |
| # Calculate coverage statistics | |
| coverage = [calculate_coverages([data], state.state_name) | |
| for data in hdxms_datas for state in data.states] | |
| coverage = np.mean(np.array(coverage), axis=0) | |
| coverage_non_zero = 1 - np.count_nonzero(coverage == 0) / len(protein_sequence) | |
| # Gather peptides and calculate statistics | |
| all_peptides = [pep for data in hdxms_datas for state in data.states for pep in state.peptides] | |
| unique_peptides = set(pep.identifier for pep in all_peptides) | |
| avg_pep_length = np.mean([len(pep.sequence) for pep in all_peptides]) | |
| # all tps | |
| all_tps = [tp for pep in all_peptides for tp in pep.timepoints if tp.deut_time != np.inf and tp.deut_time != 0.0] | |
| time_course = sorted(list(set([tp.deut_time for tp in all_tps]))) | |
| def _group_and_average(numbers, threshold=50): | |
| numbers.sort() | |
| groups, current_group = [], [numbers[0]] | |
| for number in numbers[1:]: | |
| (current_group.append(number) if number - current_group[0] <= threshold else (groups.append(current_group), current_group := [number])) | |
| groups.append(current_group) | |
| return groups, [round(sum(group) / len(group), 1) for group in groups] | |
| groups, avg_timepoints = _group_and_average(time_course) | |
| # Calculate back exchange rates and IQR | |
| peptides_with_exp = [pep for pep in all_peptides if pep.get_timepoint(np.inf) is not None] | |
| backexchange_rates = [1 - pep.max_d / pep.theo_max_d for pep in peptides_with_exp] | |
| if backexchange_rates == []: | |
| iqr_backexchange = np.nan | |
| else: | |
| iqr_backexchange = np.percentile(backexchange_rates, 75) - np.percentile(backexchange_rates, 25) | |
| # Calculate redundancy based on coverage | |
| redundancy = np.mean(coverage) | |
| # Print formatted output | |
| stats_text = ( | |
| "=" * 60 + "\n" + | |
| " " * 20 + "HDX-MS Data Statistics\n" + | |
| "=" * 60 + "\n" + | |
| f"States names: {state_names}\n" + | |
| f"Time course (s): {avg_timepoints}\n" + | |
| f"Number of time points: {len(avg_timepoints)}\n" + | |
| f"Protein sequence length: {len(protein_sequence)}\n" + | |
| f"Average coverage: {coverage_non_zero:.2f}\n" + | |
| f"Number of unique peptides: {len(unique_peptides)}\n" + | |
| f"Average peptide length: {avg_pep_length:.1f}\n" + | |
| f"Redundancy (based on average coverage): {redundancy:.1f}\n" + | |
| f"Average peptide length to redundancy ratio: {avg_pep_length / redundancy:.1f}\n" + | |
| f"Backexchange average, IQR: {np.mean(backexchange_rates):.2f}, {iqr_backexchange:.2f}\n" + | |
| "=" * 60 | |
| ) | |
| return stats_text | |
| def get_log_kex_plot(ana_objs, output_dir): | |
| first_key = list(ana_objs.keys())[0] | |
| seq_len = len(ana_objs[first_key].protein_sequence) | |
| num_len = int(np.ceil(seq_len / 150)) | |
| fig, ax = plt.subplots(1, 1, figsize=(40*num_len, 8), sharey=True, sharex=True) | |
| state_keys = list(ana_objs.keys()) | |
| for idx, (state_key, ana_obj) in enumerate(ana_objs.items()): | |
| state_name = get_display_state_name(state_key, state_keys) | |
| ana_obj.plot_kex_bar( | |
| ax=ax, resolution_indicator_pos=15-idx, label=state_name, show_seq=False, | |
| ) | |
| ax.set_xlabel("Residue", fontsize=24) | |
| spacing = int(seq_len // 150 * 1 + 1) | |
| ax.set_xticks(ax.get_xticks()[::spacing]) | |
| ax.set_xticklabels(ax.get_xticklabels(), fontdict={"fontsize": 24}) | |
| seq_pos = 17 | |
| for ii in range(0, seq_len, spacing): | |
| ax.text(ii, seq_pos, ana_objs[first_key].protein_sequence[ii], ha="center", va="center", fontsize=22) | |
| from matplotlib.colors import Normalize | |
| from matplotlib import cm | |
| coverage_max = np.nanmax(ana_objs[first_key].coverage) | |
| norm = Normalize(vmin=0, vmax=coverage_max) | |
| sm = cm.ScalarMappable(cmap=plt.cm.Blues, norm=norm) | |
| sm.set_array([]) | |
| cbar_width_inch = 100 / 72 | |
| cbar_height_inch = 20 / 72 | |
| fig_width_inch, fig_height_inch = fig.get_size_inches() | |
| last_residue_pos = seq_len - 1 | |
| ax_pos = ax.get_position() | |
| xlim = ax.get_xlim() | |
| last_residue_fig_pos = ax_pos.x0 + (last_residue_pos - xlim[0]) / (xlim[1] - xlim[0]) * ax_pos.width | |
| x_pos = last_residue_fig_pos - (cbar_width_inch / fig_width_inch) | |
| cbar_ax = fig.add_axes([ | |
| x_pos, | |
| 0.7, | |
| cbar_width_inch / fig_width_inch, | |
| cbar_height_inch / fig_height_inch | |
| ]) | |
| cbar = fig.colorbar(sm, cax=cbar_ax, orientation='horizontal') | |
| cbar.set_label('Coverage', fontsize=18, rotation=0) | |
| cbar.set_ticks([0, coverage_max]) | |
| cbar.set_ticklabels(['0', f'{int(coverage_max)}']) | |
| cbar.ax.tick_params(labelsize=18) | |
| cbar.outline.set_visible(False) | |
| ax.legend(loc='upper left', bbox_to_anchor=(0.01, 0.7)) | |
| fig.savefig(f"{output_dir}/pfnet_plots/log_kex_plot.png") | |
| return f"{output_dir}/pfnet_plots/log_kex_plot.png" | |
| def get_csv_results(ana_objs, output_dicts, output_dir): | |
| csv_path_list = [] | |
| for idx, (state_key, ana_obj) in enumerate(ana_objs.items()): | |
| df_logPF = create_logP_df(ana_obj, 0, output_dicts[state_key]["results"]["pfnet_pred_log_kex_confidence"]) | |
| csv_path = f"{output_dir}/pfnet_output/results_{state_key}.csv" | |
| df_logPF.to_csv(csv_path, index=False) | |
| csv_path_list.append(csv_path) | |
| return csv_path_list | |
| def get_heatmap(hdxms_data_list, output_dir): | |
| all_state_names = [state.state_name for data in hdxms_data_list for state in data.states] | |
| if len(hdxms_data_list) == 1: | |
| fig = create_heatmap_single_state(hdxms_data_list, colorbar_max=80) | |
| fig.savefig(f"{output_dir}/pfnet_plots/heatmap_{all_state_names[0]}.png") | |
| else: | |
| from itertools import product | |
| from pigeon_feather.data import HDXStatePeptideCompares | |
| if len(hdxms_data_list) >= 2: | |
| state1_name = hdxms_data_list[0].states[0].state_name | |
| state2_name = hdxms_data_list[1].states[0].state_name | |
| state1_list = [hdxms_data_list[0].states[0]] | |
| state2_list = [hdxms_data_list[1].states[0]] | |
| compare = HDXStatePeptideCompares(state1_list, state2_list) | |
| compare.add_all_compare() | |
| if len(compare.peptide_compares) > 0: | |
| heatmap_compare = create_heatmap_compare(compare, 20) | |
| heatmap_compare.savefig(f'{output_dir}/pfnet_plots/heatmap_{state1_name}_{state2_name}.png') | |
| else: | |
| print("Warning: No valid peptide comparisons found for heatmap generation") | |
| heatmaps = glob.glob(f"{output_dir}/pfnet_plots/heatmap_*.png") | |
| return heatmaps | |
| def create_heatmap_compare(compare, colorbar_max, colormap="RdBu"): | |
| import matplotlib.colors as col | |
| from matplotlib.patches import Rectangle | |
| from matplotlib import cm | |
| import matplotlib.pyplot as plt | |
| from matplotlib import rcParams | |
| if not compare.peptide_compares or len(compare.peptide_compares) == 0: | |
| fig, ax = plt.subplots(figsize=(20, 10)) | |
| ax.text(0.5, 0.5, 'No peptide comparison data available', | |
| ha='center', va='center', transform=ax.transAxes, fontsize=16) | |
| ax.set_xlim(0, 1) | |
| ax.set_ylim(0, 1) | |
| plt.close() | |
| return fig | |
| if not compare.peptide_compares[0].peptide1_list or len(compare.peptide_compares[0].peptide1_list) == 0: | |
| fig, ax = plt.subplots(figsize=(20, 10)) | |
| ax.text(0.5, 0.5, 'No peptide data available for comparison', | |
| ha='center', va='center', transform=ax.transAxes, fontsize=16) | |
| ax.set_xlim(0, 1) | |
| ax.set_ylim(0, 1) | |
| plt.close() | |
| return fig | |
| with plt.style.context('default'): | |
| font_config = {"family": "Arial", "weight": "normal", "size": 14} | |
| axes_config = {"titlesize": 18, "titleweight": "bold", "labelsize": 16} | |
| fig, ax = plt.subplots(figsize=(20, 10)) | |
| ax.tick_params(labelsize=font_config["size"]) | |
| ax.set_title( | |
| compare.state1_list[0].state_name + "-" + compare.state2_list[0].state_name, | |
| fontsize=axes_config["titlesize"], | |
| fontweight=axes_config["titleweight"], | |
| fontfamily=font_config["family"] | |
| ) | |
| colormap = cm.get_cmap(colormap) | |
| leftbound = compare.peptide_compares[0].peptide1_list[0].start - 10 | |
| rightbound = compare.peptide_compares[-1].peptide1_list[0].end + 10 | |
| ax.set_xlim(leftbound, rightbound) | |
| ax.xaxis.set_ticks(np.arange(round(leftbound, -1), round(rightbound, -1), 10)) | |
| ax.set_ylim(-5, 110) | |
| ax.grid(axis="x") | |
| ax.yaxis.set_ticks([]) | |
| norm = col.Normalize(vmin=-colorbar_max, vmax=colorbar_max) | |
| for i, peptide_compare in enumerate(compare.peptide_compares): | |
| for peptide in peptide_compare.peptide1_list: | |
| rect = Rectangle( | |
| (peptide.start, (i % 20) * 5 + ((i // 20) % 2) * 2.5), | |
| peptide.end - peptide.start, | |
| 4, | |
| fc=colormap(norm(peptide_compare.deut_diff_avg)), | |
| ) | |
| ax.add_patch(rect) | |
| cbar = fig.colorbar(cm.ScalarMappable(cmap=colormap, norm=norm), ax=ax) | |
| cbar.ax.tick_params(labelsize=axes_config["labelsize"]) | |
| cbar.set_label('Deuteration difference (%)', fontsize=18, fontfamily="Arial") | |
| ax.set_xlabel('Residue', fontsize=18, fontfamily="Arial") | |
| fig.tight_layout() | |
| plt.close() | |
| return fig | |
| def create_heatmap_single_state(hdxms_datas, colorbar_max, colormap="Greens"): | |
| import matplotlib.colors as col | |
| from matplotlib.patches import Rectangle | |
| from matplotlib import colormaps | |
| import matplotlib.patches as patches | |
| import matplotlib.pyplot as plt | |
| import seaborn as sns | |
| state_name = list( | |
| set([state.state_name for data in hdxms_datas for state in data.states]) | |
| ) | |
| if len(state_name) > 1: | |
| raise ValueError("More than one state name found") | |
| else: | |
| state_name = state_name[0] | |
| with plt.style.context('default'): | |
| font_config = {"family": "Arial", "weight": "normal", "size": 14} | |
| axes_config = {"titlesize": 18, "titleweight": "bold", "labelsize": 16} | |
| fig, ax = plt.subplots(1, 1, figsize=(20, 10)) | |
| ax.tick_params(labelsize=font_config["size"]) | |
| ax.set_title( | |
| state_name, | |
| fontsize=axes_config["titlesize"], | |
| fontweight=axes_config["titleweight"], | |
| fontfamily=font_config["family"] | |
| ) | |
| colormap = sns.light_palette("seagreen", as_cmap=True) | |
| all_peptides = [ | |
| pep for data in hdxms_datas for state in data.states for pep in state.peptides | |
| ] | |
| all_peptides.sort(key=lambda x: x.start) | |
| leftbound = all_peptides[0].start - 10 | |
| rightbound = all_peptides[-1].end + 10 | |
| ax.set_xlim(leftbound, rightbound) | |
| ax.xaxis.set_ticks(np.arange(round(leftbound, -1), round(rightbound, -1), 10)) | |
| ax.set_ylim(-5, 110) | |
| ax.grid(axis="x") | |
| ax.yaxis.set_ticks([]) | |
| norm = col.Normalize(vmin=0, vmax=colorbar_max) | |
| for i, peptide in enumerate(all_peptides): | |
| avg_d_percent = np.average( | |
| [tp.d_percent for tp in peptide.timepoints if tp.deut_time != np.inf] | |
| ) | |
| rect = Rectangle( | |
| (peptide.start, (i % 20) * 5 + ((i // 20) % 2) * 2.5), | |
| peptide.end - peptide.start, | |
| 4, | |
| fc=colormap(norm(avg_d_percent)), | |
| ) | |
| ax.add_patch(rect) | |
| # coverage | |
| coverage = np.zeros(len(hdxms_datas[0].states[0].hdxms_data.protein_sequence)) | |
| for pep in all_peptides: | |
| coverage[pep.start - 1 : pep.end] += 1 | |
| height = 3 | |
| for i in range(len(coverage)): | |
| color_intensity = ( | |
| coverage[i] / 20 | |
| ) # coverage.max() # Normalizing the data for color intensity | |
| rect = patches.Rectangle( | |
| (i, 105), 1, height, color=plt.cm.Blues(color_intensity) | |
| ) | |
| ax.add_patch(rect) | |
| from matplotlib import cm | |
| cbar = fig.colorbar(cm.ScalarMappable(cmap=colormap, norm=norm), ax=ax) | |
| cbar.ax.tick_params(labelsize=axes_config["labelsize"]) | |
| cbar.set_label('Deuteration (%)', fontsize=18, fontfamily="Arial") | |
| ax.set_xlabel('Residue', fontsize=18, fontfamily="Arial") | |
| fig.tight_layout() | |
| plt.close() | |
| return fig | |
| def make_BFactorPlot(pdb_file, ana_objs, output_dir): | |
| state_keys = list(ana_objs.keys()) | |
| if len(state_names) == 1: | |
| first_ana_obj = list(ana_objs.values())[0] | |
| bfactor_plot = BFactorPlot( | |
| first_ana_obj, | |
| pdb_file=pdb_file, | |
| plot_deltaG=True, | |
| temperature=float(first_ana_obj.temperature), | |
| ) | |
| dg_pdb_path = f"{output_dir}/pfnet_plots/PFNet_dG.pdb" | |
| bfactor_plot.plot(dg_pdb_path) | |
| else: | |
| ana_obj_list = list(ana_objs.values()) | |
| bfactor_plot = BFactorPlot( | |
| ana_obj_list[0], | |
| ana_obj_list[1], | |
| pdb_file=pdb_file, | |
| plot_deltaG=True, | |
| temperature=float(ana_obj_list[0].temperature), | |
| ) | |
| file_state_name_0 = get_display_state_name(state_keys[0], state_keys) | |
| file_state_name_1 = get_display_state_name(state_keys[1], state_keys) | |
| dg_pdb_path = f"{output_dir}/pfnet_plots/PFNet_ddG_{file_state_name_0}-{file_state_name_1}.pdb" | |
| bfactor_plot.plot(dg_pdb_path) | |
| return dg_pdb_path | |
| def logPF_to_deltaG(ana_obj, logPF): | |
| """ | |
| :param logPF: logP value | |
| :return: deltaG in kJ/mol, local unfolding energy | |
| """ | |
| return 8.3145 * ana_obj.temperature * np.log(10) * logPF / 1000 | |
| def create_logP_df(ana_obj, index_offset, pfnet_confidence): | |
| df_logPF = pd.DataFrame() | |
| for res_i, _ in enumerate(ana_obj.results_obj.protein_sequence): | |
| res_obj_i = ana_obj.results_obj.get_residue_by_resindex(res_i) | |
| avg_logP, std_logP, SE_logP = get_res_avg_logP(res_obj_i) | |
| log_kex = get_res_avg_log_kex(res_obj_i) * -1 | |
| df_i = pd.DataFrame( | |
| { | |
| "resid": [res_obj_i.resid - index_offset], | |
| "resname": [res_obj_i.resname], | |
| 'avg_dG (kJ/mol)': [round(logPF_to_deltaG(ana_obj, avg_logP), 3)], | |
| 'std_dG (kJ/mol)': [round(logPF_to_deltaG(ana_obj, std_logP), 3)], | |
| "avg_logP (log(sec^-1))": [round(avg_logP, 3)], | |
| "std_logP (log(sec^-1))": [round(std_logP, 3)], | |
| "log_kch (log(sec^-1))": [round(res_obj_i.log_k_init, 3)], | |
| "log_kex (log(sec^-1))": [round(log_kex, 3)], | |
| "is_nan": [res_obj_i.is_nan()], | |
| "coverage": [ana_obj.coverage[res_i]], | |
| "PFNet_confidence": [float(pfnet_confidence[res_i])] | |
| } | |
| ) | |
| if res_obj_i.is_nan(): | |
| df_i["single_resolved"] = [np.nan] | |
| df_i["min_pep logPs"] = [np.nan] | |
| df_i["min_pep log_kex"] = [np.nan] | |
| else: | |
| df_i["single_resolved"] = [res_obj_i.mini_pep.if_single_residue()] | |
| df_i["min_pep logPs"] = [[round(i, 3) for i in res_obj_i.clustering_results_logP]] | |
| df_i["min_pep log_kex"] = [-1*np.round(res_obj_i.mini_pep.clustering_results_log_kex, 3)] | |
| df_logPF = pd.concat([df_logPF, df_i]) | |
| df_logPF['is_nan'] = df_logPF['is_nan'].astype(bool) | |
| df_logPF['single_resolved'] = df_logPF['single_resolved'].astype(bool) | |
| df_logPF = df_logPF.reset_index(drop=True) | |
| return df_logPF | |
| def visualize_structure(pdb_data, pf_data=None): | |
| view = py3Dmol.view(width=700, height=500) | |
| view.addModel(pdb_data, 'pdb') | |
| view.setStyle({'cartoon': {'color': 'gray'}}) | |
| if pf_data is not None and not pf_data.empty: | |
| for _, row in pf_data.iterrows(): | |
| residue_index = row['residue_index'] | |
| pf_value = process_pf_value(row['log10_pf_value']) | |
| if np.isnan(pf_value): | |
| continue | |
| if pf_value < 2: | |
| color = 'blue' | |
| elif pf_value < 4: | |
| color = 'green' | |
| else: | |
| color = 'red' | |
| view.addStyle({'resi': str(residue_index)}, | |
| {'cartoon': {'color': color}, 'stick': {}}) | |
| view.zoomTo() | |
| return view | |
| def load_demo_data_mini_protein(): | |
| current_dir = os.path.dirname(os.path.abspath(__file__)) | |
| hxms_path = os.path.join(current_dir, "demo_data", "EEHEEEEHEE_rd4_0871.hxms") | |
| pdb_path = os.path.join(current_dir, "demo_data", "EEHEEEEHEE_rd4_0871.pdb") | |
| return hxms_path, None, pdb_path | |
| def load_demo_data_ecDHFR(): | |
| current_dir = os.path.dirname(os.path.abspath(__file__)) | |
| hxms_path1 = os.path.join(current_dir, "demo_data", "ecDHFR_APO.hxms") | |
| hxms_path2 = os.path.join(current_dir, "demo_data", "ecDHFR_MTX.hxms") | |
| pdb_path = os.path.join(current_dir, "demo_data", "1RG7_MTX.pdb") | |
| return hxms_path1, hxms_path2, pdb_path | |
| def load_pdb_data(pdb_id, pdb_file): | |
| """ | |
| Loads a PDB file either by ID (download) or from an uploaded file. | |
| Returns the path to the PDB file. | |
| """ | |
| pdb_path = None | |
| if pdb_id: | |
| try: | |
| pdb_url = f"https://files.rcsb.org/download/{pdb_id}.pdb" | |
| temp_dir = tempfile.mkdtemp() | |
| pdb_path = os.path.join(temp_dir, f"{pdb_id}.pdb") | |
| urllib.request.urlretrieve(pdb_url, pdb_path) | |
| print(f"Downloaded PDB from {pdb_url} to {pdb_path}") | |
| except Exception as e: | |
| print(f"Error downloading PDB ID {pdb_id}: {e}") | |
| traceback.print_exc() | |
| return None | |
| elif pdb_file: | |
| pdb_path = pdb_file.name | |
| print(f"Using uploaded PDB file: {pdb_path}") | |
| return pdb_path | |
| def get_colorbar_css(color_list): | |
| n = len(color_list) | |
| color_stops = [] | |
| for i, c in enumerate(color_list): | |
| pct = int(i / (n - 1) * 100) | |
| color_stops.append(f"{c} {pct}%") | |
| return ", ".join(color_stops) | |
| GA_ID = os.environ.get("GOOGLE_ANALYTICS_ID", "") | |
| ga_script = "" | |
| if GA_ID: | |
| ga_script = f""" | |
| <script async src="https://www.googletagmanager.com/gtag/js?id={GA_ID}"></script> | |
| <script> | |
| window.dataLayer = window.dataLayer || []; | |
| function gtag(){{dataLayer.push(arguments);}} | |
| gtag('js', new Date()); | |
| gtag('config', '{GA_ID}'); | |
| </script> | |
| """ | |
| with gr.Blocks( | |
| css=""" | |
| """, | |
| head=ga_script | |
| ) as demo: | |
| gr.Markdown("# PFNet") | |
| gr.Markdown( | |
| '[PFNet](https://github.com/glasgowlab/PFNet) takes an [HXMS file](https://huggingface.co/spaces/glasgow-lab/PFLink) as input and predicts ΔG<sub>op</sub> for protein residues.' | |
| ) | |
| with gr.Tabs(): | |
| height = 150 | |
| with gr.TabItem("Input"): | |
| with gr.Row(equal_height=True): | |
| with gr.Column(scale=2): | |
| file1 = gr.File(label="First HXMS file", file_types=[".hxms"], height=height) | |
| file2 = gr.File(label="Second HXMS file (optional)", file_types=[".hxms"], scale=1, height=height) | |
| with gr.Column(scale=1): | |
| pdb_id_input = gr.Textbox(label="PDB ID (Optional)", placeholder="e.g., 1A2B", scale=1) | |
| pdb_file_input = gr.File(label="Upload PDB (Optional)", scale=1, height=height) | |
| with gr.Column(scale=1): | |
| demo_btn1 = gr.Button("Load demo data (mini_protein)",) | |
| demo_btn2 = gr.Button("Load demo data (ecDHFR)") | |
| gr.Markdown("### Settings") | |
| with gr.Row(equal_height=True): | |
| with gr.Column(scale=1, min_width=220): | |
| centroid_model = gr.Checkbox(label="Use centroid model (only recommended if you are missing envelope data)", value=False) | |
| uptake_plots = gr.Checkbox(label="Generate uptake plots", value=False) | |
| with gr.Column(scale=1, min_width=220): | |
| with gr.Row(): | |
| vmin = gr.Number( | |
| value=0, | |
| label="Structure color min", | |
| ) | |
| vmax = gr.Number( | |
| value=50, | |
| label="Structure color max", | |
| ) | |
| btn = gr.Button("Run prediction") | |
| gr.Markdown("# Output") | |
| with gr.Tabs(): | |
| with gr.TabItem("Summary"): | |
| summary = gr.Textbox(label="Results summary", lines=30, autoscroll=True) | |
| with gr.TabItem("log(kₑₓ) plot"): | |
| log_kex_plot = gr.Image(label="log(kₑₓ) plot") | |
| with gr.TabItem("Molstar viewer"): | |
| molstar_view = gr.HTML(label="Molstar viewer") | |
| with gr.TabItem("Uptake plots"): | |
| uptake_gallery = gr.File( | |
| label="Download uptake plots (PDF)", | |
| file_count="multiple", | |
| visible=True | |
| ) | |
| uptake_info = gr.Textbox( | |
| label="Info", | |
| value="", | |
| visible=True, | |
| interactive=False | |
| ) | |
| with gr.TabItem("AE histogram"): | |
| ae_histogram = gr.Gallery( | |
| label="Absolute error histogram", | |
| show_label=True, | |
| allow_preview=True, | |
| columns=2, | |
| height="auto", | |
| object_fit="contain", | |
| visible=True | |
| ) | |
| ae_info = gr.Textbox( | |
| label="Info", | |
| value="", | |
| visible=True, | |
| interactive=False | |
| ) | |
| with gr.TabItem("Heatmaps"): | |
| heatmaps = gr.Gallery(label="Heatmaps", show_label=True, allow_preview=True, columns=2, height="auto", object_fit="contain") | |
| with gr.TabItem("Download results (ZIP)"): | |
| zip_file = gr.File(label="Download all results as a ZIP file") | |
| btn.click( | |
| fn=predict_gradio, | |
| inputs=[ | |
| file1, | |
| file2, | |
| pdb_id_input, | |
| pdb_file_input, | |
| centroid_model, | |
| uptake_plots, | |
| vmin, | |
| vmax, | |
| ], | |
| outputs=[summary, log_kex_plot, molstar_view, uptake_gallery, uptake_info, ae_histogram, ae_info, heatmaps, zip_file], | |
| ) | |
| file1.change( | |
| fn=get_default_structure_color_range, | |
| inputs=[file1, file2], | |
| outputs=[vmin, vmax], | |
| ) | |
| file2.change( | |
| fn=get_default_structure_color_range, | |
| inputs=[file1, file2], | |
| outputs=[vmin, vmax], | |
| ) | |
| demo_btn1.click( | |
| fn=load_demo_data_mini_protein, | |
| inputs=[], | |
| outputs=[file1, file2, pdb_file_input], | |
| ).then( | |
| fn=get_default_structure_color_range, | |
| inputs=[file1, file2], | |
| outputs=[vmin, vmax], | |
| ) | |
| demo_btn2.click( | |
| fn=load_demo_data_ecDHFR, | |
| inputs=[], | |
| outputs=[file1, file2, pdb_file_input], | |
| ).then( | |
| fn=get_default_structure_color_range, | |
| inputs=[file1, file2], | |
| outputs=[vmin, vmax], | |
| ) | |
| demo.launch( | |
| share=True, | |
| debug=False, | |
| show_error=True, | |
| ) | |