Source code for bobaT.tl

from . import utils as ut
from .plot import *

import pandas as pd
import numpy as np
import resource
import os
import os.path as op
import pickle
import seaborn as sns
from scipy import stats
from graph_tool import all as gt
from graph_tool import GraphView
from sklearn.mixture import GaussianMixture
from scipy.cluster.hierarchy import linkage, dendrogram
from graph_tool.topology import label_components
from matplotlib.patches import Patch
import matplotlib.gridspec as gridspec
import matplotlib.pyplot as plt
import sklearn.model_selection as ms
from scipy.stats import mannwhitneyu
from statsmodels.stats.multitest import multipletests
import json
from sklearn.metrics import *


### ------------ RULE FITTING ----------- ###


[docs] def reorder_binary_decision_tree(old_regulator_order, regulators): n = len(regulators) new_order = [] # Map the old order into the new order old_index_to_new = [regulators.index(i) for i in old_regulator_order] # Loop through the leaves for the new rule for leaf in range(2**n): # Get the binary for this leaf, ordered by the new order binary = ut.idx2binary(leaf, n) # Figure out what this binary would have been in the old order oldbinary = "".join([binary[i] for i in old_index_to_new]) # What leaf was that in the old order? oldleaf = ut.state2idx(oldbinary) # Map that old leaf to the current, reordered leaf new_order.append(oldleaf) return new_order
# If A=f(B,C,D), this checks whether B being ON or OFF has an impact > threshold for any combination of C={ON/OFF} D={ON/OFF} # # `heat` (cells x 2**n, the same soft leaf-membership matrix get_rules/get_rules_scvelo # already computes while fitting the rule) switches tot_dif/signed_tot_dif from a plain # sum over every combination of the *other* regulators (treating every combination as # equally likely) to a data-weighted average, weighted by how often each combination # actually occurs in the real training data (leaf_weights = heat.mean(axis=0)). Default # as of 2026-08: pass heat explicitly to get this data-weighted behavior; heat=None keeps # the original unweighted-sum behavior for any external caller that doesn't have it. # max_dif (used for irrelevant-regulator pruning) is unchanged either way.
[docs] def detect_irrelevant_regulator(regulators, rule, threshold=0.1, heat=None): n = len(regulators) max_difs = [] tot_difs = [] signed_tot_difs = [] leaf_weights = heat.mean(axis=0) if heat is not None else None irrelevant = [] for r, regulator in enumerate(regulators): print(f"...checking if {regulator} is irrelevant") max_dif = 0 tot_dif = 0 signed_tot_dif = 0 total_weight = 0 leaves = ut.get_leaves_of_regulator(2**n, r) for i, j in zip(*leaves): dif = np.abs(rule[j] - rule[i]) max_dif = max(dif, max_dif) if leaf_weights is None: tot_dif = tot_dif + dif signed_tot_dif = signed_tot_dif + rule[j] - rule[i] else: # weight = P(this combination of the OTHER regulators, marginalized over # this one) -- same context whether this regulator is on (leaf j) or off # (leaf i), so summing the two leaves' weights gives that marginal. context_weight = leaf_weights[i] + leaf_weights[j] tot_dif = tot_dif + context_weight * dif signed_tot_dif = signed_tot_dif + context_weight * (rule[j] - rule[i]) total_weight = total_weight + context_weight if leaf_weights is not None and total_weight > 0: tot_dif = tot_dif / total_weight signed_tot_dif = signed_tot_dif / total_weight max_difs.append(max_dif) tot_difs.append(tot_dif) signed_tot_difs.append(signed_tot_dif) if max_dif < threshold: irrelevant.append(regulator) return ( dict(zip(regulators, max_difs)), dict(zip(regulators, tot_difs)), dict(zip(regulators, signed_tot_difs)), ) # added signed_tot_dif to output
# return irrelevant # the user can determine if they want to just save, just show, or save and show the plot
[docs] def get_rules_scvelo( data, data_t1, vertex_dict, plot=False, show_plot=False, save_plot=True, threshold=0, save_dir="rules", hlines=None, ): v_names = dict() # Invert the vertex_dict for vertex_name in list(vertex_dict): v_names[vertex_dict[vertex_name]] = vertex_name nodes = list(vertex_dict) rules = dict() regulators_dict = dict() strengths = pd.DataFrame(index=nodes, columns=nodes) signed_strengths = pd.DataFrame(index=nodes, columns=nodes) for gene in nodes: print(gene) # for each node of the network irrelevant = [] n_irrelevant_new = 0 regulators = [ v_names[v] for v in vertex_dict[gene].in_neighbors() if not v_names[v] in irrelevant ] # Define a set of regulators as the in_neighbors of the node # This breaks when all regulators have been deemed irrelevant, or none have while True: n_irrelevant_old = n_irrelevant_new regulators_dict[gene] = regulators n = len(regulators) # We have to make sure we haven't stripped all the regulators as irrelevant if n > 0: # This becomes the eventual probabilistic rule. It has 2 rows # that describe prob(ON) and prob(OFF). At the end these rows # are normalized to sum to 1, such that the rule becomes # prob(ON) / (prob(ON) + prob(OFF) prob_01 = np.zeros((2, 2**n)) # This is the distribution of how much each sample reflects/constrains each leaf of the Binary Decision Diagram heat = np.ones((data.shape[0], 2**n)) for leaf in range(2**n): if leaf % 50 == 0: print(leaf) binary = ut.idx2binary(leaf, len(regulators)) binary = [{"0": False, "1": True}[i] for i in binary] # Binary becomes a list of lists of T and Fs to represent each column for i, idx in enumerate(data.index): # for each row in data column... # grab that row (df) and the expression value for the current node (left side of rule plot) (val) df = data.loc[idx] val = float(data_t1.loc[idx, gene]) for col, on in enumerate(binary): # for each regulator in each column in decision tree... regulator = regulators[col] # if that regulator is on in the decision tree, multiply the weight in the heatmap for that # row of data and column of tree with a weight that = probability that that node is on in the data # df(regulator) = expression value of regulator in data for that row # multiply for each regulator (parent TF) in leaf if on: heat[i, leaf] *= float(df[regulator]) else: heat[i, leaf] *= 1 - float(df[regulator]) # the probability for that leaf becomes the value of expression (val) times that square in the heatmap # this loops over the rows in the heatmap and keeps multiplying in the weight * expression value prob_01[0, leaf] += ( val * heat[i, leaf] ) # Probabilitiy of being ON prob_01[1, leaf] += (1 - val) * heat[i, leaf] # We weigh each column by adding in a sample with prob=50% and # a weight given by 1-max(weight). So leaves where no samples # had high weight will end up with a high weight of 0.5. For # instance, if the best sample has a weight 0.1 (crappy), the # rule will have a sample added with weight 0.9, and 50% prob. max_heat = 1 - np.max(heat, axis=0) for i in range(prob_01.shape[1]): prob_01[0, i] += max_heat[i] * 0.5 prob_01[1, i] += max_heat[i] * 0.5 # The rule is normalized so that prob(ON)+prob(OFF)=1 rules[gene] = prob_01[0, :] / np.sum(prob_01, axis=0) ( max_regulator_relevance, tot_regulator_relevance, signed_tot_regulator_relevance, ) = detect_irrelevant_regulator( regulators, rules[gene], threshold=threshold, heat=heat ) old_regulator_order = [i for i in regulators] regulators = sorted( regulators, key=lambda x: max_regulator_relevance[x], reverse=True ) if max_regulator_relevance[regulators[-1]] < threshold: irrelevant.append(regulators[-1]) old_regulator_order.remove(regulators[-1]) regulators.remove(regulators[-1]) regulators = sorted( regulators, key=lambda x: tot_regulator_relevance[x], reverse=True ) regulators_dict[gene] = regulators # regulators = old_regulator_order # irrelevant += detect_irrelevant_regulator(regulators, rules[gene], threshold=threshold) n_irrelevant_new = len(irrelevant) if len(regulators) == 0 and gene not in irrelevant: regulators = [ gene, ] regulators_dict[gene] = [ gene, ] elif n_irrelevant_old == n_irrelevant_new or len(regulators) == 0: break if len(regulators) > 0: importance_order = reorder_binary_decision_tree( old_regulator_order, regulators ) heat = heat[:, importance_order] rules[gene] = rules[gene][importance_order] # rules[gene] = smooth_rule(rules[gene], regulators, tot_regulator_relevance, np.max(heat,axis=0)) # strengths and signed_strengths should have child nodes as rows with columns as parent nodes strengths.loc[gene] = tot_regulator_relevance signed_strengths.loc[gene] = signed_tot_regulator_relevance if plot: plot_rule( gene, rules[gene], regulators, heat, data, save_dir=save_dir, save=save_plot, show_plot=show_plot, hlines=hlines, ) return rules, regulators_dict, strengths, signed_strengths
# data=dataframe with rows=samples, cols=genes # nodes = list of nodes in network # vertex_dict = dictionary mapping gene name to a vertex in a graph_tool Graph() # v_names - A dictionary mapping vertex in graph to name # plot = boolean - make the resulting plot # threshold = float from 0.0 to 1.0, used as threshold for removing irrelevant regulators. 0 removes nothing. 1 removes all.
[docs] def get_rules( data, vertex_dict, plot=False, threshold=0, save_dir="rules", save_plot=True, show_plot=False, hlines=None, pseudocount_mode="max_heat", pseudocount_c=1.0, pseudocount_target="uniform", ): """ pseudocount_mode: controls how much WEIGHT the pseudo-observation gets for a given leaf. "max_heat" (default, original behavior): weight = 1-max(heat[:, leaf]) -- a leaf only escapes the pull toward the anchor (see pseudocount_target) if at least one single cell is a confident (near-1 heat) match for it. "aggregate": weight = pseudocount_c / (pseudocount_c + sum(heat[:, leaf])), a Beta-style pseudocount tied to the leaf's TOTAL weighted evidence rather than its single best-matching cell -- a leaf with many moderately-confident cells accumulates enough aggregate weight to escape the pull even with no single confident match. On the SCLC 6667 network, this mode alone (isolated from an earlier, unrelated remove_selfloops network-loading mismatch -- barcode 7777) showed no measurable benefit on external validation once that mismatch was corrected; kept as an option, not a recommendation. pseudocount_c: only used when pseudocount_mode="aggregate". Roughly, the number of fully-confident-cells'-worth of aggregate evidence a leaf needs before the prior's pull toward the anchor meaningfully fades. pseudocount_target: controls what VALUE the pseudo-observation is anchored to. "uniform" (default, original behavior): 0.5, i.e. an agnostic prior with no information about which way an under-evidenced leaf should lean. "marginal_rate": the gene's own marginal rate (data[gene].mean(), its overall on/off frequency across all training cells) -- an empirical-Bayes anchor appropriate for single-cell data, where most genes are not naturally balanced 50/50, so pulling a thin-evidence leaf toward 0.5 can be a systematic bias for a gene that's rarely (or almost always) on. """ v_names = dict() for vertex_name in list(vertex_dict): # invert the vertex_dict v_names[vertex_dict[vertex_name]] = vertex_name nodes = list(vertex_dict) rules = dict() regulators_dict = dict() strengths = pd.DataFrame(index=nodes, columns=nodes) signed_strengths = pd.DataFrame(index=nodes, columns=nodes) total_nodes = len(nodes) for xx, gene in enumerate(nodes): print("Fitting ", xx, "/", total_nodes, "rules") print(gene) # for each node of the network irrelevant = [] n_irrelevant_new = 0 regulators = [ v_names[v] for v in vertex_dict[gene].in_neighbors() if not v_names[v] in irrelevant ] # define a set of regulators as the in_neighbors of the node while ( True ): # This breaks when all regulators have been deemed irrelevant, or none have n_irrelevant_old = n_irrelevant_new regulators_dict[gene] = regulators n = len(regulators) # we have to make sure we haven't stripped all the regulators as irrelevant if n > 0: # This becomes the eventual probabilistic rule. It has 2 rows # that describe prob(ON) and prob(OFF). At the end these rows # are normalized to sum to 1, such that the rule becomes # prob(ON) / (prob(ON) + prob(OFF) prob_01 = np.zeros((2, 2**n)) # This is the distribution of how much each sample reflects/constrains each leaf of the Binary Decision Diagram heat = np.ones((data.shape[0], 2**n)) for leaf in range(2**n): if leaf % 100 == 0: print(leaf) binary = ut.idx2binary(leaf, len(regulators)) binary = [{"0": False, "1": True}[i] for i in binary] # binary becomes a list of lists of T and Fs to represent each column for i, idx in enumerate(data.index): # for each row in data column... # grab that row (df) and the expression value for the current node (left side of rule plot) (val) df = data.loc[idx] val = float(data.loc[idx, gene]) for col, on in enumerate(binary): # for each regulator in each column in decision tree... regulator = regulators[col] # if that regulator is on in the decision tree, multiply the weight in the heatmap for that # row of data and column of tree with a weight that = probability that that node is on in the data # df(regulator) = expression value of regulator in data for that row # multiply for each regulator (parent TF) in leaf if on: heat[i, leaf] *= float(df[regulator]) else: heat[i, leaf] *= 1 - float(df[regulator]) # the probability for that leaf becomes the value of expression (val) times that square in the heatmap # this loops over the rows in the heatmap and keeps multiplying in the weight * expression value prob_01[0, leaf] += ( val * heat[i, leaf] ) # Probabilitiy of being ON prob_01[1, leaf] += (1 - val) * heat[i, leaf] # We weigh each column by adding in a sample with prob=50% and # a weight given by 1-max(weight). So leaves where no samples # had high weight will end up with a high weight of 0.5. For # instance, if the best sample has a weight 0.1 (crappy), the # rule will have a sample added with weight 0.9, and 50% prob. if pseudocount_mode == "max_heat": pseudo_weight = 1 - np.max(heat, axis=0) elif pseudocount_mode == "aggregate": total_heat = np.sum(heat, axis=0) pseudo_weight = pseudocount_c / (pseudocount_c + total_heat) else: raise ValueError( f"Unknown pseudocount_mode: {pseudocount_mode!r} " "(expected 'max_heat' or 'aggregate')" ) if pseudocount_target == "uniform": anchor = 0.5 elif pseudocount_target == "marginal_rate": anchor = float(data[gene].mean()) else: raise ValueError( f"Unknown pseudocount_target: {pseudocount_target!r} " "(expected 'uniform' or 'marginal_rate')" ) for i in range(prob_01.shape[1]): prob_01[0, i] += pseudo_weight[i] * anchor prob_01[1, i] += pseudo_weight[i] * (1 - anchor) # The rule is normalized so that prob(ON)+prob(OFF)=1 rules[gene] = prob_01[0, :] / np.sum(prob_01, axis=0) ( max_regulator_relevance, tot_regulator_relevance, signed_tot_regulator_relevance, ) = detect_irrelevant_regulator( regulators, rules[gene], threshold=threshold, heat=heat ) old_regulator_order = [i for i in regulators] regulators = sorted( regulators, key=lambda x: max_regulator_relevance[x], reverse=True ) if max_regulator_relevance[regulators[-1]] < threshold: irrelevant.append(regulators[-1]) old_regulator_order.remove(regulators[-1]) regulators.remove(regulators[-1]) regulators = sorted( regulators, key=lambda x: tot_regulator_relevance[x], reverse=True ) regulators_dict[gene] = regulators # regulators = old_regulator_order # irrelevant += detect_irrelevant_regulator(regulators, rules[gene], threshold=threshold) n_irrelevant_new = len(irrelevant) if len(regulators) == 0 and gene not in irrelevant: regulators = [ gene, ] regulators_dict[gene] = [ gene, ] elif n_irrelevant_old == n_irrelevant_new or len(regulators) == 0: break if len(regulators) > 0: importance_order = reorder_binary_decision_tree( old_regulator_order, regulators ) heat = heat[:, importance_order] rules[gene] = rules[gene][importance_order] # rules[gene] = smooth_rule(rules[gene], regulators, tot_regulator_relevance, np.max(heat,axis=0)) strengths.loc[gene] = tot_regulator_relevance signed_strengths.loc[gene] = signed_tot_regulator_relevance if plot: plot_rule( gene, rules[gene], regulators, heat, data, save_dir=save_dir, save=save_plot, show_plot=show_plot, hlines=hlines, ) return rules, regulators_dict, strengths, signed_strengths
[docs] def save_rules(rules, regulators_dict, fname="rules.txt", delimiter="|"): lines = [] for k in regulators_dict.keys(): rule = ",".join(["%f" % i for i in rules[k]]) regulators = ",".join(regulators_dict[k]) lines.append("%s|%s|%s" % (k, regulators, rule)) outfile = open(fname, "w") outfile.write("\n".join(lines)) outfile.close()
### ------------ FIT VALIDATION ------------ ###
[docs] def parent_heatmap(data, regulators_dict, gene): regulators = [i for i in regulators_dict[gene]] n = len(regulators) # This is the distribution of how much each sample reflects/constrains each leaf of the Binary Decision Diagram heat = np.ones((data.shape[0], 2**n)) for leaf in range(2**n): binary = ut.idx2binary(leaf, len(regulators)) # Binary becomes a list of lists of T and Fs to represent each column binary = [{"0": False, "1": True}[i] for i in binary] for i, idx in enumerate(data.index): # for each row in data column... # grab that row (df) and the expression value for the current node (left side of rule plot) (val) df = data.loc[idx] val = float(data.loc[idx, gene]) for col, on in enumerate(binary): # for each regulator in each column in decision tree... regulator = regulators[col] # if that regulator is on in the decision tree, multiply the weight in the heatmap for that # row of data and column of tree with a weight that = probability that that node is on in the data # df(regulator) = expression value of regulator in data for that row # multiply for each regulator (parent TF) in leaf if on: heat[i, leaf] *= float(df[regulator]) else: heat[i, leaf] *= 1 - float(df[regulator]) regulator_order = [i for i in regulators] return heat, regulator_order
[docs] def roc( validation, node, n_thresholds=10, plot=False, show_plot=False, save=False, save_dir=None, ): tprs = [] fprs = [] for i in np.linspace(0, 1, n_thresholds, endpoint=False): p, r = calc_roc(validation, i) tprs.append(p) fprs.append(r) # area = auc(fprs, tprs) # AUC function wasn't working... replace with np.trapezoid area = np.abs(np.trapezoid(x=fprs, y=tprs)) if plot == True: plot_roc( fprs, tprs, area, node, save=save, save_dir=save_dir, show_plot=show_plot, ) return tprs, fprs, area
# TODO: replace this function with sklearn.metrics.ROC_curve
[docs] def calc_roc(validation, threshold): # P: True positive over predicted condition positive (of the ones predicted positive, how many are actually # positive?) # R: True positive over all condition positive (of the actually positive, how many are predicted to be positive?) predicted = validation.loc[validation["predicted"] > threshold] actual = validation.loc[validation["actual"] > 0.5] predicted_neg = validation.loc[validation["predicted"] <= threshold] actual_neg = validation.loc[validation["actual"] <= 0.5] true_positive = len(set(actual.index).intersection(set(predicted.index))) false_positive = len( set(actual_neg.index).intersection(set(predicted.index))) true_negative = len( set(actual_neg.index).intersection(set(predicted_neg.index))) if len(actual.index.values) == 0 or len(actual_neg.index.values) == 0: return -1, -1 else: # print((true_positive+true_negative)/(len(validation))) tpr = true_positive / len(actual.index) fpr = false_positive / len(actual_neg.index) return tpr, fpr
# this function is broken for some reason UGH
[docs] def auc(fpr, tpr): # fpr is x axis, tpr is y axis print("Calculating area under discrete ROC curve") area = 0 i_old, j_old = 0, 0 for c, i in enumerate(fpr): j = tpr[c] if c == 0: i_old = i j_old = j else: area += np.abs(i - i_old) * j_old + 0.5 * np.abs(i - i_old) * np.abs( j - j_old ) i_old = i j_old = j return area
[docs] def save_auc_by_gene(area_all, nodes, save_dir): outfile = open(f"{save_dir}/aucs.csv", "w+") for n, a in enumerate(area_all): outfile.write(f"{nodes[n]},{a} \n") outfile.close()
# Function to run fit validation # first runs plot_acuracy(_scvelo)() to get validation dataframe # then it calculates roc giving the user the option to plot and save the dataframes and/or plots # val_type = validation type; use the plot scvelo accuracy function or just the plot_accuracy function # fname = optional name to append to default file save name # returns: validation, tprs_all, fprs_all, and area_all
[docs] def fit_validation( data_test, nodes, regulators_dict, rules, data_test_t1=None, save=False, save_dir=None, fname="", clusters=None, plot=True, plot_clusters=False, show_plots=False, save_df=False, n_thresholds=50, customPalette=sns.color_palette("Set2"), ): # create output file if fname != "": outfile = open(f"{save_dir}/tprs_fprs_{fname}.csv", "w+") else: outfile = open(f"{save_dir}/tprs_fprs.csv", "w+") ind = [x for x in np.linspace(0, 1, 50)] tpr_all = pd.DataFrame(index=ind) fpr_all = pd.DataFrame(index=ind) area_all = [] outfile.write(f",,") for j in ind: outfile.write(str(j) + ",") outfile.write("\n") for node in nodes: # print(node) validation = plot_accuracy( data=data_test, node=node, regulators_dict=regulators_dict, rules=rules, data_t1=data_test_t1, plot_clusters=plot_clusters, clusters=clusters, save=save, save_dir=save_dir, show_plot=show_plots, save_df=save_df, customPalette=customPalette, ) tprs, fprs, area = roc( validation, node, n_thresholds=n_thresholds, plot=plot, show_plot=show_plots, save=save, save_dir=save_dir, ) tpr_all[node] = tprs fpr_all[node] = fprs outfile.write(f"{node},tprs,{tprs}\n") outfile.write(f"{node},fprs,{fprs}\n") area_all.append(area) outfile.close() return validation, tpr_all, fpr_all, area_all
# validation_dir = where the validation files were saved; to be read for each node # get roc output from validation files if they exist already? could be helpful for testing so # fit validation doesn't have to be run again
[docs] def roc_from_file( validation_dir, nodes, n_thresholds=50, plot=False, show_plots=False, save=False, save_dir=None, ): ind = [i for i in np.linspace(0, 1, 50)] tpr_all = pd.DataFrame(index=ind) fpr_all = pd.DataFrame(index=ind) area_all = [] for node in nodes: validation = pd.read_csv( f"{validation_dir}/{node}_validation.csv", index_col=0, header=0 ) tprs, fprs, area = roc( validation, node, n_thresholds=n_thresholds, plot=plot, show_plot=show_plots, save=save, save_dir=save_dir, ) tpr_all[node] = tprs fpr_all[node] = fprs area_all.append(area) return tpr_all, fpr_all, area_all
[docs] def get_sklearn_metrics(VAL_DIR, plot_cm=True, show=False, save=True, save_stats=True, verbose=False): files = glob.glob(f"{VAL_DIR}/accuracy_plots/*.csv") summary_stats = pd.DataFrame(columns=['gene', 'accuracy', 'balanced_accuracy_score', 'f1', 'roc_auc_score', "precision", "recall", "explained_variance", 'max_error', 'r2', 'log-loss']) if len(files) == 0: print("You must first run tl.fit_validation() to generate the appropriate files.") return summary_stats else: for f in files: val_df = pd.read_csv(f, header=0, index_col=0) val_df['actual_binary'] = [{True: 1, False: 0}[x] for x in val_df['actual'] > 0.5] val_df['predicted_binary'] = [{True: 1, False: 0}[ x] for x in val_df['predicted'] > 0.5] gene = os.path.basename(f).removesuffix("_validation.csv") if verbose: print(gene) # classification stats acc = accuracy_score( val_df['actual_binary'], val_df['predicted_binary']) bal_acc = balanced_accuracy_score( val_df['actual_binary'], val_df['predicted_binary']) f1 = f1_score(val_df['actual_binary'], val_df['predicted_binary']) prec = precision_score( val_df['actual_binary'], val_df['predicted_binary']) rec = recall_score(val_df['actual_binary'], val_df['predicted_binary']) # regression stats try: roc_auc = roc_auc_score( val_df['actual_binary'], val_df['predicted']) except ValueError: print( f"ValueError for {gene}. Only one class present in y_true. ROC AUC score is not defined in that case.") roc_auc = np.nan expl_var = explained_variance_score( val_df['actual'], val_df['predicted']) max_err = max_error(val_df['actual'], val_df['predicted']) r2 = r2_score(val_df['actual'], val_df['predicted']) try: ll = log_loss(val_df['actual_binary'], val_df['predicted']) except ValueError: print( f"ValueError for {gene}: y_true contains only one label (0). Log-loss is not defined in that case.") ll = np.nan new_row = { 'gene': gene, 'accuracy': acc, 'balanced_accuracy_score': bal_acc, 'f1': f1, 'roc_auc_score': roc_auc, 'precision': prec, 'recall': rec, 'explained_variance': expl_var, 'max_error': max_err, 'r2': r2, 'log-loss': ll } summary_stats = pd.concat( [summary_stats, pd.DataFrame([new_row])], ignore_index=True) # summary_stats = pd.concat([summary_stats, pd.Series([gene, acc, bal_acc, f1, roc_auc, prec, rec, expl_var, max_err, r2, ll], # index=summary_stats.columns)], ignore_index=True) if plot_cm: plt.figure() cm = confusion_matrix( val_df['actual_binary'], val_df['predicted_binary']) disp = ConfusionMatrixDisplay(confusion_matrix=cm) disp.plot(cmap="Blues") plt.title(gene) if show: plt.show() if save: plt.savefig( f"{VAL_DIR}/accuracy_plots/{gene}_confusion_matrix.pdf") plt.close() summary_stats = summary_stats.sort_values( 'gene').reset_index().drop("index", axis=1) if save_stats: summary_stats.to_csv(f"{VAL_DIR}/summary_stats.csv") return summary_stats
#plotting the average AUC curve across ALL SAMPLES in a group (after averaging TPS and FPR for all genes)
[docs] def get_sample_avg_curve(fpr_all, tpr_all, num_nodes, area_all, remove_sources=True, vertex_dict=None, graph=None): if remove_sources: if vertex_dict is None or graph is None: raise ValueError("vertex_dict and graph must be provided if remove_sources is True") v_names_dict = dict() for vertex_name in list(vertex_dict): v_names_dict[vertex_dict[vertex_name]] = vertex_name # root_nodes = [v for v in graph.vertices() if v.in_degree() == 0] sources = [] for v in graph.vertices(): in_neighbors = graph.get_in_neighbors(v) # v can be an int index or a is_self_only = len(in_neighbors) == 1 and in_neighbors[0] == int(v) if len(in_neighbors) == 0 or is_self_only: sources.append(v_names_dict[v]) print("Removing sources (artificially high ROC) from averaged ROCplot: ", sources) ind = ~fpr_all.columns.isin(sources) fpr_all = fpr_all.loc[:, ind] tpr_all = tpr_all.loc[:, ind] area_all = [a for x, a in enumerate(area_all) if ind[x]] num_nodes = num_nodes - len(sources) fpr_avg = (fpr_all.sum(axis=1) / num_nodes).values tpr_avg = (tpr_all.sum(axis=1) / num_nodes).values auc_avg = np.sum(area_all) / num_nodes return fpr_avg, tpr_avg, auc_avg
[docs] def plot_cohort_roc_with_ci( sample_fprs, sample_tprs, n_boot=2000, ci=95, save=False, save_dir=None, show_plot=False, fname="", ): fpr_matrix = np.vstack(sample_fprs) tpr_matrix = np.vstack(sample_tprs) n_samples = fpr_matrix.shape[0] mean_fpr = fpr_matrix.mean(axis=0) mean_tpr = tpr_matrix.mean(axis=0) rng = np.random.default_rng() boot_tpr = np.zeros((n_boot, tpr_matrix.shape[1])) for b in range(n_boot): idx = rng.integers(0, n_samples, n_samples) boot_tpr[b] = tpr_matrix[idx].mean(axis=0) lower_tpr = np.percentile(boot_tpr, (100 - ci) / 2, axis=0) upper_tpr = np.percentile(boot_tpr, 100 - (100 - ci) / 2, axis=0) plt.figure() ax = plt.subplot() plt.plot(mean_fpr, mean_tpr, "-o", label="Mean ROC (across samples)") plt.fill_between(mean_fpr, lower_tpr, upper_tpr, alpha=0.3, label=f"{ci}% CI") ax.plot(ax.get_xlim(), ax.get_ylim(), ls="--", c=".3") plt.xlim(0, 1) plt.ylim(0, 1) plt.ylabel("True Positive Rate") plt.xlabel("False Positive Rate") plt.title(f"Cohort ROC (n={n_samples} samples)\nMean AUC = {np.trapezoid(mean_tpr, mean_fpr):.3f}") plt.legend() if save == True: suffix = f"_{fname}" if fname != "" else "" plt.savefig(f"{save_dir}/ROC_AUC_cohort_CI{suffix}.pdf") if show_plot == True: plt.show() plt.close() return mean_fpr, mean_tpr, lower_tpr, upper_tpr
### ------------ ATTRACTORS ------------ ### # tf_basin --> if -1, use average distance between clusters. otherwise use the same size basin for all phenotypes
[docs] def find_attractors( binarized_data, rules, nodes, regulators_dict, tf_basin, save_dir=None, threshold=0.5, on_nodes=[], off_nodes=[], ): att = dict() n = len(nodes) dist_dict = None if tf_basin < 0: dist_dict = ut.get_avg_min_distance(binarized_data, n) for k in binarized_data.keys(): print(k) att[k] = [] outfile = open( f"{save_dir}/attractors_{k}.txt", "w+" ) # if there are no clusters, comment this out outfile.write("start-state,dist-to-start,attractor\n") start_states = list(binarized_data[k]) cnt = 0 for i in start_states: start_states = [i] # print(start_states) # print("Getting partial STG...") # getting entire stg is too costly, so just get stg out to 5 TF neighborhood if type(tf_basin) == int and tf_basin >= 0: if len(on_nodes) == 0 and len(off_nodes) == 0: stg, edge_weights = ut.get_partial_stg( start_states, rules, nodes, regulators_dict, tf_basin ) else: stg, edge_weights = ut.get_partial_stg( start_states, rules, nodes, regulators_dict, tf_basin, on_nodes=on_nodes, off_nodes=off_nodes, ) # elif type(tf_basin) == dict: elif dist_dict is not None: if len(on_nodes) == 0 and len(off_nodes) == 0: stg, edge_weights = ut.get_partial_stg( start_states, rules, nodes, regulators_dict, dist_dict[k] ) else: stg, edge_weights = ut.get_partial_stg( start_states, rules, nodes, regulators_dict, dist_dict[k], on_nodes=on_nodes, off_nodes=off_nodes, ) else: print( "tf_basin needs to be an integer or < 0 to indicate dictionary of integers for each subtype." ) # Directed stg pruned with threshold .5 # n = number of nodes that can change (each TF gets chosen with equal probability) # EXCEPT nodes that are held ON or OFF (no chance of changing) # each edge actually has a probability of being selected * chance of changing # print("Pruning STG edges...") d_stg = ut.prune_stg_edges( stg, edge_weights, n - len(on_nodes) - len(off_nodes), threshold=threshold, ) # Each strongly connected component becomes a single node # components[2] tells if its an attractor # components[0] of v tells what components does v belong to # print('Condensing STG...') c_stg, c_vertex_dict, components = ut.condense(d_stg) vidx = stg.vertex_properties["idx"] # maps graph_tools made up index in partial stg to an index of state that means something to us # print("Checking for attractors...") for v in stg.vertices(): # loop through every state in stg and if it's an attractor if components[2][components[0][v]]: if v != 0: outfile.write( f"{start_states[0]},{ut.hamming_idx(vidx[v],start_states[0],n)}, {vidx[v]}\n" ) # print(i, int(v), vidx[v]) att[k].append(vidx[v]) cnt += 1 if cnt % 100 == 0: print( "...", np.round( cnt / (len(binarized_data[k])) * 100, 2), "% done" ) outfile.close() for k in att.keys(): att[k] = list(set(att[k])) return att
[docs] def write_attractor_dict(attractor_dict, nodes, outfile): for j in nodes: outfile.write(f",{j}") outfile.write("\n") for k in attractor_dict.keys(): att = [ut.idx2binary(x, len(nodes)) for x in attractor_dict[k]] for i, a in zip(att, attractor_dict[k]): outfile.write(f"{k}") for c in i: outfile.write(f",{c}") outfile.write("\n") outfile.close()
[docs] def filter_attractors( attractor_dir, nodes, clusters ): average_states = {} attractor_dict = {} average_states_df = pd.read_csv( f'{attractor_dir}/average_states.txt', sep=',', header=0, index_col=0) for i, r in average_states_df.iterrows(): s = "" cnt = 0 for letter in list(r): if cnt > 0: s = s+str(letter) cnt += 1 average_states[i] = ut.state2idx(s) for phen in clusters['class'].unique(): d = pd.read_csv( f'{attractor_dir}/attractors_{phen}.txt', sep=',', header=0) attractor_dict[f'{phen}'] = list(np.unique(d['attractor'])) # ##### Below code compares each attractor to average state for each subtype instead of closest single binarized data point a = attractor_dict.copy() # attractor_dict = a.copy() for p in attractor_dict.keys(): print(p) for q in attractor_dict.keys(): if p == q: continue n_same = list(set(attractor_dict[p]).intersection( set(attractor_dict[q]))) if len(n_same) != 0: for x in n_same: p_dist = ut.hamming_idx(x, average_states[p], len(nodes)) q_dist = ut.hamming_idx(x, average_states[q], len(nodes)) if p_dist < q_dist: a[q].remove(x) elif q_dist < p_dist: a[p].remove(x) else: a[q].remove(x) a[p].remove(x) try: a[f'{q}_{p}'].append(x) except KeyError: a[f'{q}_{p}'] = [x] attractor_dict = a print(attractor_dict) file = open(f"{attractor_dir}/attractors_filtered.txt", 'w+') # plot attractors for j in nodes: file.write(f",{j}") file.write("\n") for k in attractor_dict.keys(): att = [ut.idx2binary(x, len(nodes)) for x in attractor_dict[k]] for i, a in zip(att, attractor_dict[k]): file.write(f"{k}") for c in i: file.write(f",{c}") file.write("\n") file.close() return attractor_dict
### ------------ AVG STATES ------------ ###
[docs] def find_avg_states(binarized_data, nodes, save_dir): n = len(nodes) average_states = dict() for k in binarized_data.keys(): ave = ut.average_state(binarized_data[k], n) state = ave.copy() state[state < 0.5] = 0 state[state >= 0.5] = 1 state = [int(i) for i in state] idx = ut.state2idx("".join(["%d" % i for i in state])) average_states[k] = idx file = open(f"{save_dir}/average_states.txt", "w+") ut.get_avg_state_index(nodes, average_states, file, save_dir=save_dir) return average_states
### ------------ PERTURBATIONS SUMMARY ------------ ###
[docs] def perturbations_summary(attractor_dict, perturbations_dir, show=False, save=True, plot_by_attractor=False, save_dir="clustered_perturb_plots", save_full=True, significance='both', fname="", ncols=5, mean_threshold=-0.3): if plot_by_attractor: plot_destabilization_scores( attractor_dict, perturbations_dir, show=False, save=True, clustered=False) print("Plotting perturbation summary plots...") plot_destabilization_scores( attractor_dict, perturbations_dir, show=show, save=save, save_dir=save_dir) print("Testing significance of TF perturbations...") perturb_dict, full = ut.get_perturbation_dict( attractor_dict, perturbations_dir, significance=significance, save_full=False, mean_threshold=mean_threshold) perturb_gene_dict = ut.reverse_perturb_dictionary(perturb_dict) if save_full: ut.write_dict_of_dicts(perturb_gene_dict, file=f"{perturbations_dir}/{save_dir}/perturbation_TF_dictionary{fname}.txt") full_sig = ut.get_ci_sig(full, group_cols=['cluster', 'gene', 'perturb']) if save_full: full_sig.to_csv( f"{perturbations_dir}/{save_dir}/perturbation_stats.csv") plot_perturb_gene_dictionary( perturb_gene_dict, full, perturbations_dir, show=False, save=True, ncols=ncols, fname=fname) return perturb_gene_dict, full, full_sig