wfloat 2.0.0__py3-none-win_amd64.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- wfloat/__init__.py +57 -0
- wfloat/__main__.py +5 -0
- wfloat/_assets.py +421 -0
- wfloat/_cache.py +208 -0
- wfloat/_cli.py +74 -0
- wfloat/_constants.py +147 -0
- wfloat/_core.py +1795 -0
- wfloat/_download.py +130 -0
- wfloat/_generated_model_urls.py +46 -0
- wfloat/_llm.py +84 -0
- wfloat/_llm_assets.py +142 -0
- wfloat/_llm_load.py +61 -0
- wfloat/_model.py +394 -0
- wfloat/_native.py +17 -0
- wfloat/_results.py +163 -0
- wfloat/_stt.py +122 -0
- wfloat/_stt_assets.py +144 -0
- wfloat/_stt_load.py +75 -0
- wfloat/_vad.py +119 -0
- wfloat/_vad_assets.py +130 -0
- wfloat/_vad_load.py +97 -0
- wfloat/_version.py +1 -0
- wfloat/native/wfloat-core.dll +0 -0
- wfloat-2.0.0.dist-info/METADATA +251 -0
- wfloat-2.0.0.dist-info/RECORD +28 -0
- wfloat-2.0.0.dist-info/WHEEL +5 -0
- wfloat-2.0.0.dist-info/entry_points.txt +2 -0
- wfloat-2.0.0.dist-info/top_level.txt +1 -0
wfloat/__init__.py
ADDED
|
@@ -0,0 +1,57 @@
|
|
|
1
|
+
from ._constants import SPEAKER_IDS, VALID_EMOTIONS, VALID_SIDS
|
|
2
|
+
from ._llm import LlmModel
|
|
3
|
+
from ._llm_load import load_llm_model
|
|
4
|
+
from ._model import Model, TtsModel, load, load_tts_model
|
|
5
|
+
from ._stt import SttModel, SttSession
|
|
6
|
+
from ._stt_load import load_moonshine_tiny_en, load_stt_model, load_whisper_tiny_en
|
|
7
|
+
from ._vad import VadModel
|
|
8
|
+
from ._vad_load import load_silero_vad, load_vad_model
|
|
9
|
+
from ._results import (
|
|
10
|
+
Audio,
|
|
11
|
+
AudioResult,
|
|
12
|
+
GenerationResult,
|
|
13
|
+
LlmGenerationResult,
|
|
14
|
+
StreamingTranscriptionResult,
|
|
15
|
+
TranscriptionResult,
|
|
16
|
+
TranscriptionSegment,
|
|
17
|
+
TranscriptionToken,
|
|
18
|
+
Timeline,
|
|
19
|
+
TimelineChunk,
|
|
20
|
+
TtsSynthesisResult,
|
|
21
|
+
VadDetectionResult,
|
|
22
|
+
VadSegment,
|
|
23
|
+
)
|
|
24
|
+
from ._version import __version__
|
|
25
|
+
|
|
26
|
+
__all__ = [
|
|
27
|
+
"Audio",
|
|
28
|
+
"AudioResult",
|
|
29
|
+
"GenerationResult",
|
|
30
|
+
"LlmGenerationResult",
|
|
31
|
+
"LlmModel",
|
|
32
|
+
"Model",
|
|
33
|
+
"SPEAKER_IDS",
|
|
34
|
+
"SttModel",
|
|
35
|
+
"SttSession",
|
|
36
|
+
"StreamingTranscriptionResult",
|
|
37
|
+
"TtsModel",
|
|
38
|
+
"Timeline",
|
|
39
|
+
"TimelineChunk",
|
|
40
|
+
"TranscriptionResult",
|
|
41
|
+
"TranscriptionSegment",
|
|
42
|
+
"TranscriptionToken",
|
|
43
|
+
"TtsSynthesisResult",
|
|
44
|
+
"VALID_EMOTIONS",
|
|
45
|
+
"VALID_SIDS",
|
|
46
|
+
"VadDetectionResult",
|
|
47
|
+
"VadModel",
|
|
48
|
+
"VadSegment",
|
|
49
|
+
"load",
|
|
50
|
+
"load_llm_model",
|
|
51
|
+
"load_moonshine_tiny_en",
|
|
52
|
+
"load_silero_vad",
|
|
53
|
+
"load_stt_model",
|
|
54
|
+
"load_whisper_tiny_en",
|
|
55
|
+
"load_tts_model",
|
|
56
|
+
"load_vad_model",
|
|
57
|
+
]
|
wfloat/__main__.py
ADDED
wfloat/_assets.py
ADDED
|
@@ -0,0 +1,421 @@
|
|
|
1
|
+
from dataclasses import dataclass
|
|
2
|
+
from pathlib import Path
|
|
3
|
+
from typing import Dict, Mapping, Optional
|
|
4
|
+
from urllib.parse import urlparse
|
|
5
|
+
|
|
6
|
+
from ._generated_model_urls import MODEL_ASSETS, REGISTRY_ORIGIN, SHARED_ASSETS
|
|
7
|
+
|
|
8
|
+
REGISTRY_BASE_URL = REGISTRY_ORIGIN
|
|
9
|
+
WFLOAT_TTS_MODEL_ID = "wfloat/wfloat-tts"
|
|
10
|
+
SILERO_VAD_MODEL_ID = "snakers4/silero-vad"
|
|
11
|
+
SMOLLM2_360M_INSTRUCT_MODEL_ID = "HuggingFaceTB/SmolLM2-360M-Instruct"
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
@dataclass(frozen=True)
|
|
15
|
+
class ModelAssets:
|
|
16
|
+
model_onnx: str
|
|
17
|
+
model_onnx_checksum: str
|
|
18
|
+
model_tokens: str
|
|
19
|
+
model_tokens_checksum: str
|
|
20
|
+
espeak_data: str
|
|
21
|
+
espeak_checksum: str
|
|
22
|
+
|
|
23
|
+
@classmethod
|
|
24
|
+
def from_dict(cls, data: Dict[str, object]) -> "ModelAssets":
|
|
25
|
+
required_fields = (
|
|
26
|
+
"model_onnx",
|
|
27
|
+
"model_onnx_checksum",
|
|
28
|
+
"model_tokens",
|
|
29
|
+
"model_tokens_checksum",
|
|
30
|
+
"espeak_data",
|
|
31
|
+
"espeak_checksum",
|
|
32
|
+
)
|
|
33
|
+
|
|
34
|
+
missing = [
|
|
35
|
+
field_name
|
|
36
|
+
for field_name in required_fields
|
|
37
|
+
if not isinstance(data.get(field_name), str) or not str(data.get(field_name)).strip()
|
|
38
|
+
]
|
|
39
|
+
if missing:
|
|
40
|
+
raise ValueError(
|
|
41
|
+
"Model asset response is missing required fields: %s"
|
|
42
|
+
% ", ".join(missing)
|
|
43
|
+
)
|
|
44
|
+
|
|
45
|
+
return cls(
|
|
46
|
+
model_onnx=str(data["model_onnx"]),
|
|
47
|
+
model_onnx_checksum=str(data["model_onnx_checksum"]),
|
|
48
|
+
model_tokens=str(data["model_tokens"]),
|
|
49
|
+
model_tokens_checksum=str(data["model_tokens_checksum"]),
|
|
50
|
+
espeak_data=str(data["espeak_data"]),
|
|
51
|
+
espeak_checksum=str(data["espeak_checksum"]),
|
|
52
|
+
)
|
|
53
|
+
|
|
54
|
+
def to_dict(self) -> Dict[str, str]:
|
|
55
|
+
return {
|
|
56
|
+
"model_onnx": self.model_onnx,
|
|
57
|
+
"model_onnx_checksum": self.model_onnx_checksum,
|
|
58
|
+
"model_tokens": self.model_tokens,
|
|
59
|
+
"model_tokens_checksum": self.model_tokens_checksum,
|
|
60
|
+
"espeak_data": self.espeak_data,
|
|
61
|
+
"espeak_checksum": self.espeak_checksum,
|
|
62
|
+
}
|
|
63
|
+
|
|
64
|
+
|
|
65
|
+
@dataclass(frozen=True)
|
|
66
|
+
class SttModelAssets:
|
|
67
|
+
family: str
|
|
68
|
+
tokens: str
|
|
69
|
+
tokens_checksum: Optional[str] = None
|
|
70
|
+
model: Optional[str] = None
|
|
71
|
+
model_checksum: Optional[str] = None
|
|
72
|
+
preprocessor: Optional[str] = None
|
|
73
|
+
preprocessor_checksum: Optional[str] = None
|
|
74
|
+
encoder: Optional[str] = None
|
|
75
|
+
encoder_checksum: Optional[str] = None
|
|
76
|
+
decoder: Optional[str] = None
|
|
77
|
+
decoder_checksum: Optional[str] = None
|
|
78
|
+
joiner: Optional[str] = None
|
|
79
|
+
joiner_checksum: Optional[str] = None
|
|
80
|
+
uncached_decoder: Optional[str] = None
|
|
81
|
+
uncached_decoder_checksum: Optional[str] = None
|
|
82
|
+
cached_decoder: Optional[str] = None
|
|
83
|
+
cached_decoder_checksum: Optional[str] = None
|
|
84
|
+
|
|
85
|
+
@classmethod
|
|
86
|
+
def from_dict(cls, data: Dict[str, object]) -> "SttModelAssets":
|
|
87
|
+
family = str(data.get("family") or "").strip()
|
|
88
|
+
if not family:
|
|
89
|
+
raise ValueError("STT asset response is missing required field: family")
|
|
90
|
+
|
|
91
|
+
files = data.get("files")
|
|
92
|
+
if isinstance(files, Mapping):
|
|
93
|
+
normalized: Dict[str, str] = {}
|
|
94
|
+
for key, value in files.items():
|
|
95
|
+
if not isinstance(key, str) or not isinstance(value, Mapping):
|
|
96
|
+
continue
|
|
97
|
+
url = value.get("url")
|
|
98
|
+
checksum = value.get("checksum")
|
|
99
|
+
if isinstance(url, str) and url.strip():
|
|
100
|
+
normalized[key] = url
|
|
101
|
+
if isinstance(checksum, str) and checksum.strip():
|
|
102
|
+
normalized[f"{key}_checksum"] = checksum
|
|
103
|
+
merged = dict(data)
|
|
104
|
+
merged.update(normalized)
|
|
105
|
+
data = merged
|
|
106
|
+
|
|
107
|
+
required_fields = ("tokens",)
|
|
108
|
+
missing = [
|
|
109
|
+
field_name
|
|
110
|
+
for field_name in required_fields
|
|
111
|
+
if not isinstance(data.get(field_name), str)
|
|
112
|
+
or not str(data.get(field_name)).strip()
|
|
113
|
+
]
|
|
114
|
+
if missing:
|
|
115
|
+
raise ValueError(
|
|
116
|
+
"STT asset response is missing required fields: %s"
|
|
117
|
+
% ", ".join(missing)
|
|
118
|
+
)
|
|
119
|
+
|
|
120
|
+
def optional_string(key: str) -> Optional[str]:
|
|
121
|
+
value = data.get(key)
|
|
122
|
+
if isinstance(value, str) and value.strip():
|
|
123
|
+
return value.strip()
|
|
124
|
+
return None
|
|
125
|
+
|
|
126
|
+
return cls(
|
|
127
|
+
family=family,
|
|
128
|
+
tokens=str(data["tokens"]).strip(),
|
|
129
|
+
tokens_checksum=optional_string("tokens_checksum"),
|
|
130
|
+
model=optional_string("model"),
|
|
131
|
+
model_checksum=optional_string("model_checksum"),
|
|
132
|
+
preprocessor=optional_string("preprocessor"),
|
|
133
|
+
preprocessor_checksum=optional_string("preprocessor_checksum"),
|
|
134
|
+
encoder=optional_string("encoder"),
|
|
135
|
+
encoder_checksum=optional_string("encoder_checksum"),
|
|
136
|
+
decoder=optional_string("decoder"),
|
|
137
|
+
decoder_checksum=optional_string("decoder_checksum"),
|
|
138
|
+
joiner=optional_string("joiner"),
|
|
139
|
+
joiner_checksum=optional_string("joiner_checksum"),
|
|
140
|
+
uncached_decoder=optional_string("uncached_decoder"),
|
|
141
|
+
uncached_decoder_checksum=optional_string("uncached_decoder_checksum"),
|
|
142
|
+
cached_decoder=optional_string("cached_decoder"),
|
|
143
|
+
cached_decoder_checksum=optional_string("cached_decoder_checksum"),
|
|
144
|
+
)
|
|
145
|
+
|
|
146
|
+
def to_dict(self) -> Dict[str, str]:
|
|
147
|
+
data = {
|
|
148
|
+
"family": self.family,
|
|
149
|
+
"tokens": self.tokens,
|
|
150
|
+
}
|
|
151
|
+
optional_fields = {
|
|
152
|
+
"tokens_checksum": self.tokens_checksum,
|
|
153
|
+
"model": self.model,
|
|
154
|
+
"model_checksum": self.model_checksum,
|
|
155
|
+
"preprocessor": self.preprocessor,
|
|
156
|
+
"preprocessor_checksum": self.preprocessor_checksum,
|
|
157
|
+
"encoder": self.encoder,
|
|
158
|
+
"encoder_checksum": self.encoder_checksum,
|
|
159
|
+
"decoder": self.decoder,
|
|
160
|
+
"decoder_checksum": self.decoder_checksum,
|
|
161
|
+
"joiner": self.joiner,
|
|
162
|
+
"joiner_checksum": self.joiner_checksum,
|
|
163
|
+
"uncached_decoder": self.uncached_decoder,
|
|
164
|
+
"uncached_decoder_checksum": self.uncached_decoder_checksum,
|
|
165
|
+
"cached_decoder": self.cached_decoder,
|
|
166
|
+
"cached_decoder_checksum": self.cached_decoder_checksum,
|
|
167
|
+
}
|
|
168
|
+
for key, value in optional_fields.items():
|
|
169
|
+
if value:
|
|
170
|
+
data[key] = value
|
|
171
|
+
return data
|
|
172
|
+
|
|
173
|
+
|
|
174
|
+
@dataclass(frozen=True)
|
|
175
|
+
class VadModelAssets:
|
|
176
|
+
family: str
|
|
177
|
+
model: str
|
|
178
|
+
model_checksum: Optional[str] = None
|
|
179
|
+
|
|
180
|
+
@classmethod
|
|
181
|
+
def from_dict(cls, data: Dict[str, object]) -> "VadModelAssets":
|
|
182
|
+
family = str(data.get("family") or "").strip()
|
|
183
|
+
if not family:
|
|
184
|
+
raise ValueError("VAD asset response is missing required field: family")
|
|
185
|
+
|
|
186
|
+
files = data.get("files")
|
|
187
|
+
if isinstance(files, Mapping):
|
|
188
|
+
model_file = files.get("model")
|
|
189
|
+
if isinstance(model_file, Mapping):
|
|
190
|
+
merged = dict(data)
|
|
191
|
+
url = model_file.get("url")
|
|
192
|
+
checksum = model_file.get("checksum")
|
|
193
|
+
if isinstance(url, str) and url.strip():
|
|
194
|
+
merged["model"] = url
|
|
195
|
+
if isinstance(checksum, str) and checksum.strip():
|
|
196
|
+
merged["model_checksum"] = checksum
|
|
197
|
+
data = merged
|
|
198
|
+
|
|
199
|
+
model = str(data.get("model") or "").strip()
|
|
200
|
+
if not model:
|
|
201
|
+
raise ValueError("VAD asset response is missing required field: model")
|
|
202
|
+
|
|
203
|
+
def optional_string(key: str) -> Optional[str]:
|
|
204
|
+
value = data.get(key)
|
|
205
|
+
if isinstance(value, str) and value.strip():
|
|
206
|
+
return value.strip()
|
|
207
|
+
return None
|
|
208
|
+
|
|
209
|
+
return cls(
|
|
210
|
+
family=family,
|
|
211
|
+
model=model,
|
|
212
|
+
model_checksum=optional_string("model_checksum"),
|
|
213
|
+
)
|
|
214
|
+
|
|
215
|
+
def to_dict(self) -> Dict[str, str]:
|
|
216
|
+
data = {
|
|
217
|
+
"family": self.family,
|
|
218
|
+
"model": self.model,
|
|
219
|
+
}
|
|
220
|
+
if self.model_checksum:
|
|
221
|
+
data["model_checksum"] = self.model_checksum
|
|
222
|
+
return data
|
|
223
|
+
|
|
224
|
+
|
|
225
|
+
@dataclass(frozen=True)
|
|
226
|
+
class LlmModelAssets:
|
|
227
|
+
family: str
|
|
228
|
+
model: str
|
|
229
|
+
model_checksum: Optional[str] = None
|
|
230
|
+
context_size: Optional[int] = None
|
|
231
|
+
chat_template: Optional[str] = None
|
|
232
|
+
chat_template_format: Optional[str] = None
|
|
233
|
+
|
|
234
|
+
@classmethod
|
|
235
|
+
def from_dict(cls, data: Dict[str, object]) -> "LlmModelAssets":
|
|
236
|
+
family = str(data.get("family") or "").strip()
|
|
237
|
+
if not family:
|
|
238
|
+
raise ValueError("LLM asset response is missing required field: family")
|
|
239
|
+
|
|
240
|
+
files = data.get("files")
|
|
241
|
+
if isinstance(files, Mapping):
|
|
242
|
+
model_file = files.get("model")
|
|
243
|
+
if isinstance(model_file, Mapping):
|
|
244
|
+
merged = dict(data)
|
|
245
|
+
url = model_file.get("url")
|
|
246
|
+
checksum = model_file.get("checksum")
|
|
247
|
+
if isinstance(url, str) and url.strip():
|
|
248
|
+
merged["model"] = url
|
|
249
|
+
if isinstance(checksum, str) and checksum.strip():
|
|
250
|
+
merged["model_checksum"] = checksum
|
|
251
|
+
data = merged
|
|
252
|
+
|
|
253
|
+
model = str(data.get("model") or "").strip()
|
|
254
|
+
if not model:
|
|
255
|
+
raise ValueError("LLM asset response is missing required field: model")
|
|
256
|
+
|
|
257
|
+
def optional_string(key: str) -> Optional[str]:
|
|
258
|
+
value = data.get(key)
|
|
259
|
+
if isinstance(value, str) and value.strip():
|
|
260
|
+
return value.strip()
|
|
261
|
+
return None
|
|
262
|
+
|
|
263
|
+
context_size_value = data.get("context_size")
|
|
264
|
+
context_size = None
|
|
265
|
+
if context_size_value is not None:
|
|
266
|
+
context_size = int(context_size_value)
|
|
267
|
+
|
|
268
|
+
return cls(
|
|
269
|
+
family=family,
|
|
270
|
+
model=model,
|
|
271
|
+
model_checksum=optional_string("model_checksum"),
|
|
272
|
+
context_size=context_size,
|
|
273
|
+
chat_template=optional_string("chat_template"),
|
|
274
|
+
chat_template_format=optional_string("chat_template_format"),
|
|
275
|
+
)
|
|
276
|
+
|
|
277
|
+
def to_dict(self) -> Dict[str, object]:
|
|
278
|
+
data: Dict[str, object] = {
|
|
279
|
+
"family": self.family,
|
|
280
|
+
"model": self.model,
|
|
281
|
+
}
|
|
282
|
+
if self.model_checksum:
|
|
283
|
+
data["model_checksum"] = self.model_checksum
|
|
284
|
+
if self.context_size is not None:
|
|
285
|
+
data["context_size"] = self.context_size
|
|
286
|
+
if self.chat_template:
|
|
287
|
+
data["chat_template"] = self.chat_template
|
|
288
|
+
if self.chat_template_format:
|
|
289
|
+
data["chat_template_format"] = self.chat_template_format
|
|
290
|
+
return data
|
|
291
|
+
|
|
292
|
+
|
|
293
|
+
def filename_from_url(url: str, fallback: str) -> str:
|
|
294
|
+
parsed = urlparse(url)
|
|
295
|
+
filename = Path(parsed.path).name
|
|
296
|
+
return filename or fallback
|
|
297
|
+
|
|
298
|
+
|
|
299
|
+
def _registry_url(asset: Mapping[str, object]) -> str:
|
|
300
|
+
path = asset.get("path")
|
|
301
|
+
if not isinstance(path, str) or not path.startswith("/"):
|
|
302
|
+
raise RuntimeError("Registry asset is missing a valid path.")
|
|
303
|
+
return REGISTRY_ORIGIN + path
|
|
304
|
+
|
|
305
|
+
|
|
306
|
+
def _registry_checksum(asset: Mapping[str, object]) -> Optional[str]:
|
|
307
|
+
checksum = asset.get("sha256")
|
|
308
|
+
if isinstance(checksum, str) and checksum.strip():
|
|
309
|
+
return checksum
|
|
310
|
+
return None
|
|
311
|
+
|
|
312
|
+
|
|
313
|
+
def _required_checksum(asset: Mapping[str, object], name: str) -> str:
|
|
314
|
+
checksum = _registry_checksum(asset)
|
|
315
|
+
if checksum is None:
|
|
316
|
+
raise RuntimeError(f"Registry asset is missing required checksum: {name}")
|
|
317
|
+
return checksum
|
|
318
|
+
|
|
319
|
+
|
|
320
|
+
def _model_assets(model_name: str) -> Mapping[str, object]:
|
|
321
|
+
data = MODEL_ASSETS.get(model_name)
|
|
322
|
+
if not isinstance(data, Mapping):
|
|
323
|
+
raise ValueError("Unsupported model: %s" % model_name)
|
|
324
|
+
return data
|
|
325
|
+
|
|
326
|
+
|
|
327
|
+
def _file_asset(data: Mapping[str, object], name: str) -> Mapping[str, object]:
|
|
328
|
+
asset = data.get(name)
|
|
329
|
+
if not isinstance(asset, Mapping):
|
|
330
|
+
raise RuntimeError(f"Registry model entry is missing asset: {name}")
|
|
331
|
+
return asset
|
|
332
|
+
|
|
333
|
+
|
|
334
|
+
def _optional_file_url(data: Mapping[str, object], name: str) -> Optional[str]:
|
|
335
|
+
asset = data.get(name)
|
|
336
|
+
if not isinstance(asset, Mapping):
|
|
337
|
+
return None
|
|
338
|
+
return _registry_url(asset)
|
|
339
|
+
|
|
340
|
+
|
|
341
|
+
def _optional_file_checksum(data: Mapping[str, object], name: str) -> Optional[str]:
|
|
342
|
+
asset = data.get(name)
|
|
343
|
+
if not isinstance(asset, Mapping):
|
|
344
|
+
return None
|
|
345
|
+
return _registry_checksum(asset)
|
|
346
|
+
|
|
347
|
+
|
|
348
|
+
def fetch_model_assets(model_name: str) -> ModelAssets:
|
|
349
|
+
if model_name != WFLOAT_TTS_MODEL_ID:
|
|
350
|
+
raise ValueError("Unsupported TTS model: %s" % model_name)
|
|
351
|
+
|
|
352
|
+
model_assets = _model_assets(model_name)
|
|
353
|
+
model_onnx = _file_asset(model_assets, "model_onnx")
|
|
354
|
+
model_tokens = _file_asset(model_assets, "model_tokens")
|
|
355
|
+
espeak_data = SHARED_ASSETS["espeak_ng_data_zip"]
|
|
356
|
+
|
|
357
|
+
return ModelAssets(
|
|
358
|
+
model_onnx=_registry_url(model_onnx),
|
|
359
|
+
model_onnx_checksum=_required_checksum(model_onnx, "model_onnx"),
|
|
360
|
+
model_tokens=_registry_url(model_tokens),
|
|
361
|
+
model_tokens_checksum=_required_checksum(model_tokens, "model_tokens"),
|
|
362
|
+
espeak_data=_registry_url(espeak_data),
|
|
363
|
+
espeak_checksum=_required_checksum(espeak_data, "espeak_ng_data_zip"),
|
|
364
|
+
)
|
|
365
|
+
|
|
366
|
+
|
|
367
|
+
def fetch_stt_assets(model_name: str) -> SttModelAssets:
|
|
368
|
+
model_assets = _model_assets(model_name)
|
|
369
|
+
family = model_assets.get("family")
|
|
370
|
+
if family not in {"whisper", "zipformer-transducer", "moonshine"}:
|
|
371
|
+
raise ValueError("Unsupported STT model: %s" % model_name)
|
|
372
|
+
|
|
373
|
+
return SttModelAssets(
|
|
374
|
+
family=str(family),
|
|
375
|
+
model=_optional_file_url(model_assets, "model"),
|
|
376
|
+
model_checksum=_optional_file_checksum(model_assets, "model"),
|
|
377
|
+
preprocessor=_optional_file_url(model_assets, "preprocessor"),
|
|
378
|
+
preprocessor_checksum=_optional_file_checksum(model_assets, "preprocessor"),
|
|
379
|
+
encoder=_optional_file_url(model_assets, "encoder"),
|
|
380
|
+
encoder_checksum=_optional_file_checksum(model_assets, "encoder"),
|
|
381
|
+
decoder=_optional_file_url(model_assets, "decoder"),
|
|
382
|
+
decoder_checksum=_optional_file_checksum(model_assets, "decoder"),
|
|
383
|
+
joiner=_optional_file_url(model_assets, "joiner"),
|
|
384
|
+
joiner_checksum=_optional_file_checksum(model_assets, "joiner"),
|
|
385
|
+
uncached_decoder=_optional_file_url(model_assets, "uncached_decoder"),
|
|
386
|
+
uncached_decoder_checksum=_optional_file_checksum(model_assets, "uncached_decoder"),
|
|
387
|
+
cached_decoder=_optional_file_url(model_assets, "cached_decoder"),
|
|
388
|
+
cached_decoder_checksum=_optional_file_checksum(model_assets, "cached_decoder"),
|
|
389
|
+
tokens=_registry_url(_file_asset(model_assets, "tokens")),
|
|
390
|
+
tokens_checksum=_optional_file_checksum(model_assets, "tokens"),
|
|
391
|
+
)
|
|
392
|
+
|
|
393
|
+
|
|
394
|
+
def fetch_vad_assets(model_name: str) -> VadModelAssets:
|
|
395
|
+
if model_name != SILERO_VAD_MODEL_ID:
|
|
396
|
+
raise ValueError("Unsupported VAD model: %s" % model_name)
|
|
397
|
+
|
|
398
|
+
model_assets = _model_assets(model_name)
|
|
399
|
+
model = _file_asset(model_assets, "model")
|
|
400
|
+
|
|
401
|
+
return VadModelAssets(
|
|
402
|
+
family=str(model_assets.get("family") or "silero-vad"),
|
|
403
|
+
model=_registry_url(model),
|
|
404
|
+
model_checksum=_registry_checksum(model),
|
|
405
|
+
)
|
|
406
|
+
|
|
407
|
+
|
|
408
|
+
def fetch_llm_assets(model_name: str) -> LlmModelAssets:
|
|
409
|
+
if model_name != SMOLLM2_360M_INSTRUCT_MODEL_ID:
|
|
410
|
+
raise ValueError("Unsupported LLM model: %s" % model_name)
|
|
411
|
+
|
|
412
|
+
model_assets = _model_assets(model_name)
|
|
413
|
+
model = _file_asset(model_assets, "model")
|
|
414
|
+
|
|
415
|
+
return LlmModelAssets(
|
|
416
|
+
family=str(model_assets.get("family") or "smollm"),
|
|
417
|
+
model=_registry_url(model),
|
|
418
|
+
model_checksum=_registry_checksum(model),
|
|
419
|
+
context_size=8192,
|
|
420
|
+
chat_template_format="chatml",
|
|
421
|
+
)
|
wfloat/_cache.py
ADDED
|
@@ -0,0 +1,208 @@
|
|
|
1
|
+
import json
|
|
2
|
+
import os
|
|
3
|
+
import shutil
|
|
4
|
+
import tempfile
|
|
5
|
+
import uuid
|
|
6
|
+
from dataclasses import dataclass
|
|
7
|
+
from pathlib import Path
|
|
8
|
+
from typing import Optional
|
|
9
|
+
|
|
10
|
+
from ._assets import ModelAssets, filename_from_url
|
|
11
|
+
from ._download import (
|
|
12
|
+
download_file,
|
|
13
|
+
extract_archive,
|
|
14
|
+
normalize_checksum,
|
|
15
|
+
resolve_extracted_data_directory,
|
|
16
|
+
verify_checksum,
|
|
17
|
+
)
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
@dataclass(frozen=True)
|
|
21
|
+
class CachedModelAssets:
|
|
22
|
+
model_name: str
|
|
23
|
+
cache_dir: Path
|
|
24
|
+
model_path: Path
|
|
25
|
+
tokens_path: Path
|
|
26
|
+
espeak_data_dir: Path
|
|
27
|
+
manifest_path: Path
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
def get_default_cache_dir() -> Path:
|
|
31
|
+
if os.name == "nt":
|
|
32
|
+
local_appdata = os.environ.get("LOCALAPPDATA")
|
|
33
|
+
if local_appdata:
|
|
34
|
+
return Path(local_appdata) / "wfloat" / "Cache"
|
|
35
|
+
return Path.home() / "AppData" / "Local" / "wfloat" / "Cache"
|
|
36
|
+
|
|
37
|
+
if sys_platform_startswith("darwin"):
|
|
38
|
+
return Path.home() / "Library" / "Caches" / "wfloat"
|
|
39
|
+
|
|
40
|
+
xdg_cache_home = os.environ.get("XDG_CACHE_HOME")
|
|
41
|
+
if xdg_cache_home:
|
|
42
|
+
return Path(xdg_cache_home) / "wfloat"
|
|
43
|
+
|
|
44
|
+
return Path.home() / ".cache" / "wfloat"
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
def sys_platform_startswith(prefix: str) -> bool:
|
|
48
|
+
import sys
|
|
49
|
+
|
|
50
|
+
return sys.platform.startswith(prefix)
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
def normalize_model_name(model_name: str) -> str:
|
|
54
|
+
normalized = model_name.strip().replace("\\", "/")
|
|
55
|
+
normalized = normalized.replace("/", "--")
|
|
56
|
+
normalized = normalized.replace(" ", "-")
|
|
57
|
+
return normalized
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
def _ensure_directory(path: Path) -> None:
|
|
61
|
+
path.mkdir(parents=True, exist_ok=True)
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
def _cleanup_stale_model_files(model_dir: Path, active_names) -> None:
|
|
65
|
+
if not model_dir.exists():
|
|
66
|
+
return
|
|
67
|
+
|
|
68
|
+
for child in model_dir.iterdir():
|
|
69
|
+
if child.name in active_names:
|
|
70
|
+
continue
|
|
71
|
+
if child.is_file():
|
|
72
|
+
child.unlink()
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
def _write_manifest(manifest_path: Path, model_name: str, assets: ModelAssets) -> None:
|
|
76
|
+
manifest_payload = {
|
|
77
|
+
"model_name": model_name,
|
|
78
|
+
"assets": assets.to_dict(),
|
|
79
|
+
}
|
|
80
|
+
manifest_path.write_text(json.dumps(manifest_payload, indent=2, sort_keys=True) + "\n")
|
|
81
|
+
|
|
82
|
+
|
|
83
|
+
def _ensure_cached_file(
|
|
84
|
+
source_url: str,
|
|
85
|
+
checksum: str,
|
|
86
|
+
destination: Path,
|
|
87
|
+
downloads_dir: Path,
|
|
88
|
+
*,
|
|
89
|
+
force_download: bool,
|
|
90
|
+
) -> Path:
|
|
91
|
+
if not force_download and verify_checksum(destination, checksum):
|
|
92
|
+
return destination
|
|
93
|
+
|
|
94
|
+
suffix = destination.suffix or ".bin"
|
|
95
|
+
temp_download = downloads_dir / (uuid.uuid4().hex + suffix)
|
|
96
|
+
download_file(source_url, temp_download, expected_checksum=checksum)
|
|
97
|
+
_ensure_directory(destination.parent)
|
|
98
|
+
os.replace(str(temp_download), str(destination))
|
|
99
|
+
return destination
|
|
100
|
+
|
|
101
|
+
|
|
102
|
+
def _install_espeak_data(
|
|
103
|
+
assets: ModelAssets,
|
|
104
|
+
cache_root: Path,
|
|
105
|
+
downloads_dir: Path,
|
|
106
|
+
*,
|
|
107
|
+
force_download: bool,
|
|
108
|
+
) -> Path:
|
|
109
|
+
checksum = normalize_checksum(assets.espeak_checksum)
|
|
110
|
+
espeak_root = cache_root / "espeak" / checksum
|
|
111
|
+
data_dir = espeak_root / "espeak-ng-data"
|
|
112
|
+
ready_marker = espeak_root / ".ready"
|
|
113
|
+
|
|
114
|
+
if (
|
|
115
|
+
not force_download
|
|
116
|
+
and ready_marker.is_file()
|
|
117
|
+
and data_dir.is_dir()
|
|
118
|
+
):
|
|
119
|
+
return data_dir
|
|
120
|
+
|
|
121
|
+
if espeak_root.exists():
|
|
122
|
+
shutil.rmtree(espeak_root)
|
|
123
|
+
|
|
124
|
+
_ensure_directory(espeak_root)
|
|
125
|
+
archive_name = filename_from_url(assets.espeak_data, checksum + ".zip")
|
|
126
|
+
temp_archive = downloads_dir / (uuid.uuid4().hex + "-" + archive_name)
|
|
127
|
+
download_file(
|
|
128
|
+
assets.espeak_data,
|
|
129
|
+
temp_archive,
|
|
130
|
+
expected_checksum=assets.espeak_checksum,
|
|
131
|
+
)
|
|
132
|
+
|
|
133
|
+
extraction_root = Path(
|
|
134
|
+
tempfile.mkdtemp(prefix="wfloat-espeak-", dir=str(downloads_dir))
|
|
135
|
+
)
|
|
136
|
+
try:
|
|
137
|
+
extract_archive(temp_archive, extraction_root)
|
|
138
|
+
resolved_data_dir = resolve_extracted_data_directory(extraction_root)
|
|
139
|
+
if data_dir.exists():
|
|
140
|
+
shutil.rmtree(data_dir)
|
|
141
|
+
shutil.copytree(str(resolved_data_dir), str(data_dir))
|
|
142
|
+
ready_marker.write_text("ready\n")
|
|
143
|
+
finally:
|
|
144
|
+
if temp_archive.exists():
|
|
145
|
+
temp_archive.unlink()
|
|
146
|
+
if extraction_root.exists():
|
|
147
|
+
shutil.rmtree(extraction_root)
|
|
148
|
+
|
|
149
|
+
return data_dir
|
|
150
|
+
|
|
151
|
+
|
|
152
|
+
def cache_model_assets(
|
|
153
|
+
model_name: str,
|
|
154
|
+
assets: ModelAssets,
|
|
155
|
+
*,
|
|
156
|
+
cache_dir: Optional[Path] = None,
|
|
157
|
+
force_download: bool = False,
|
|
158
|
+
) -> CachedModelAssets:
|
|
159
|
+
cache_root = Path(cache_dir) if cache_dir is not None else get_default_cache_dir()
|
|
160
|
+
models_dir = cache_root / "models"
|
|
161
|
+
downloads_dir = cache_root / "downloads"
|
|
162
|
+
|
|
163
|
+
_ensure_directory(models_dir)
|
|
164
|
+
_ensure_directory(downloads_dir)
|
|
165
|
+
|
|
166
|
+
model_dir = models_dir / normalize_model_name(model_name)
|
|
167
|
+
_ensure_directory(model_dir)
|
|
168
|
+
|
|
169
|
+
model_filename = filename_from_url(assets.model_onnx, "model.onnx")
|
|
170
|
+
tokens_filename = filename_from_url(assets.model_tokens, "tokens.txt")
|
|
171
|
+
manifest_path = model_dir / "manifest.json"
|
|
172
|
+
|
|
173
|
+
_cleanup_stale_model_files(
|
|
174
|
+
model_dir,
|
|
175
|
+
active_names={model_filename, tokens_filename, "manifest.json"},
|
|
176
|
+
)
|
|
177
|
+
|
|
178
|
+
model_path = _ensure_cached_file(
|
|
179
|
+
assets.model_onnx,
|
|
180
|
+
assets.model_onnx_checksum,
|
|
181
|
+
model_dir / model_filename,
|
|
182
|
+
downloads_dir,
|
|
183
|
+
force_download=force_download,
|
|
184
|
+
)
|
|
185
|
+
tokens_path = _ensure_cached_file(
|
|
186
|
+
assets.model_tokens,
|
|
187
|
+
assets.model_tokens_checksum,
|
|
188
|
+
model_dir / tokens_filename,
|
|
189
|
+
downloads_dir,
|
|
190
|
+
force_download=force_download,
|
|
191
|
+
)
|
|
192
|
+
espeak_data_dir = _install_espeak_data(
|
|
193
|
+
assets,
|
|
194
|
+
cache_root,
|
|
195
|
+
downloads_dir,
|
|
196
|
+
force_download=force_download,
|
|
197
|
+
)
|
|
198
|
+
|
|
199
|
+
_write_manifest(manifest_path, model_name, assets)
|
|
200
|
+
|
|
201
|
+
return CachedModelAssets(
|
|
202
|
+
model_name=model_name,
|
|
203
|
+
cache_dir=cache_root,
|
|
204
|
+
model_path=model_path,
|
|
205
|
+
tokens_path=tokens_path,
|
|
206
|
+
espeak_data_dir=espeak_data_dir,
|
|
207
|
+
manifest_path=manifest_path,
|
|
208
|
+
)
|