Team Ai
Modelpublic

OneScience-Group/flex_ddG_tutorial

sourceHugging Facemitupdated 1mo agoView on Hugging Face
0likes27downloads
reprocess_per_chain.py254 linesDownload Raw Back to scripts
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