OneScience-Group/TemStaPro-main
019
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 