forked from enviPath/enviPy
[Fix] Export IUCLID Properties, Show Model Params and Statistics, Adjust Batch Prediction Settings (#434)
Co-authored-by: Tim Lorsbach <tim@lorsba.ch> Reviewed-on: enviPath/enviPy#434
This commit is contained in:
@ -2162,7 +2162,7 @@ class Pathway(EnviPathModel, AliasMixin, ScenarioMixin, AdditionalInformationMix
|
||||
|
||||
row += [cs.smiles, cs.get_name(), n.depth]
|
||||
|
||||
edges = self.edges.filter(end_nodes__in=[n])
|
||||
edges = self.edges.filter(end_nodes=n)
|
||||
if len(edges):
|
||||
for e in edges:
|
||||
_row = row.copy()
|
||||
@ -2847,6 +2847,58 @@ class PackageBasedModel(EPModel):
|
||||
|
||||
return res
|
||||
|
||||
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": [
|
||||
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}"
|
||||
),
|
||||
],
|
||||
"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}"
|
||||
),
|
||||
],
|
||||
"Area under PR Curve": [
|
||||
auc(recall, precision),
|
||||
auc(mg_recall, mg_precision) if self.multigen_eval else None,
|
||||
],
|
||||
}
|
||||
|
||||
@cached_property
|
||||
def applicable_rules(self) -> List["Rule"]:
|
||||
"""
|
||||
@ -3027,8 +3079,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}
|
||||
|
||||
@ -3077,14 +3127,17 @@ 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
|
||||
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:
|
||||
@ -4216,6 +4269,9 @@ class ClassifierPluginModel(PackageBasedModel):
|
||||
instance = impl(conf)
|
||||
return instance
|
||||
|
||||
def parameters(self):
|
||||
return self.instance().parameters()
|
||||
|
||||
def build_dataset(self):
|
||||
"""
|
||||
Required by general model contract but actual implementation resides in plugin.
|
||||
@ -4436,6 +4492,9 @@ class PropertyPluginModel(PackageBasedModel):
|
||||
instance = impl()
|
||||
return instance
|
||||
|
||||
def parameters(self):
|
||||
return self.instance().parameters()
|
||||
|
||||
def build_dataset(self):
|
||||
"""
|
||||
Required by general model contract but actual implementation resides in plugin.
|
||||
|
||||
Reference in New Issue
Block a user