forked from enviPath/enviPy
minor
This commit is contained in:
@ -2815,18 +2815,55 @@ class PackageBasedModel(EPModel):
|
||||
)
|
||||
multigen_eval = models.BooleanField(null=False, blank=False, default=False)
|
||||
|
||||
def parameters(self):
|
||||
params = {
|
||||
"Model Evaluation Threshold": f"{self.threshold:.2f}",
|
||||
"Multi Gen Evaluation": "Yes" if self.multigen_eval else "No",
|
||||
}
|
||||
|
||||
if self.app_domain:
|
||||
params["Applicability Domain Num Neighbors"] = f"{self.app_domain.num_neighbours:.2f}"
|
||||
params["Applicability Domain Reliability Threshold"] = (
|
||||
f"{self.app_domain.reliability_threshold:.2f}"
|
||||
)
|
||||
params["Applicability Domain Local Compatibility Threshold"] = (
|
||||
f"{self.app_domain.local_compatibilty_threshold:.2f}"
|
||||
)
|
||||
|
||||
return params
|
||||
|
||||
def statistics(self):
|
||||
from sklearn.metrics import auc
|
||||
|
||||
recall = list(self.eval_results["average_recall_per_threshold"].values())
|
||||
precision = list(self.eval_results["average_precision_per_threshold"].values())
|
||||
mg_recall = list(
|
||||
self.eval_results.get("multigen_average_recall_per_threshold", {}).values()
|
||||
)
|
||||
mg_precision = list(
|
||||
self.eval_results.get("multigen_average_precision_per_threshold", {}).values()
|
||||
)
|
||||
|
||||
return {
|
||||
'accuracy': [
|
||||
"accuracy": [
|
||||
self.eval_results["average_accuracy"],
|
||||
self.eval_results.get("multigen_average_accuracy")],
|
||||
'precision': [
|
||||
self.eval_results["average_precision_per_threshold"][f"{self.threshold:.2f}"],
|
||||
self.eval_results.get("multigen_average_precision_per_threshold", {}).get(f"{self.threshold:.2f}")
|
||||
self.eval_results.get("multigen_average_accuracy"),
|
||||
],
|
||||
'recall': [
|
||||
"precision": [
|
||||
self.eval_results["average_precision_per_threshold"][f"{self.threshold:.2f}"],
|
||||
self.eval_results.get("multigen_average_precision_per_threshold", {}).get(
|
||||
f"{self.threshold:.2f}"
|
||||
),
|
||||
],
|
||||
"recall": [
|
||||
self.eval_results["average_recall_per_threshold"][f"{self.threshold:.2f}"],
|
||||
self.eval_results.get("multigen_average_recall_per_threshold", {}).get(f"{self.threshold:.2f}")
|
||||
self.eval_results.get("multigen_average_recall_per_threshold", {}).get(
|
||||
f"{self.threshold:.2f}"
|
||||
),
|
||||
],
|
||||
"Area under PR Curve": [
|
||||
auc(recall, precision),
|
||||
auc(mg_recall, mg_precision) if self.multigen_eval else None,
|
||||
],
|
||||
}
|
||||
|
||||
@ -3053,8 +3090,6 @@ class PackageBasedModel(EPModel):
|
||||
thresholds.append(np.float64(threshold))
|
||||
thresholds.sort()
|
||||
|
||||
logger.info(f"Thresholds: {thresholds}")
|
||||
|
||||
precision = {f"{t:.2f}": [] for t in thresholds}
|
||||
recall = {f"{t:.2f}": [] for t in thresholds}
|
||||
|
||||
@ -3103,14 +3138,18 @@ class PackageBasedModel(EPModel):
|
||||
for t in thresholds:
|
||||
for true, pred in zip(pathways, pred_pathways):
|
||||
acc, pre, rec = multigen_eval(true, pred, t)
|
||||
if abs(t - threshold) < 0.01:
|
||||
mg_acc = acc
|
||||
|
||||
if f"{t:.2f}" == f"{threshold:.2f}":
|
||||
mg_acc += acc
|
||||
|
||||
precision[f"{t:.2f}"].append(pre)
|
||||
recall[f"{t:.2f}"].append(rec)
|
||||
|
||||
avg_mg_acc = mg_acc / len(root_compounds)
|
||||
precision = {k: sum(v) / len(v) if len(v) > 0 else 0 for k, v in precision.items()}
|
||||
recall = {k: sum(v) / len(v) if len(v) > 0 else 0 for k, v in recall.items()}
|
||||
return mg_acc, precision, recall
|
||||
logger.info("Average Multigen Accuracy: {:.2f}".format(avg_mg_acc))
|
||||
return avg_mg_acc, precision, recall
|
||||
|
||||
# If there are eval packages perform single generation evaluation on them instead of random splits
|
||||
if self.eval_packages.count() > 0:
|
||||
|
||||
@ -477,8 +477,7 @@ def batch_predict(
|
||||
limit=None,
|
||||
setting_overrides={
|
||||
"max_nodes": num_tps,
|
||||
"max_depth": num_tps,
|
||||
"model_threshold": 0.001,
|
||||
"model_threshold": 0.0,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user