File size: 3,273 Bytes
95456ed
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
import sys
import unittest
from pprint import pprint
from pathlib import Path
from unittest.mock import patch

SCANDL2_ROOT = Path(__file__).resolve().parents[1]
PROJECT_ROOT = SCANDL2_ROOT.parent
if str(PROJECT_ROOT) not in sys.path:
    sys.path.insert(0, str(PROJECT_ROOT))


class ScanDL2SmokeTests(unittest.TestCase):
    def test_sentence_model_runs_real_inference(self):
        _skip_if_sentence_assets_are_missing()

        from ScanDL2.model import ScanDL2

        try:
            with patch.object(sys, "argv", [sys.argv[0]]):
                model = ScanDL2(text_type="sentence", bsz=1, save=None, filename=None)
                model.eval()
                output = model(["The quick brown fox jumps."])
        except PermissionError as exc:
            raise unittest.SkipTest(f"Real ScanDL2 smoke test needs socket access: {exc}") from exc

        print("\nScanDL2 output:")
        pprint(output)
        self.assertIsInstance(output, dict)
        self.assertIn("predicted_sp_words", output)
        self.assertIn("predicted_sp_ids", output)
        self.assertIn("original_sn", output)
        self.assertIn("predicted_fix_durs", output)
        self.assertIn("unique_idx", output)

    def test_scandl_and_fixdur_modules_run_real_inference(self):
        _skip_if_sentence_assets_are_missing()

        from ScanDL2.model import FixdurModule, ScanDLModule

        try:
            with patch.object(sys, "argv", [sys.argv[0]]):
                scandl_module = ScanDLModule(text_type="sentence", bsz=1)
                scandl_output = scandl_module(texts=["The quick brown fox jumps."])

                fixdur_module = FixdurModule(text_type="sentence", bsz=1)
                fixdur_output = fixdur_module(scandl_module_output=scandl_output)
        except PermissionError as exc:
            raise unittest.SkipTest(f"Real ScanDL2 smoke test needs socket access: {exc}") from exc

        print("\nScanDLModule output:")
        pprint(scandl_output)
        print("\nFixdurModule output:")
        pprint(fixdur_output)

        self.assertIsInstance(scandl_output, dict)
        self.assertIn("predicted_sp_words", scandl_output)
        self.assertIn("predicted_sp_ids", scandl_output)
        self.assertIn("original_sn", scandl_output)
        self.assertIn("unique_idx", scandl_output)

        self.assertIsInstance(fixdur_output, dict)
        self.assertIn("predicted_sp_words", fixdur_output)
        self.assertIn("predicted_sp_ids", fixdur_output)
        self.assertIn("original_sn", fixdur_output)
        self.assertIn("predicted_fix_durs", fixdur_output)
        self.assertIn("unique_idx", fixdur_output)


def _skip_if_sentence_assets_are_missing():
    required_dirs = [
        SCANDL2_ROOT / "models" / "sentence" / "scandl-module",
        SCANDL2_ROOT / "models" / "sentence" / "fixdur-module",
    ]
    missing = [path for path in required_dirs if not path.exists() or not any(path.iterdir())]
    if missing:
        raise unittest.SkipTest(
            "Missing ScanDL2 sentence model assets: "
            + ", ".join(str(path.relative_to(PROJECT_ROOT)) for path in missing)
        )


if __name__ == "__main__":
    unittest.main()