OneScience-Group/flex_ddG_tutorial
015
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 