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 "
Error: PDB file not found or invalid path.
" # Read PDB content with error handling try: with open(dg_pdb_path, 'r') as f: pdb_content = f.read() if pdb_content == "": return "
No prediction data available for visualization.
" except IOError as e: return f"
Error reading PDB file: {str(e)}
" # 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"
Error generating color mapping: {str(e)}
" # 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"
Error preparing visualization data: {str(e)}
" # 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 "
Error: Molstar CSS file not found. Please ensure 'static/molstar.css' exists.
" try: with open('static/molstar.js', 'r', encoding='utf-8') as f: molstar_js = f.read() except IOError: return "
Error: Molstar JS file not found. Please ensure 'static/molstar.js' exists.
" # Generate legend title try: if len(state_names) == 1: legend_title = f"ΔGop, (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"ΔΔGop,{display_state_name_0}-{display_state_name_1} (kJ/mol)" except (NameError, IndexError): legend_title = "ΔG (kJ/mol)" if len(state_names) == 1: hover_label = f"ΔGop" else: hover_label = f"ΔΔGop" hover_label_json = json.dumps(hover_label) # Generate HTML with error handling try: full_html = f""" Mol* Viewer with Per-Residue Coloring
{legend_title}
{vmin}{vmax}
Gray: Proline/NaN
""" except Exception as e: return f"
Error generating HTML visualization: {str(e)}
" # Generate final iframe with error handling try: return f"" except Exception as e: return f"
Error creating iframe: {str(e)}
" except Exception as e: # Catch any unexpected errors in the entire function return f"
Unexpected error in structure visualization: {str(e)}

Traceback: {html.escape(traceback.format_exc())}
" 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""" """ 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 ΔGop 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, )