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: # pragma: no cover 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()