| from __future__ import annotations |
|
|
| import sys |
| import unittest |
| from pathlib import Path |
|
|
| sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src")) |
|
|
| try: |
| import torch |
| except ImportError: |
| torch = None |
|
|
|
|
| @unittest.skipUnless(torch is not None, "PyTorch is not installed") |
| class MultiTaskLossTest(unittest.TestCase): |
| def test_missing_auxiliary_labels_are_masked(self) -> None: |
| from turn_detection.models.common import TurnDetectionOutput |
| from turn_detection.training.losses import MultiTaskTurnLoss |
|
|
| output = TurnDetectionOutput( |
| endpoint_logits=torch.tensor([0.1, -0.2], requires_grad=True), |
| midfiller_logits=torch.tensor([0.3, 0.4], requires_grad=True), |
| endfiller_logits=torch.tensor([0.5, -0.1], requires_grad=True), |
| ) |
| losses = MultiTaskTurnLoss()( |
| output, |
| torch.tensor([1.0, 0.0]), |
| torch.tensor([-1.0, float("nan")]), |
| torch.tensor([1.0, -1.0]), |
| ) |
| self.assertAlmostEqual(float(losses["midfiller"].detach()), 0.0) |
| self.assertGreater(float(losses["endfiller"].detach()), 0.0) |
| losses["total"].backward() |
| self.assertIsNotNone(output.endpoint_logits.grad) |
|
|
|
|
| if __name__ == "__main__": |
| unittest.main() |
|
|