Team Ai
Modelpublic

OneScience-Group/flex_ddG_tutorial

sourceHugging Facemitupdated 1mo agoView on Hugging Face
0likes15downloads
analyze_flex_ddG.py345 linesDownload Raw Back to scripts
1#!/usr/bin/python32 3import sys4import os5import sqlite36import shutil7import tempfile8from pprint import pprint9import pandas as pd10import numpy as np11import re12import argparse13import datetime14import sys15import collections16import threading17 18import flex_ddg_db319 20rosetta_output_file_name = 'rosetta.out'21output_database_name = 'ddG.db3'22script_output_folder = 'analysis_output'23 24# Only a fallback. The stride each run actually used is read back out of its own ddG.db3, so25# runs with different strides analyze correctly and nothing here needs editing to match a run.26# This value is used only when the database does not record it, and a warning is printed.27default_trajectory_stride = 528 29zemu_gam_params = {30    'fa_sol' :      (6.940, -6.722),31    'hbond_sc' :    (1.902, -1.999),32    'hbond_bb_sc' : (0.063,  0.452),33    'fa_rep' :      (1.659, -0.836),34    'fa_elec' :     (0.697, -0.122),35    'hbond_lr_bb' : (2.738, -1.179),36    'fa_atr' :      (2.313, -1.649),37}38 39def gam_function(x, score_term = None ):40    return -1.0 * np.exp( zemu_gam_params[score_term][0] ) + 2.0 * np.exp( zemu_gam_params[score_term][0] ) / ( 1.0 + np.exp( -1.0 * x * np.exp( zemu_gam_params[score_term][1] ) ) )41 42def apply_zemu_gam(scores):43    new_columns = list(scores.columns)44    new_columns.remove('total_score')45    scores = scores.copy()[ new_columns ]46    for score_term in zemu_gam_params:47        assert( score_term in scores.columns )48        scores[score_term] = scores[score_term].apply( gam_function, score_term = score_term )49    scores[ 'total_score' ] = scores[ list(zemu_gam_params.keys()) ].sum( axis = 1 )50    scores[ 'score_function_name' ] = scores[ 'score_function_name' ] + '-gam'51    return scores52 53def rosetta_output_succeeded( potential_struct_dir ):54    path_to_rosetta_output = os.path.join( potential_struct_dir, rosetta_output_file_name )55    if not os.path.isfile(path_to_rosetta_output):56        return False57 58    db3_file = os.path.join( potential_struct_dir, output_database_name )59    if not os.path.isfile( db3_file ):60        return False61 62    success_line_found = False63    no_more_batches_line_found = False64    with open( path_to_rosetta_output, 'r' ) as f:65        for line in f:66            if line.startswith( 'protocols.jd2.JobDistributor' ) and 'reported success in' in line:67                success_line_found = True68            if line.startswith( 'protocols.jd2.JobDistributor' ) and 'no more batches to process' in line:69                no_more_batches_line_found = True70 71    return no_more_batches_line_found and success_line_found72 73def find_finished_jobs( output_folder ):74    return_dict = {}75    job_dirs = [ os.path.abspath(os.path.join(output_folder, d)) for d in os.listdir(output_folder) if os.path.isdir( os.path.join(output_folder, d) )]76    for job_dir in job_dirs:77        completed_struct_dirs = []78        for potential_struct_dir in sorted([ os.path.abspath(os.path.join(job_dir, d)) for d in os.listdir(job_dir) if os.path.isdir( os.path.join(job_dir, d) )]):79            if rosetta_output_succeeded( potential_struct_dir ):80                completed_struct_dirs.append( potential_struct_dir )81        return_dict[job_dir] = completed_struct_dirs82 83    return return_dict84 85def get_scores_from_db3_file(db3_file, struct_number, case_name, trajectory_stride):86    conn = sqlite3.connect(db3_file)87    conn.row_factory = sqlite3.Row88    c = conn.cursor()89 90    num_batches = c.execute('SELECT max(batch_id) from batches').fetchone()[0]91 92    scores = pd.read_sql_query('''93    SELECT batches.name, structure_scores.struct_id, score_types.score_type_name, structure_scores.score_value, score_function_method_options.score_function_name from structure_scores94    INNER JOIN batches ON batches.batch_id=structure_scores.batch_id95    INNER JOIN score_function_method_options ON score_function_method_options.batch_id=batches.batch_id96    INNER JOIN score_types ON score_types.batch_id=structure_scores.batch_id AND score_types.score_type_id=structure_scores.score_type_id97    ''', conn)98 99    def renumber_struct_id( struct_id ):100        return trajectory_stride * ( 1 + (int(struct_id-1) // num_batches) )101 102    scores['struct_id'] = scores['struct_id'].apply( renumber_struct_id )103    scores['name'] = scores['name'].apply( lambda x: x[:-9] if x.endswith('_dbreport') else x )104    scores = scores.pivot_table( index = ['name', 'struct_id', 'score_function_name'], columns = 'score_type_name', values = 'score_value' ).reset_index()105    scores.rename( columns = {106        'name' : 'state',107        'struct_id' : 'backrub_steps',108    }, inplace=True)109    scores['struct_num'] = struct_number110    scores['case_name'] = case_name111 112    conn.close()113 114    return scores115 116def get_per_chain_scores_from_db3_file(db3_file, struct_number, case_name, trajectory_stride):117    '''Read the per-chain intramolecular energies written by the per-chain protocol variant118    (see per_chain_protocol.py). Returns None if the run did not report them.119 120    Only the unbound states are meaningful here: the chains are 1000 A apart, so there are no121    cross-chain pair energies and each value is exactly that chain's intramolecular energy. On122    the bound states the value additionally carries roughly half the interface energy, because123    Rosetta splits each two-body term between its two residues.124 125    Note that chain IDs come back lowercased, because Rosetta lowercases database table names.126    Chain "A" appears here as "a". per_chain_protocol.py refuses to set up a run whose chain IDs127    differ only by case, so this stays unambiguous.128    '''129    conn = sqlite3.connect(db3_file)130    conn.row_factory = sqlite3.Row131    c = conn.cursor()132 133    chain_tables = [ row[0] for row in c.execute(134        "SELECT name FROM sqlite_master WHERE type='table' AND name LIKE 'chain\\_%\\_energy' ESCAPE '\\'"135    ).fetchall() ]136    if len(chain_tables) == 0:137        conn.close()138        return None139 140    num_batches = c.execute('SELECT max(batch_id) from batches').fetchone()[0]141 142    dfs = []143    for table in chain_tables:144        chain = table[len('chain_'):-len('_energy')]145        df = pd.read_sql_query('''146        SELECT batches.name, %s.struct_id, %s.total_energy from %s147        INNER JOIN structures ON structures.struct_id=%s.struct_id148        INNER JOIN batches ON batches.batch_id=structures.batch_id149        ''' % (table, table, table, table), conn)150        df['chain'] = chain151        dfs.append(df)152    conn.close()153 154    scores = pd.concat( dfs )155    scores['struct_id'] = scores['struct_id'].apply(156        lambda struct_id: trajectory_stride * ( 1 + (int(struct_id-1) // num_batches) ) )157    scores['name'] = scores['name'].apply( lambda x: x[:-9] if x.endswith('_dbreport') else x )158    scores.rename( columns = {159        'name' : 'state',160        'struct_id' : 'backrub_steps',161        'total_energy' : 'intra_energy',162    }, inplace=True)163    scores['struct_num'] = struct_number164    scores['case_name'] = case_name165 166    return scores167 168def calc_per_chain_ddg( scores ):169    '''Per-chain intramolecular ddG, read off the unbound states and averaged over nstruct.'''170    unbound = scores.loc[ scores['state'].isin(['unbound_wt', 'unbound_mut']) ].copy()171    if len(unbound) == 0:172        return None173 174    wide = unbound.pivot_table(175        index = ['case_name', 'chain', 'backrub_steps', 'struct_num'],176        columns = 'state', values = 'intra_energy' ).reset_index()177    if 'unbound_wt' not in wide.columns or 'unbound_mut' not in wide.columns:178        return None179    wide['ddG'] = wide['unbound_mut'] - wide['unbound_wt']180 181    summary = wide.groupby( ['case_name', 'chain', 'backrub_steps'] ).agg(182        nstruct = ('ddG', 'size'),183        wt_intra = ('unbound_wt', 'mean'),184        mut_intra = ('unbound_mut', 'mean'),185        ddG = ('ddG', 'mean'),186        ddG_sd = ('ddG', 'std'),187    ).reset_index()188    summary['ddG_sem'] = summary['ddG_sd'] / np.sqrt( summary['nstruct'] )189    return summary.round(decimals=5)190 191def resolve_trajectory_stride( db3_file, stride_override = None ):192    '''Stride to label this database's checkpoints with, preferring what the run recorded.'''193    if stride_override is not None:194        return stride_override195 196    stride = flex_ddg_db3.trajectory_stride_from_db3( db3_file )197    if stride is not None:198        return stride199 200    print( 'WARNING: %s does not record backrub_trajectory_stride; assuming %d.' % (201        db3_file, default_trajectory_stride ) )202    print( '         If the run used a different stride, pass --stride to label the' )203    print( '         checkpoints correctly. This affects labels only, not any energy.' )204    return default_trajectory_stride205 206def process_finished_struct( output_path, case_name, stride_override = None ):207    db3_file = os.path.join( output_path, output_database_name )208    assert( os.path.isfile( db3_file ) )209    struct_number = int( os.path.basename(output_path) )210    trajectory_stride = resolve_trajectory_stride( db3_file, stride_override )211    scores_df = get_scores_from_db3_file( db3_file, struct_number, case_name, trajectory_stride )212    per_chain_df = get_per_chain_scores_from_db3_file( db3_file, struct_number, case_name, trajectory_stride )213 214    return scores_df, per_chain_df215 216def calc_ddg( scores ):217    total_structs = np.max( scores['struct_num'] )218 219    nstructs_to_analyze = set([total_structs])220    for x in range(10, total_structs):221        if x % 10 == 0:222            nstructs_to_analyze.add(x)223    nstructs_to_analyze = sorted(nstructs_to_analyze)224 225    all_ddg_scores = []226    for nstructs in nstructs_to_analyze:227        ddg_scores = scores.loc[ ((scores['state'] == 'unbound_mut') | (scores['state'] == 'bound_wt')) & (scores['struct_num'] <= nstructs) ].copy()228        for column in ddg_scores.columns:229            if column not in ['state', 'case_name', 'backrub_steps', 'struct_num', 'score_function_name']:230                ddg_scores.loc[:,column] *= -1.0231        ddg_scores = pd.concat( [ ddg_scores, scores.loc[ ((scores['state'] == 'unbound_wt') | (scores['state'] == 'bound_mut')) & (scores['struct_num'] <= nstructs) ].copy() ] )232        ddg_scores = ddg_scores.groupby( ['case_name', 'backrub_steps', 'struct_num', 'score_function_name'] ).sum( numeric_only = True ).reset_index()233 234        if nstructs == total_structs:235            struct_scores = ddg_scores.copy()236 237        ddg_scores = ddg_scores.groupby( ['case_name', 'backrub_steps', 'score_function_name'] ).mean( numeric_only = True ).round(decimals=5).reset_index()238        new_columns = list(ddg_scores.columns.values)239        new_columns.remove( 'struct_num' )240        ddg_scores = ddg_scores[new_columns]241        ddg_scores[ 'scored_state' ] = 'ddG'242        ddg_scores[ 'nstruct' ] = nstructs243        all_ddg_scores.append(ddg_scores)244 245    return (pd.concat(all_ddg_scores), struct_scores)246 247def calc_dgs( scores ):248    l = []249 250    total_structs = np.max( scores['struct_num'] )251 252    nstructs_to_analyze = set([total_structs])253    for x in range(10, total_structs):254        if x % 10 == 0:255            nstructs_to_analyze.add(x)256    nstructs_to_analyze = sorted(nstructs_to_analyze)257 258    for state in ['mut', 'wt']:259        for nstructs in nstructs_to_analyze:260            dg_scores = scores.loc[ (scores['state'].str.endswith(state)) & (scores['state'].str.startswith('unbound')) & (scores['struct_num'] <= nstructs) ].copy()261            for column in dg_scores.columns:262                if column not in ['state', 'case_name', 'backrub_steps', 'struct_num', 'score_function_name']:263                    dg_scores.loc[:,column] *= -1.0264            dg_scores = pd.concat( [ dg_scores, scores.loc[ (scores['state'].str.endswith(state)) & (scores['state'].str.startswith('bound')) & (scores['struct_num'] <= nstructs) ].copy() ] )265            dg_scores = dg_scores.groupby( ['case_name', 'backrub_steps', 'struct_num', 'score_function_name'] ).sum( numeric_only = True ).reset_index()266            dg_scores = dg_scores.groupby( ['case_name', 'backrub_steps', 'score_function_name'] ).mean( numeric_only = True ).round(decimals=5).reset_index()267            new_columns = list(dg_scores.columns.values)268            new_columns.remove( 'struct_num' )269            dg_scores = dg_scores[new_columns]270            dg_scores[ 'scored_state' ] = state + '_dG'271            dg_scores[ 'nstruct' ] = nstructs272            l.append( dg_scores )273    return l274 275def analyze_output_folder( output_folder, stride_override = None ):276    # Pass in an outer output folder. Subdirectories are considered different mutation cases, with subdirectories of different structures.277    finished_jobs = find_finished_jobs( output_folder )278    if len(finished_jobs) == 0:279        print( 'No finished jobs found' )280        return281 282    ddg_scores_dfs = []283    struct_scores_dfs = []284    per_chain_dfs = []285    for finished_job, finished_structs in finished_jobs.items():286        inner_scores_list = []287        inner_per_chain_list = []288        for finished_struct in finished_structs:289            inner_scores, inner_per_chain = process_finished_struct( finished_struct, os.path.basename(finished_job), stride_override )290            inner_scores_list.append( inner_scores )291            if inner_per_chain is not None:292                inner_per_chain_list.append( inner_per_chain )293        scores = pd.concat( inner_scores_list )294        if len(inner_per_chain_list) > 0:295            per_chain_summary = calc_per_chain_ddg( pd.concat( inner_per_chain_list ) )296            if per_chain_summary is not None:297                per_chain_dfs.append( per_chain_summary )298        ddg_scores, struct_scores = calc_ddg( scores )299        struct_scores_dfs.append( struct_scores )300        ddg_scores_dfs.append( ddg_scores )301        ddg_scores_dfs.append( apply_zemu_gam(ddg_scores) )302        ddg_scores_dfs.extend( calc_dgs( scores ) )303 304    if not os.path.isdir(script_output_folder):305        os.makedirs(script_output_folder)306    basename = os.path.basename(output_folder)307 308    pd.concat( struct_scores_dfs ).to_csv( os.path.join(script_output_folder, basename + '-struct_scores_results.csv' ) )309 310    df = pd.concat( ddg_scores_dfs )311    df.to_csv( os.path.join(script_output_folder, basename + '-results.csv') )312 313    display_columns = ['backrub_steps', 'case_name', 'nstruct', 'score_function_name', 'scored_state', 'total_score']314    for score_type in ['mut_dG', 'wt_dG', 'ddG']:315        print( score_type )316        print( df.loc[ df['scored_state'] == score_type ][display_columns].head( n = 20 ) )317        print( '' )318 319    if len(per_chain_dfs) > 0:320        per_chain = pd.concat( per_chain_dfs )321        per_chain.to_csv( os.path.join(script_output_folder, basename + '-per_chain_results.csv'), index = False )322        print( 'per-chain intramolecular ddG (from the unbound states)' )323        print( per_chain.head( n = 40 ).to_string(index = False) )324        print( '' )325        print( 'NOTE: this is the intramolecular strain difference in the *bound* backbone' )326        print( '      conformation, not a folding ddG -- the unbound state is never relaxed.' )327        print( '      A chain you did not mutate should come out at ~0 +/- ddG_sem; if it does' )328        print( '      not, nstruct is too low to average out the whole-pose minimization noise.' )329        print( '' )330 331if __name__ == '__main__':332    parser = argparse.ArgumentParser(333        description = 'Analyze one or more flex ddG output folders (e.g. "output").' )334    parser.add_argument( 'output_folders', nargs = '+', help = 'flex ddG output folder(s)' )335    parser.add_argument( '--stride', type = int, default = None,336                         help = 'override backrub_trajectory_stride instead of reading it from'337                                ' each ddG.db3. Affects checkpoint labels only, not any energy.' )338    parsed_args = parser.parse_args()339 340    for folder_to_analyze in parsed_args.output_folders:341        if os.path.isdir( folder_to_analyze ):342            analyze_output_folder( folder_to_analyze, parsed_args.stride )343        else:344            print( 'ERROR: %s is not a valid directory' % folder_to_analyze )345