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