| """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) | |