File size: 1,100 Bytes
35d483e | 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 | """Training objectives, metrics, and deterministic trainer."""
# Keep metric imports lightweight: importing ``turn_detection.training.metrics``
# must not require torch. Heavy objects are exposed lazily.
from .metrics import (
binary_classification_metrics,
confusion_counts,
metrics_at_fpr_budgets,
operational_metrics,
sliced_metrics,
threshold_at_max_fpr,
)
__all__ = [
"DistillationLoss",
"DistillationLossConfig",
"MultiTaskLossConfig",
"MultiTaskTurnLoss",
"Trainer",
"TrainerConfig",
"binary_classification_metrics",
"confusion_counts",
"metrics_at_fpr_budgets",
"operational_metrics",
"sliced_metrics",
"threshold_at_max_fpr",
]
def __getattr__(name: str):
if name in {
"DistillationLoss",
"DistillationLossConfig",
"MultiTaskLossConfig",
"MultiTaskTurnLoss",
}:
from . import losses
return getattr(losses, name)
if name in {"Trainer", "TrainerConfig"}:
from . import trainer
return getattr(trainer, name)
raise AttributeError(name)
|