voidful commited on
Commit
68d787f
·
1 Parent(s): 347b364

Pin default built-in speaker contract

Browse files
Files changed (1) hide show
  1. tests/test_release_pins.py +92 -0
tests/test_release_pins.py CHANGED
@@ -1,10 +1,24 @@
1
  import ast
2
  import hashlib
3
  from pathlib import Path
 
4
 
5
 
6
  ROOT = Path(__file__).resolve().parents[1]
7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
8
 
9
  def _string_constants(path: Path) -> dict[str, str]:
10
  tree = ast.parse(path.read_text(encoding="utf-8"))
@@ -19,6 +33,84 @@ def _string_constants(path: Path) -> dict[str, str]:
19
  return values
20
 
21
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
22
  def test_remote_model_and_speaker_encoder_are_revision_pinned():
23
  app_path = ROOT / "app.py"
24
  source = app_path.read_text(encoding="utf-8")
 
1
  import ast
2
  import hashlib
3
  from pathlib import Path
4
+ from types import SimpleNamespace
5
 
6
 
7
  ROOT = Path(__file__).resolve().parents[1]
8
 
9
+ FROZEN_SPEAKER_ANCHORS = {
10
+ "aaf1a0878e37875382bb0e5c8a3a2ba43be67297": {
11
+ "speaker_id": "female_voice",
12
+ "speaker_index": 1,
13
+ "ui_label": "內建語者 B",
14
+ "dtype": "float32",
15
+ "shape": (192,),
16
+ "sha256": (
17
+ "e33e4cb6a741d4d1237aa4ff557f1e663d0a6427dccda1f51e83bc149d4188ca"
18
+ ),
19
+ }
20
+ }
21
+
22
 
23
  def _string_constants(path: Path) -> dict[str, str]:
24
  tree = ast.parse(path.read_text(encoding="utf-8"))
 
33
  return values
34
 
35
 
36
+ def _isolated_load_speakers(*, metadata: dict, speaker_ids: tuple[str, ...]):
37
+ """Execute only ``_load_speakers`` without importing the GPU application."""
38
+
39
+ app_path = ROOT / "app.py"
40
+ tree = ast.parse(app_path.read_text(encoding="utf-8"))
41
+ function = next(
42
+ node
43
+ for node in tree.body
44
+ if isinstance(node, ast.FunctionDef) and node.name == "_load_speakers"
45
+ )
46
+ module = ast.Module(
47
+ body=[
48
+ ast.ImportFrom(
49
+ module="__future__",
50
+ names=[ast.alias(name="annotations")],
51
+ level=0,
52
+ ),
53
+ function,
54
+ ],
55
+ type_ignores=[],
56
+ )
57
+ ast.fix_missing_locations(module)
58
+
59
+ load_calls = []
60
+ centroids = tuple(f"test-centroid-{index}" for index in range(len(speaker_ids)))
61
+
62
+ def fake_load(path, **kwargs):
63
+ load_calls.append((path, kwargs))
64
+ return {"speaker_ids": speaker_ids, "centroids": centroids}
65
+
66
+ namespace = {
67
+ "MODEL_DIR": "/pinned/model",
68
+ "METADATA": metadata,
69
+ "os": SimpleNamespace(
70
+ path=SimpleNamespace(
71
+ join=lambda *parts: "/".join(part.strip("/") for part in parts),
72
+ exists=lambda _path: True,
73
+ )
74
+ ),
75
+ "torch": SimpleNamespace(load=fake_load),
76
+ }
77
+ exec(compile(module, app_path, "exec"), namespace)
78
+ result = namespace["_load_speakers"]()
79
+ return result, load_calls
80
+
81
+
82
+ def test_missing_metadata_speaker_id_selects_frozen_female_voice_as_builtin_b():
83
+ constants = _string_constants(ROOT / "app.py")
84
+ contract = FROZEN_SPEAKER_ANCHORS[constants["MODEL_REVISION"]]
85
+ speaker_ids = ("hung_yi_lee", contract["speaker_id"])
86
+
87
+ (labels, default_label), load_calls = _isolated_load_speakers(
88
+ metadata={},
89
+ speaker_ids=speaker_ids,
90
+ )
91
+
92
+ assert contract == {
93
+ "speaker_id": "female_voice",
94
+ "speaker_index": 1,
95
+ "ui_label": "內建語者 B",
96
+ "dtype": "float32",
97
+ "shape": (192,),
98
+ "sha256": (
99
+ "e33e4cb6a741d4d1237aa4ff557f1e663d0a6427dccda1f51e83bc149d4188ca"
100
+ ),
101
+ }
102
+ assert speaker_ids[contract["speaker_index"]] == contract["speaker_id"]
103
+ assert tuple(labels) == ("內建語者 A", contract["ui_label"])
104
+ assert labels[contract["ui_label"]] == "test-centroid-1"
105
+ assert default_label == contract["ui_label"]
106
+ assert load_calls == [
107
+ (
108
+ "pinned/model/checkpoints/speaker_centroids.pt",
109
+ {"map_location": "cpu", "weights_only": True},
110
+ )
111
+ ]
112
+
113
+
114
  def test_remote_model_and_speaker_encoder_are_revision_pinned():
115
  app_path = ROOT / "app.py"
116
  source = app_path.read_text(encoding="utf-8")