Spaces:
Paused
Paused
File size: 2,935 Bytes
5716801 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 | from typing import Any, Type
from chemprop.nn import loss, predictors
__all__ = ["pop_attr"]
def pop_attr(o: object, attr: str, *args) -> Any | None:
"""like ``pop()`` but for attribute maps"""
match len(args):
case 0:
return _pop_attr(o, attr)
case 1:
return _pop_attr_d(o, attr, args[0])
case _:
raise TypeError(f"Expected at most 2 arguments! got: {len(args)}")
def _pop_attr(o: object, attr: str) -> Any:
val = getattr(o, attr)
delattr(o, attr)
return val
def _pop_attr_d(o: object, attr: str, default: Any | None = None) -> Any | None:
try:
val = getattr(o, attr)
delattr(o, attr)
except AttributeError:
val = default
return val
def validate_loss_function(
predictor_ffn: Type[predictors._FFNPredictorBase], criterion: Type[loss.LossFunction]
):
match predictor_ffn:
case predictors.RegressionFFN:
if criterion not in (loss.MSELoss, loss.BoundedMSELoss):
raise ValueError(f"Expected a regression loss function! got: {criterion.__name__}")
case predictors.MveFFN:
if criterion is not loss.MVELoss:
raise ValueError(f"Expected a MVE loss function! got: {criterion.__name__}")
case predictors.EvidentialFFN:
if criterion is not loss.EvidentialLoss:
raise ValueError(f"Expected an evidential loss function! got: {criterion.__name__}")
case predictors.BinaryClassificationFFN:
if criterion not in (loss.BCELoss, loss.BinaryMCCLoss):
raise ValueError(
f"Expected a binary classification loss function! got: {criterion.__name__}"
)
case predictors.BinaryDirichletFFN:
if loss is not loss.BinaryDirichletLoss:
raise ValueError(
f"Expected a binary Dirichlet loss function! got: {criterion.__name__}"
)
case predictors.MulticlassClassificationFFN:
if loss not in (loss.CrossEntropyLoss, loss.MulticlassMCCLoss):
raise ValueError(
f"Expected a multiclass classification loss function! got: {criterion.__name__}"
)
case predictors.MulticlassDirichletFFN:
if loss is not loss.MulticlassDirichletLoss:
raise ValueError(
f"Expected a multiclass Dirichlet loss function! got: {criterion.__name__}"
)
case predictors.SpectralFFN:
if loss not in (loss.SIDLoss, loss.WassersteinLoss):
raise ValueError(f"Expected a spectral loss function! got: {criterion.__name__}")
case _:
raise ValueError(
f"Unknown predictor function! got: {predictor_ffn}. "
f"Expected one of: {tuple(predictors.PredictorRegistry.values())}"
)
|