tiny-hinglish-turn-detector / tests /test_model_losses.py
suvradeepp's picture
Publish Tiny Hinglish Turn Detector development preview
35d483e verified
Raw
History Blame Contribute Delete
1.31 kB
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()