OneScience-Group/flex_ddG_tutorial
027
1#!/usr/bin/env python32 3import os4import sys5import argparse6import functools7import subprocess8import re9import shutil10import datetime11import math12import collections13import threading14 15import flex_ddg_db316 17use_multiprocessing = False18if use_multiprocessing:19 import multiprocessing20 21# The Reporter class is useful for printing output for tasks which will take a long time22# Really, you should just use tqdm now, but I used this before I knew about tqdm and it removes a dependency23 24# Time in seconds function25# Converts datetime timedelta object to number of seconds26def ts(td):27 return (td.microseconds + (td.seconds + td.days * 24 * 3600) * 1e6) / 1e628 29def mean(l):30 # Not using numpy mean to avoid dependency31 return float( sum(l) ) / float( len(l) )32 33class Reporter:34 def __init__( self, task, entries = 'files', print_output = True, eol_char = '\r' ):35 self._lock = threading.Lock()36 self.print_output = print_output37 self.start = datetime.datetime.now()38 self.entries = entries39 self.lastreport = self.start40 self.task = task41 self.report_interval = datetime.timedelta( seconds = 1 ) # Interval to print progress42 self.n = 043 self.completion_time = None44 if self.print_output:45 print('\nStarting ' + task)46 self.total_count = None # Total tasks to be processed47 self.maximum_output_string_length = 048 self.rolling_est_total_time = collections.deque( maxlen = 50 )49 self.kv_callback_results = {}50 self.list_results = []51 self.eol_char = eol_char52 53 def set_total_count(self, x):54 self.total_count = x55 self.rolling_est_total_time = collections.deque( maxlen = max(1, int( .05 * x )) )56 57 def decrement_total_count(self):58 if self.total_count:59 self.total_count -= 160 61 def report(self, n):62 with self._lock:63 self.n = n64 time_now = datetime.datetime.now()65 if self.print_output and self.lastreport < (time_now - self.report_interval):66 self.lastreport = time_now67 if self.total_count:68 percent_done = float(self.n) / float(self.total_count)69 est_total_time_seconds = ts(time_now - self.start) * (1.0 / percent_done)70 self.rolling_est_total_time.append( est_total_time_seconds )71 est_total_time = datetime.timedelta( seconds = mean(self.rolling_est_total_time) )72 time_remaining = est_total_time - (time_now - self.start)73 eta = time_now + time_remaining74 time_remaining_str = 'ETA: %s Est. time remaining: ' % eta.strftime("%Y-%m-%d %H:%M:%S")75 76 time_remaining_str += str( datetime.timedelta( seconds = int(ts(time_remaining)) ) )77 78 output_string = " Processed: %d %s (%.1f%%) %s" % (n, self.entries, percent_done*100.0, time_remaining_str)79 else:80 output_string = " Processed: %d %s" % (n, self.entries)81 82 output_string += self.eol_char83 84 if len(output_string) > self.maximum_output_string_length:85 self.maximum_output_string_length = len(output_string)86 elif len(output_string) < self.maximum_output_string_length:87 output_string = output_string.ljust(self.maximum_output_string_length)88 sys.stdout.write( output_string )89 sys.stdout.flush()90 91 def increment_report(self):92 self.report(self.n + 1)93 94 def increment_report_callback(self, cb_value):95 self.increment_report()96 97 def increment_report_keyval_callback(self, kv_pair):98 key, value = kv_pair99 self.kv_callback_results[key] = value100 self.increment_report()101 102 def increment_report_list_callback(self, new_list_items):103 self.list_results.extend(new_list_items)104 self.increment_report()105 106 def decrement_report(self):107 self.report(self.n - 1)108 109 def add_to_report(self, x):110 self.report(self.n + x)111 112 def done(self):113 self.completion_time = datetime.datetime.now()114 if self.print_output:115 print('Done %s, processed %d %s, took %s\n' % (self.task, self.n, self.entries, self.completion_time-self.start))116 117 def elapsed_time(self):118 if self.completion_time:119 return self.completion_time - self.start120 else:121 return datetime.datetime.now() - self.start122 123 124struct_db3_file = 'struct.db3'125 126# Extraction uses the score_jd2 binary, not rosetta_scripts. It is built alongside127# rosetta_scripts by the standard Rosetta build, but is a separate executable.128#score_jd2_path = os.path.expanduser( '~/rosetta/source/bin/score_jd2' )129score_jd2_path = os.path.expanduser(130"/public/home/scnb9biwet/jiangqq/flex_ddG_tutorial-master/software/rosetta3.9/main/source/bin/score_jd2.default.linuxgccrelease"131)132# Only a fallback. Extracted structures are named by how many backrub steps produced them, so133# the stride each run used is read back out of its own struct.db3. This value is used only when134# the database does not record it, and a warning is printed.135default_trajectory_stride = 5136 137def resolve_trajectory_stride( struct_db, stride_override = None ):138 '''Stride to name this database's extracted PDBs with, preferring what the run recorded.'''139 if stride_override is not None:140 return stride_override141 142 stride = flex_ddg_db3.trajectory_stride_from_db3( struct_db )143 if stride is not None:144 return stride145 146 print( 'WARNING: %s does not record backrub_trajectory_stride; assuming %d.' % (147 struct_db, default_trajectory_stride ) )148 print( ' If the run used a different stride, pass --stride, or the extracted PDBs' )149 print( ' will be named with the wrong backrub step counts.' )150 return default_trajectory_stride151 152def recursive_find_struct_dbs( input_dir ):153 return_list = []154 155 for path in [os.path.join(input_dir, x) for x in os.listdir( input_dir )]:156 if os.path.isdir( path ):157 return_list.extend( recursive_find_struct_dbs( path ) )158 elif os.path.isfile( path ) and os.path.basename( path ) == struct_db3_file:159 return_list.append( path )160 161 return return_list162 163def extract_structures( struct_db, rename_function = None ):164 args = [165 os.path.abspath( score_jd2_path ),166 '-inout:dbms:database_name', struct_db3_file,167 '-in:use_database',168 '-out:pdb',169 ]170 171 working_directory = os.path.dirname( struct_db )172 rosetta_outfile_path = os.path.join(working_directory, 'structure_output.txt' )173 if not use_multiprocessing:174 print(rosetta_outfile_path)175 rosetta_outfile = open( rosetta_outfile_path, 'w')176 if not use_multiprocessing:177 print( ' '.join( args ) )178 # No shell: joining the arguments into a string breaks as soon as a path contains a space.179 rosetta_process = subprocess.Popen(180 args,181 stdout=rosetta_outfile, stderr=subprocess.STDOUT, close_fds = True, cwd = working_directory,182 )183 return_code = rosetta_process.wait()184 rosetta_outfile.close()185 186 if return_code == 0:187 os.remove( rosetta_outfile_path )188 else:189 print( 'ERROR: score_jd2 failed on %s (exit %d) -- see %s' % (190 struct_db, return_code, rosetta_outfile_path ) )191 return return_code192 193 if rename_function != None:194 for path in [ os.path.join( working_directory, x ) for x in os.listdir( working_directory ) ]:195 m = re.match( r'(\d+)_0001\.pdb$', os.path.basename(path) )196 if m:197 dest_path = os.path.join( working_directory, rename_function( int(m.group(1)) ) )198 shutil.move( path, dest_path )199 200 return return_code201 202def flex_ddG_rename(struct_id, trajectory_stride):203 steps = [204 'backrub',205 'wt',206 'mut',207 ]208 209 return '%s_%05d.pdb' % ( steps[ (struct_id-1) % len(steps) ], (((struct_id-1) // len(steps)) + 1) * trajectory_stride )210 211def main( input_dir, stride_override = None ):212 struct_dbs = recursive_find_struct_dbs( input_dir )213 print( 'Found {:d} structure database files to extract'.format( len(struct_dbs) ) )214 215 if use_multiprocessing:216 pool = multiprocessing.Pool()217 r = Reporter('extracting structure database files', entries = '.db3 files')218 r.set_total_count( len(struct_dbs) )219 220 for struct_db in struct_dbs:221 # Each database is named using the stride its own run was launched with.222 # functools.partial rather than a lambda, so that this stays picklable for the223 # multiprocessing path below.224 stride = resolve_trajectory_stride( struct_db, stride_override )225 rename_function = functools.partial( flex_ddG_rename, trajectory_stride = stride )226 if use_multiprocessing:227 pool.apply_async(228 extract_structures,229 args = (struct_db,),230 kwds = {'rename_function' : rename_function},231 callback = r.increment_report_callback232 )233 else:234 r.increment_report_callback(235 extract_structures( struct_db, rename_function = rename_function )236 )237 238 if use_multiprocessing:239 pool.close()240 pool.join()241 r.done()242 243if __name__ == '__main__':244 parser = argparse.ArgumentParser(245 description = 'Extract PDBs from the struct.db3 files under a flex ddG output folder.' )246 parser.add_argument( 'output_folders', nargs = '+', help = 'flex ddG output folder(s)' )247 parser.add_argument( '--stride', type = int, default = None,248 help = 'override backrub_trajectory_stride instead of reading it from'249 ' each struct.db3. Affects extracted PDB names only.' )250 parsed_args = parser.parse_args()251 252 if not os.path.isfile( score_jd2_path ):253 print( 'ERROR: "score_jd2_path" variable must be set to the location of the "score_jd2" binary executable' )254 print( 'This file might look something like: "score_jd2.linuxgccrelease"' )255 print( 'Note that this is a different executable from the "rosetta_scripts" binary used to run flex ddG' )256 raise Exception( 'score_jd2 missing' )257 258 for x in parsed_args.output_folders:259 if os.path.isdir(x):260 main( x, parsed_args.stride )261 else:262 print( 'ERROR: %s is not a valid directory' % x )263 