Team Ai
Modelpublic

OneScience-Group/flex_ddG_tutorial

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