PFNet / app.py
lucl13
bug fix due to the new HXMS io func
f943630
Raw History Blame Contribute Delete
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,
)