Team Ai
Modelpublic

OneScience-Group/TemStaPro-main

sourceHugging Facemitupdated 1mo agoView on Hugging Face
0likes19downloads
results.py245 linesDownload Raw Back to scripts
1"""2Representing the output of the program.3"""4 5import numpy6import matplotlib.pyplot as plt7 8def get_temperature_label(predictions, temperature_ranges, left_hand=True):9    """10    Process the raw output of the inference model to get temperature range 11    labels.12 13    predictions - LIST that contains predictions for each temperature range14    temperature_ranges - LIST with temperature ranges' labels15    left_hand - BOOLEAN that indicates to find the left-hand 16        (True) or right-hand (False) limit17 18    returns STRING that is the label of the limiting temperature range19    """20    if(left_hand):21        for j, pred in enumerate(predictions):22            if(float(pred) < 0.5):23                return temperature_ranges[j]24            elif(float(pred) >= 0.5 and j != len(predictions)-1):25                continue26            else:27                return temperature_ranges[-1]28    else:29        for j, pred in enumerate(predictions[::-1]):30            if(float(pred) >= 0.5):31                return temperature_ranges[len(predictions)-j]32            elif(float(pred) < 0.5 and j != len(predictions)-1):33                continue34            else:35                return temperature_ranges[0]36 37def detect_clash(predictions, left_hand=True):38    """39    Detecting the conflicting predictions of the ensemble.40 41    predictions - LIST that contains predictions for each temperature range42    left_hand - BOOLEAN that indicates to find the clash from left-hand 43        (True) or right-hand (False)44    45    returns STRING '-' if clash was not detected, '*' if it was46    """	47    if(left_hand):48        for j, pred in enumerate(predictions):49            if(j and round(float(predictions[j-1])) <50                round(float(predictions[j]))):51                return "*"52            elif(j and round(float(predictions[j-1])) >=53                round(float(predictions[j])) and j != len(predictions)-1):54                continue55            elif(j and round(float(predictions[j-1])) >=56                round(float(predictions[j])) and j == len(predictions)-1):57                return "-"58            elif(len(predictions) == 1):59                return "-"60                61    else:62        for j, pred in enumerate(predictions[::-1]):63            if(j != len(predictions)-1 and round(float(predictions[j-1])) <64                round(float(predictions[j]))):65                return "*"66            elif(j != len(predictions)-1 and round(float(predictions[j-1])) >=67                round(float(predictions[j])) and j != len(predictions)-2):68                continue69            elif(j != len(predictions)-1 and round(float(predictions[j-1])) >=70                round(float(predictions[j])) and j == len(predictions)-2):71                return "-"72            elif(len(predictions) == 1):73                return "-"74 75def print_inferences_header(file_handle, thresholds, 76    print_thermophilicity=False):77    """78    Print inferences table header.79 80    file_handle - FILE to which the results will be printed81    thresholds - LIST of thresholds that are used82    print_thermophilicity - BOOLEAN that determines whether to print the 83        thermophilicity column84    """85 86    predictions_columns_names = ""87    for threshold in thresholds:88        predictions_columns_names += f"t{threshold}_binary\tt{threshold}_raw\t"89 90    header = f"protein_id\tposition\tsequence\tlength\t{predictions_columns_names}"+\91        f"left_hand_label\tright_hand_label\tclash"92    if(print_thermophilicity): header += "\tthermophilicity"93 94    print(header, file=file_handle)95 96def print_inferences(averaged_inferences, binary_inferences, original_headers,97    labels, clashes, thermophilicity_labels, file_handle, sequences=None, 98    run_mode='mean', print_thermophilicity=False):99    """100    Print results.101 102    averaged_inferences - LIST of DICT that keeps each sequence's mean inferences103    binary_inferences - LIST of DICT that keeps each sequence's binary inferences104    original_headers - DICT of original sequences' headers for printing105    labels - LIST of DICT that keeps each sequence's left-hand and right-hand 106        temperature prediction labels107    clashes - LIST of DICT that keeps each sequence's clash labels108    thermophilicity_labels - DICT with possible thermophilicity labels109    sequences - LIST of DICT that keeps sequence ids as keys and sequences as values110    file_handle - FILE to which the results will be printed111    run_mode - STRING that determines which run mode is executed:112        'mean', 'per-res', 'per-segment'113    print_thermophilicity - BOOLEAN that determines to print the 114        thermophilicity column115    """116 117    if(sequences is None): return118 119    for proc_header in averaged_inferences.keys():120        merged_inferences = []121        for i, inf in enumerate(binary_inferences[proc_header]):122            merged_inferences.append("%d" % binary_inferences[proc_header][i])123            merged_inferences.append("%.3e" % averaged_inferences[proc_header][i])124 125        # Setting the default values for run_mode 'mean'126        if(run_mode == "mean"):127            out_header = original_headers[proc_header]128            position = '-'129        elif(run_mode == "per-segment"):130            out_header = original_headers["_".join(proc_header.split("_")[0:-1])]131            pos_range = proc_header.split("_")[-1].split("-")132            range_length = int(pos_range[1])-int(pos_range[0])133            134            # Calculating the position (numerated from 1)135            position = str(int(pos_range[0])+int(range_length/2)+1)136        elif(run_mode == "per-res"):137            out_header = original_headers["_".join(proc_header.split("_")[0:-1])]138            position = str(int(proc_header.split("_")[-1])+1)139       140        output_line = "%s\t%s\t%s\t%d\t%s\t%s\t%s" % (out_header, position, 141            sequences[proc_header],142            len(sequences[proc_header]), "\t".join(merged_inferences),143            "\t".join(labels[proc_header]), clashes[proc_header][0])144        145        # Choosing the thermophilicity label146        if(print_thermophilicity):147            thermophilicity = "undetermined"148            if(labels[proc_header][0] == labels[proc_header][1]):149                for t in list(thermophilicity_labels.keys()):150                    if(labels[proc_header][0] in thermophilicity_labels[t]):151                        thermophilicity = t152                        break153            output_line += f"\t{thermophilicity}"154        155        print(output_line, file=file_handle)156 157def plot_per_res_inferences(averaged_inferences, thresholds, plot_dir, 158    smoothen=True, window_size=21, x_label="residue index", 159    title="Per-residue predictions"):160    """161    Plotting per-residue inferences.162 163    averaged_inferences - DICT that keeps each sequence's inferences 164        (averaged of all threshold models))165    thresholds - LIST with binary models' temperature thresholds166    plot_dir - STRING that determines the directory where plots should 167        be saved168    smoothen - BOOL indicates to plot smoothened curve169    """170    WINDOW_SIZE = window_size171 172    original_seq_ids = set()173    for seq_id in averaged_inferences.keys():174        original_seq_ids.add("_".join(seq_id.split("_")[0:-1]))175    original_seq_ids = list(original_seq_ids)176  177    offset = 0 178    for or_seq_id in sorted(original_seq_ids):179        x_values = []180        y_values = []181 182        # Python3.7+: DICT has the keys sorted by the insertion order183        for i, seq_id in enumerate(list(averaged_inferences.keys())):184            if(or_seq_id == "_".join(seq_id.split("_")[0:-1])):185                x_values.append(i-offset)186                y_values.append(averaged_inferences[seq_id])187 188        y_values = numpy.array(y_values).T189        190        for i, threshold in enumerate(thresholds):191            plt.figure(f"t{threshold} models' per-residue inferences for {seq_id}")192            color = "lightgrey" if(smoothen) else "navy"193            plt.plot(x_values, y_values[i], linewidth=1, color=color)194            plt.xlabel(x_label)195            plt.ylabel("prediction")196            plt.title(f"{title} of {or_seq_id} using threshold {threshold}", wrap=True)197            plt.ylim(bottom=0, top=1)198            j = 0199            y_smoothened_values = []200  201            if(smoothen):202                while j < len(y_values[i])-WINDOW_SIZE+1:203                    window_average = round(numpy.sum(204                        y_values[i][j:j+WINDOW_SIZE])/WINDOW_SIZE, 2)205          206                    y_smoothened_values.append(window_average)207                    j += 1208      209                plt.plot(x_values[int(WINDOW_SIZE/2):-int(WINDOW_SIZE/2)], 210                    y_smoothened_values, linewidth=1, color="navy")211            212            plt.savefig(f"{plot_dir}/{or_seq_id}_per_residue_plot_t{threshold}.svg", format="svg")213 214        offset += len(x_values)215 216def plot_inferences(per_res_out, per_segment_out, averaged_inferences, thresholds, plot_dir,217    window_size, segment_size, smoothen):218    """219    Deciding and calling, which inferences to plot.220 221    per_res_out - STRING or None to determine whether per-residue predictions 222        are required223    per_segment_out - STRING or None to determine whether per-segment 224        predictions are required225    averaged_inferences - DICT that keeps each sequence's inferences 226        (averaged of all threshold models))227    thresholds - LIST with binary models' temperature thresholds228    plot_dir - STRING that determines the directory where plots should 229        be saved230    window_size - INT of the window size for curve smoothening231    segment_size - INT of the segment size of combined residues232    smoothen - BOOL indicates to plot smoothened curve233    """234    if(plot_dir is None): return 235    if(per_res_out):236        plot_per_res_inferences(averaged_inferences, thresholds,237            plot_dir, window_size=window_size)238 239    if(per_segment_out):240        plot_per_res_inferences(averaged_inferences, thresholds,241            plot_dir, smoothen=smoothen,242            window_size=window_size,243            x_label=f"segment (k={segment_size}) index",244        title="Per-segment predictions")245