OneScience-Group/flex_ddG_tutorial
027
1#!/usr/bin/env python32"""3Recover per-chain intramolecular energies from an EXISTING flex ddG run.4 5No re-sampling is required. The backrub trajectory is the expensive part and you already ran6it: struct.db3 holds the coordinates of every backrub, wild type minimized and mutant minimized7pose. This script reads those poses back into Rosetta, isolates one chain at a time, rescores,8and reports per-chain ddG. Verified to reproduce the in-protocol per-chain metric (see9per_chain_protocol.py) to nine decimal places.10 11Why the chains have to be physically isolated rather than just selected: struct.db3 stores the12*bound* poses. A Chain residue selector scoped over a bound pose still picks up cross-chain pair13energies, because Rosetta's residue_total_energies splits each two-body term between its two14residues, so roughly half the interface energy leaks into each chain. Deleting the other chains15reproduces the separated unbound state exactly, since intra-chain energy is invariant under the16rigid-body translation InterfaceDdGMover uses to unbind.17 18What this does and does not give you19------------------------------------20It gives you the intramolecular energy difference between the mutant and wild type in the21*bound* backbone conformation, per chain. That is a strain term. It is NOT a folding ddG of the22free monomer: the unbound state is never relaxed here, in this script or in flex ddG itself. For23a true monomer stability ddG use a dedicated protocol such as cartesian_ddg on the isolated chain.24 25Built-in control: a chain you did not mutate should come out at 0 within its SEM. If it does not,26nstruct is too low to average out the whole-pose minimization noise, and the mutated chain's27number is not trustworthy either.28 29Usage30-----31 python3 reprocess_per_chain.py <output_folder> [--stride N] [--chains A,B] [--csv out.csv]32 33The backrub_trajectory_stride is read back out of each struct.db3, so --stride is only needed34for databases that do not record it. It affects the checkpoint labels, not any energy.35"""36 37import argparse38import glob39import os40import re41import sqlite342import subprocess43import sys44 45import numpy as np46import pandas as pd47 48import flex_ddg_db349 50rosetta_scripts_path = os.path.expanduser('~/rosetta/source/bin/rosetta_scripts')51 52# Must match the score function the original run used.53rosetta_flags = [54 '-restore_talaris_behavior',55 '-in:file:fullatom',56 '-out:nooutput',57]58 59struct_db3_name = 'struct.db3'60per_chain_db3_name = 'per_chain.db3'61 62# flex ddG writes three poses per checkpoint, in this order.63pose_order = ['backrub', 'wt', 'mut']64 65 66def chains_in_struct_db3(struct_db3):67 conn = sqlite3.connect(struct_db3)68 try:69 chains = [row[0] for row in conn.execute(70 'SELECT DISTINCT chain_id FROM residue_pdb_identification ORDER BY chain_id')]71 finally:72 conn.close()73 return [c for c in chains if c and c.strip()]74 75 76def write_rescore_protocol(chains, out_path, scorefxn='fa_talaris2014'):77 """Emit a RosettaScripts protocol that isolates and rescores each chain in turn."""78 lowered = [c.lower() for c in chains]79 if len(set(lowered)) != len(lowered):80 raise ValueError('Chain IDs differ only by case, which collides in batch names: %s' % chains)81 82 selectors = '\n'.join(83 ' <Chain name="chain_%s" chains="%s"/>\n'84 ' <Not name="not_chain_%s" selector="chain_%s"/>' % (c, c, c, c) for c in chains)85 86 movers = '\n'.join(87 ' <DeleteRegionMover name="isolate_chain_%s" residue_selector="not_chain_%s"/>\n'88 ' <ReportToDB name="chain_%s_report" batch_description="per_chain" database_name="%s">\n'89 ' <ScoreTypeFeatures/>\n'90 ' <ScoreFunctionFeatures scorefxn="%s"/>\n'91 ' <StructureScoresFeatures scorefxn="%s"/>\n'92 ' </ReportToDB>' % (c, c, c, per_chain_db3_name, scorefxn, scorefxn) for c in chains)93 94 steps = [' <Add mover_name="save_full"/>']95 for i, c in enumerate(chains):96 if i > 0:97 steps.append(' <Add mover_name="restore_full"/>')98 steps.append(' <Add mover_name="isolate_chain_%s"/>' % c)99 steps.append(' <Add mover_name="chain_%s_report"/>' % c)100 101 xml = '''<ROSETTASCRIPTS>102 <SCOREFXNS>103 <ScoreFunction name="%s" weights="talaris2014"/>104 </SCOREFXNS>105 106 <RESIDUE_SELECTORS>107%s108 </RESIDUE_SELECTORS>109 110 <MOVERS>111 <SavePoseMover name="save_full" reference_name="full_pose" restore_pose="0"/>112 <SavePoseMover name="restore_full" reference_name="full_pose" restore_pose="1"/>113%s114 </MOVERS>115 116 <PROTOCOLS>117%s118 </PROTOCOLS>119 <OUTPUT />120</ROSETTASCRIPTS>121''' % (scorefxn, selectors, movers, '\n'.join(steps))122 123 with open(out_path, 'w') as f:124 f.write(xml)125 return out_path126 127 128def rescore_one(struct_db3, protocol_path):129 """Run the isolate-and-rescore protocol on one struct.db3, writing per_chain.db3 beside it."""130 working_dir = os.path.dirname(os.path.abspath(struct_db3))131 out_db3 = os.path.join(working_dir, per_chain_db3_name)132 if os.path.isfile(out_db3):133 os.remove(out_db3)134 135 args = [136 os.path.abspath(rosetta_scripts_path),137 '-inout:dbms:database_name', struct_db3_name,138 '-in:use_database',139 '-parser:protocol', os.path.abspath(protocol_path),140 ] + rosetta_flags141 142 log_path = os.path.join(working_dir, 'per_chain_rescore.log')143 with open(log_path, 'w') as log:144 proc = subprocess.Popen(args, stdout=log, stderr=subprocess.STDOUT, cwd=working_dir)145 returncode = proc.wait()146 147 if returncode != 0 or not os.path.isfile(out_db3):148 raise RuntimeError('Rescoring failed for %s -- see %s' % (struct_db3, log_path))149 return out_db3150 151 152def read_per_chain_db3(per_chain_db3, struct_number, case_name, stride):153 conn = sqlite3.connect(per_chain_db3)154 df = pd.read_sql_query('''155 SELECT batches.name AS batch, structures.tag AS tag, structure_scores.score_value AS energy156 FROM structure_scores157 INNER JOIN structures ON structures.struct_id=structure_scores.struct_id158 INNER JOIN batches ON batches.batch_id=structure_scores.batch_id159 INNER JOIN score_types ON score_types.batch_id=structure_scores.batch_id160 AND score_types.score_type_id=structure_scores.score_type_id161 WHERE score_types.score_type_name="total_score"162 ''', conn)163 conn.close()164 165 # batch name is "chain_<X>_report"; tag is "<original struct.db3 struct_id>_0001"166 df['chain'] = df['batch'].apply(lambda b: b[len('chain_'):-len('_report')])167 original_id = df['tag'].apply(lambda t: int(re.match(r'(\d+)', t).group(1)))168 df['pose'] = original_id.apply(lambda i: pose_order[(i - 1) % len(pose_order)])169 df['backrub_steps'] = original_id.apply(lambda i: stride * (((i - 1) // len(pose_order)) + 1))170 df['struct_num'] = struct_number171 df['case_name'] = case_name172 return df[['case_name', 'struct_num', 'backrub_steps', 'chain', 'pose', 'energy']]173 174 175def main():176 parser = argparse.ArgumentParser(description=__doc__,177 formatter_class=argparse.RawDescriptionHelpFormatter)178 parser.add_argument('output_folder', help='flex ddG output folder (e.g. "output")')179 parser.add_argument('--stride', type=int, default=None,180 help='override backrub_trajectory_stride instead of reading it from each '181 'struct.db3. Affects checkpoint labels only, not any energy.')182 parser.add_argument('--chains', default=None,183 help='comma-separated chains (default: auto-detect from struct.db3)')184 parser.add_argument('--csv', default=None, help='write the full per-structure table here')185 parser.add_argument('--reuse', action='store_true',186 help='skip Rosetta where per_chain.db3 already exists')187 args = parser.parse_args()188 189 if not os.path.isfile(rosetta_scripts_path):190 sys.exit('ERROR: set rosetta_scripts_path to your compiled rosetta_scripts binary')191 192 struct_db3s = sorted(glob.glob(os.path.join(args.output_folder, '*', '*', struct_db3_name)))193 if not struct_db3s:194 sys.exit('ERROR: no %s found under %s' % (struct_db3_name, args.output_folder))195 print('Found %d %s files' % (len(struct_db3s), struct_db3_name))196 197 chains = args.chains.split(',') if args.chains else chains_in_struct_db3(struct_db3s[0])198 print('Chains: %s' % ', '.join(chains))199 200 protocol_path = os.path.join(args.output_folder, 'per_chain_rescore.generated.xml')201 write_rescore_protocol(chains, protocol_path)202 print('Wrote protocol %s\n' % protocol_path)203 204 frames = []205 for i, struct_db3 in enumerate(struct_db3s, start=1):206 struct_dir = os.path.dirname(struct_db3)207 case_name = os.path.basename(os.path.dirname(struct_dir))208 struct_number = os.path.basename(struct_dir)209 out_db3 = os.path.join(struct_dir, per_chain_db3_name)210 211 if args.reuse and os.path.isfile(out_db3):212 print(' [%d/%d] %s (reusing)' % (i, len(struct_db3s), struct_dir))213 else:214 print(' [%d/%d] %s' % (i, len(struct_db3s), struct_dir))215 out_db3 = rescore_one(struct_db3, protocol_path)216 217 stride = args.stride218 if stride is None:219 stride = flex_ddg_db3.trajectory_stride_from_db3(struct_db3)220 if stride is None:221 stride = 5222 print(' WARNING: %s does not record backrub_trajectory_stride; assuming %d. '223 'Pass --stride to label the checkpoints correctly.' % (struct_db3, stride))224 225 frames.append(read_per_chain_db3(out_db3, struct_number, case_name, stride))226 227 per_structure = pd.concat(frames)228 wide = per_structure[per_structure['pose'].isin(['wt', 'mut'])].pivot_table(229 index=['case_name', 'chain', 'backrub_steps', 'struct_num'],230 columns='pose', values='energy').reset_index()231 wide['ddG'] = wide['mut'] - wide['wt']232 233 if args.csv:234 wide.to_csv(args.csv, index=False)235 print('\nWrote %s' % args.csv)236 237 summary = wide.groupby(['case_name', 'chain', 'backrub_steps']).agg(238 nstruct=('ddG', 'size'),239 wt_intra=('wt', 'mean'),240 mut_intra=('mut', 'mean'),241 ddG=('ddG', 'mean'),242 ddG_sd=('ddG', 'std'),243 ).reset_index()244 summary['ddG_sem'] = summary['ddG_sd'] / np.sqrt(summary['nstruct'])245 246 print('\n=== per-chain intramolecular ddG ===')247 print(summary.round(4).to_string(index=False))248 print('\nA chain you did NOT mutate should read ~0 within ddG_sem.')249 print('This is bound-conformation strain, not a folding ddG (see the module docstring).')250 251 252if __name__ == '__main__':253 main()254 