Validating rule accuracy
Evaluate how well the fitted rules predict held-out data using ROC/AUC and classification metrics.
Starter notebook. The cells below are a scaffold: real function calls with placeholder paths/arguments and
TODOmarkers. Fill in your own data and run top-to-bottom. Anything markedTODOis a choice you need to make for your dataset.
Setup and imports
[ ]:
import os
import os.path as op
import time
import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
import seaborn as sns
import bobaT as bb
sns.set_style('whitegrid')
plt.rcParams['figure.dpi'] = 100
import warnings
warnings.filterwarnings('ignore')
Configure paths
[ ]:
# Input data
DATA_DIR = './test_data'
# Output directories
OUTPUT_DIR = './output'
VAL_DIR = './output/validation'
ATTRACTOR_DIR = './output/attractors'
PERTURB_DIR = './output/perturbations'
os.makedirs(OUTPUT_DIR, exist_ok=True)
os.makedirs(VAL_DIR, exist_ok=True)
Load network, rules, and the test set
[ ]:
# Load the base network as a graph-tool graph
network_path = 'tf-lit-network.csv' # TODO: your base network CSV
graph, vertex_dict = bb.load.load_network(
f'{DATA_DIR}/{network_path}',
remove_sinks=False, remove_selfloops=True, remove_sources=False,
)
v_names, nodes = bb.utils.get_nodes(vertex_dict, graph)
rules, regulators_dict = bb.load.load_rules(fname=f'{OUTPUT_DIR}/rules.txt')
# TODO: load the held-out test data saved during inference
data_test = None # e.g. reload from f'{OUTPUT_DIR}/data_split/...'
Score the test set and compute ROC/AUC per node
[ ]:
validation, tprs_all, fprs_all, area_all = bb.tl.fit_validation(
data_test, data_test_t1=None, nodes=nodes,
regulators_dict=regulators_dict, rules=rules,
save=True, save_dir=VAL_DIR, plot=True, show_plots=False, save_df=True,
)
Plot AUCs and averaged ROC
[ ]:
bb.plot.plot_aucs(VAL_DIR, save=True, show_plot=True)
bb.plot.plot_validation_avgs(fprs_all, tprs_all, len(nodes), area_all,
save=True, save_dir=VAL_DIR, show_plot=True)
scikit-learn classification metrics
[ ]:
summary_stats = bb.tl.get_sklearn_metrics(VAL_DIR)
bb.plot.plot_sklearn_metrics(VAL_DIR)
bb.plot.plot_sklearn_summ_stats(summary_stats.drop('max_error', axis=1), VAL_DIR, fname='')