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/_download.py
ADDED
|
@@ -0,0 +1,130 @@
|
|
|
1
|
+
import hashlib
|
|
2
|
+
import os
|
|
3
|
+
import shutil
|
|
4
|
+
import tarfile
|
|
5
|
+
import uuid
|
|
6
|
+
import zipfile
|
|
7
|
+
from pathlib import Path
|
|
8
|
+
from typing import Optional
|
|
9
|
+
from urllib.request import Request, urlopen
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
def normalize_checksum(checksum: str) -> str:
|
|
13
|
+
normalized = checksum.strip().lower()
|
|
14
|
+
if normalized.startswith("sha256:"):
|
|
15
|
+
return normalized.split(":", 1)[1]
|
|
16
|
+
if normalized.startswith("sha256-"):
|
|
17
|
+
return normalized.split("-", 1)[1]
|
|
18
|
+
return normalized
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
def sha256_file(path: Path) -> str:
|
|
22
|
+
digest = hashlib.sha256()
|
|
23
|
+
with path.open("rb") as file_obj:
|
|
24
|
+
while True:
|
|
25
|
+
chunk = file_obj.read(1024 * 1024)
|
|
26
|
+
if not chunk:
|
|
27
|
+
break
|
|
28
|
+
digest.update(chunk)
|
|
29
|
+
return digest.hexdigest()
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
def verify_checksum(path: Path, checksum: str) -> bool:
|
|
33
|
+
if not path.is_file():
|
|
34
|
+
return False
|
|
35
|
+
return sha256_file(path) == normalize_checksum(checksum)
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
def download_file(
|
|
39
|
+
url: str,
|
|
40
|
+
destination: Path,
|
|
41
|
+
*,
|
|
42
|
+
expected_checksum: Optional[str] = None,
|
|
43
|
+
timeout: float = 60.0,
|
|
44
|
+
) -> Path:
|
|
45
|
+
destination.parent.mkdir(parents=True, exist_ok=True)
|
|
46
|
+
temp_path = destination.with_name(destination.name + ".tmp-" + uuid.uuid4().hex)
|
|
47
|
+
|
|
48
|
+
request = Request(
|
|
49
|
+
url,
|
|
50
|
+
headers={
|
|
51
|
+
"Accept": "*/*",
|
|
52
|
+
"User-Agent": "wfloat-python/0.0.0",
|
|
53
|
+
},
|
|
54
|
+
method="GET",
|
|
55
|
+
)
|
|
56
|
+
|
|
57
|
+
try:
|
|
58
|
+
with urlopen(request, timeout=timeout) as response:
|
|
59
|
+
with temp_path.open("wb") as file_obj:
|
|
60
|
+
shutil.copyfileobj(response, file_obj)
|
|
61
|
+
|
|
62
|
+
if expected_checksum is not None and not verify_checksum(temp_path, expected_checksum):
|
|
63
|
+
raise RuntimeError(
|
|
64
|
+
"Downloaded file checksum did not match expected value for %s." % url
|
|
65
|
+
)
|
|
66
|
+
|
|
67
|
+
os.replace(str(temp_path), str(destination))
|
|
68
|
+
return destination
|
|
69
|
+
finally:
|
|
70
|
+
if temp_path.exists():
|
|
71
|
+
temp_path.unlink()
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
def _assert_within_destination(path: Path, destination: Path) -> None:
|
|
75
|
+
resolved_destination = destination.resolve()
|
|
76
|
+
resolved_path = path.resolve()
|
|
77
|
+
destination_str = str(resolved_destination)
|
|
78
|
+
path_str = str(resolved_path)
|
|
79
|
+
if path_str != destination_str and not path_str.startswith(destination_str + os.sep):
|
|
80
|
+
raise RuntimeError("Archive member would extract outside destination: %s" % path)
|
|
81
|
+
|
|
82
|
+
|
|
83
|
+
def _extract_zip(archive_path: Path, destination: Path) -> None:
|
|
84
|
+
with zipfile.ZipFile(archive_path) as archive:
|
|
85
|
+
for member in archive.infolist():
|
|
86
|
+
member_path = destination / member.filename
|
|
87
|
+
_assert_within_destination(member_path, destination)
|
|
88
|
+
archive.extractall(destination)
|
|
89
|
+
|
|
90
|
+
|
|
91
|
+
def _extract_tar(archive_path: Path, destination: Path) -> None:
|
|
92
|
+
with tarfile.open(archive_path) as archive:
|
|
93
|
+
for member in archive.getmembers():
|
|
94
|
+
member_path = destination / member.name
|
|
95
|
+
_assert_within_destination(member_path, destination)
|
|
96
|
+
archive.extractall(destination)
|
|
97
|
+
|
|
98
|
+
|
|
99
|
+
def extract_archive(archive_path: Path, destination: Path) -> None:
|
|
100
|
+
destination.mkdir(parents=True, exist_ok=True)
|
|
101
|
+
if zipfile.is_zipfile(archive_path):
|
|
102
|
+
_extract_zip(archive_path, destination)
|
|
103
|
+
return
|
|
104
|
+
|
|
105
|
+
if tarfile.is_tarfile(archive_path):
|
|
106
|
+
_extract_tar(archive_path, destination)
|
|
107
|
+
return
|
|
108
|
+
|
|
109
|
+
raise RuntimeError(
|
|
110
|
+
"Unsupported archive format for %s. Python model assets should provide a zip or tar archive."
|
|
111
|
+
% archive_path
|
|
112
|
+
)
|
|
113
|
+
|
|
114
|
+
|
|
115
|
+
def resolve_extracted_data_directory(extraction_root: Path) -> Path:
|
|
116
|
+
visible_contents = [
|
|
117
|
+
path
|
|
118
|
+
for path in extraction_root.iterdir()
|
|
119
|
+
if not path.name.startswith(".") and path.name != "__MACOSX"
|
|
120
|
+
]
|
|
121
|
+
|
|
122
|
+
named_directory = extraction_root / "espeak-ng-data"
|
|
123
|
+
if named_directory.is_dir():
|
|
124
|
+
return named_directory
|
|
125
|
+
|
|
126
|
+
child_directories = [path for path in visible_contents if path.is_dir()]
|
|
127
|
+
if len(visible_contents) == 1 and len(child_directories) == 1:
|
|
128
|
+
return child_directories[0]
|
|
129
|
+
|
|
130
|
+
return extraction_root
|
|
@@ -0,0 +1,46 @@
|
|
|
1
|
+
# Generated from wfloat/assets/registry.json. Do not edit.
|
|
2
|
+
|
|
3
|
+
REGISTRY_ORIGIN = 'https://registry.wfloat.com'
|
|
4
|
+
|
|
5
|
+
MODEL_ASSETS = {'HuggingFaceTB/SmolLM2-360M-Instruct': {'family': 'smollm',
|
|
6
|
+
'model': {'path': '/models/huggingfacetb/smollm2-360m-instruct/model.Q4_K_M.gguf',
|
|
7
|
+
'sha256': '75c4346ef9e855ed630f80078a2430cf63aaca599e340360998a313070fcdc47'}},
|
|
8
|
+
'UsefulSensors/moonshine-tiny': {'cached_decoder': {'path': '/models/usefulsensors/moonshine-tiny/cached_decoder.int8.onnx',
|
|
9
|
+
'sha256': '2aff28bba6a03d8dcf5c9feac45462629bae37317442299f28115ad09da773f6'},
|
|
10
|
+
'encoder': {'path': '/models/usefulsensors/moonshine-tiny/encoder.int8.onnx',
|
|
11
|
+
'sha256': '8774dfba578de027ec6595c2c654a0836434489bc963a0db124a7f181f571acb'},
|
|
12
|
+
'family': 'moonshine',
|
|
13
|
+
'preprocessor': {'path': '/models/usefulsensors/moonshine-tiny/preprocessor.onnx',
|
|
14
|
+
'sha256': 'f33addce61a143460fe753b5ee5b7db255e5140b5b779c065b94f6c83ff0bf4e'},
|
|
15
|
+
'tokens': {'path': '/models/usefulsensors/moonshine-tiny/tokens.txt',
|
|
16
|
+
'sha256': '1165c2aeb9f72f457a83be2d459a09054f27490acd9b41bd43794dfd25e296ea'},
|
|
17
|
+
'uncached_decoder': {'path': '/models/usefulsensors/moonshine-tiny/uncached_decoder.int8.onnx',
|
|
18
|
+
'sha256': '216737000dd5881a17aa043f6bbd286add33e4c3b0ae257153e2ec15438bdc41'}},
|
|
19
|
+
'k2-fsa/streaming-zipformer-en': {'decoder': {'path': '/models/k2-fsa/streaming-zipformer-en/decoder.onnx',
|
|
20
|
+
'sha256': '9da02b77cb08826756ec6a88635f35a40374e4164e7c6359121a9145958a6ceb'},
|
|
21
|
+
'encoder': {'path': '/models/k2-fsa/streaming-zipformer-en/encoder.int8.onnx',
|
|
22
|
+
'sha256': '32c98281c7bd8b63e3e142d007251b37f120572e8fdea9a4f5a79ce22b10ec4f'},
|
|
23
|
+
'family': 'zipformer-transducer',
|
|
24
|
+
'joiner': {'path': '/models/k2-fsa/streaming-zipformer-en/joiner.onnx',
|
|
25
|
+
'sha256': 'bd5c26ad6a41cbd90c2cfa239c0b55b145af878ce1d79b4739d90f8be93359ba'},
|
|
26
|
+
'tokens': {'path': '/models/k2-fsa/streaming-zipformer-en/tokens.txt',
|
|
27
|
+
'sha256': '49e3c2646595fd907228b3c6787069658f67b17377c60aeb8619c4551b2316fb'}},
|
|
28
|
+
'openai/whisper-tiny-en': {'decoder': {'path': '/models/openai/whisper-tiny-en/decoder.int8.onnx',
|
|
29
|
+
'sha256': '06c0e6ff6348d427e51839219d1c886c18cfdf411e629e33f5e1679bff9c1527'},
|
|
30
|
+
'encoder': {'path': '/models/openai/whisper-tiny-en/encoder.int8.onnx',
|
|
31
|
+
'sha256': '0ce578b827c94a961aacb8fa14b02f096504b337e5c94be37c36238cbe3e8bc6'},
|
|
32
|
+
'family': 'whisper',
|
|
33
|
+
'tokens': {'path': '/models/openai/whisper-tiny-en/tokens.txt',
|
|
34
|
+
'sha256': '306cd27f03c1a714eca7108e03d66b7dc042abe8c258b44c199a7ed9838dd930'}},
|
|
35
|
+
'snakers4/silero-vad': {'family': 'silero-vad',
|
|
36
|
+
'model': {'path': '/models/snakers4/silero-vad/model.onnx',
|
|
37
|
+
'sha256': '9e2449e1087496d8d4caba907f23e0bd3f78d91fa552479bb9c23ac09cbb1fd6'}},
|
|
38
|
+
'wfloat/wfloat-tts': {'model_onnx': {'path': '/models/wfloat/wfloat-tts/model.onnx',
|
|
39
|
+
'sha256': 'a7e65773a29499b80a393bbe08af3507e18f6ef95faa0eaf7cb4ba353c8693ae'},
|
|
40
|
+
'model_tokens': {'path': '/models/wfloat/wfloat-tts/tokens.txt',
|
|
41
|
+
'sha256': '96fd291bede0544469d4d8935d462fdd6dc947f22ad47369753e1a82db3d748e'}}}
|
|
42
|
+
|
|
43
|
+
SHARED_ASSETS = {'espeak_ng_data_aar': {'path': '/assets/espeak-ng-data/espeak-ng-data-2023.9.7-4.aar',
|
|
44
|
+
'sha256': 'a526b72e81cb1a17e07f55ca0117bba8fbcac7ccd2fa502c61be926eafeaf64e'},
|
|
45
|
+
'espeak_ng_data_zip': {'path': '/assets/espeak-ng-data/espeak-ng-data-2023.9.7-4.zip',
|
|
46
|
+
'sha256': '56c2879ab1ab44c594c78f34e76c50cf1dd7b8f6ca0ca2634b6766a6edb32add'}}
|
wfloat/_llm.py
ADDED
|
@@ -0,0 +1,84 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from dataclasses import dataclass
|
|
4
|
+
from typing import Callable, Dict, Optional, Sequence
|
|
5
|
+
|
|
6
|
+
from ._results import LlmGenerationResult
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
@dataclass
|
|
10
|
+
class LlmModel:
|
|
11
|
+
model_id: str
|
|
12
|
+
family: str
|
|
13
|
+
_native_llm: object
|
|
14
|
+
context_size: int = 0
|
|
15
|
+
|
|
16
|
+
def generate(
|
|
17
|
+
self,
|
|
18
|
+
prompt: str,
|
|
19
|
+
*,
|
|
20
|
+
max_tokens: int = 128,
|
|
21
|
+
temperature: float = 0.8,
|
|
22
|
+
top_p: float = 0.95,
|
|
23
|
+
top_k: int = 40,
|
|
24
|
+
repeat_penalty: float = 1.0,
|
|
25
|
+
seed: int = 0,
|
|
26
|
+
on_token: Optional[Callable[[str], None]] = None,
|
|
27
|
+
) -> LlmGenerationResult:
|
|
28
|
+
if not hasattr(self._native_llm, "generate"):
|
|
29
|
+
raise RuntimeError("Native LLM backend does not support generate().")
|
|
30
|
+
|
|
31
|
+
return self._native_llm.generate(
|
|
32
|
+
prompt,
|
|
33
|
+
max_tokens=max_tokens,
|
|
34
|
+
temperature=temperature,
|
|
35
|
+
top_p=top_p,
|
|
36
|
+
top_k=top_k,
|
|
37
|
+
repeat_penalty=repeat_penalty,
|
|
38
|
+
seed=seed,
|
|
39
|
+
on_token=on_token,
|
|
40
|
+
)
|
|
41
|
+
|
|
42
|
+
def chat(
|
|
43
|
+
self,
|
|
44
|
+
messages: Sequence[Dict[str, str]],
|
|
45
|
+
*,
|
|
46
|
+
max_tokens: int = 128,
|
|
47
|
+
temperature: float = 0.8,
|
|
48
|
+
top_p: float = 0.95,
|
|
49
|
+
top_k: int = 40,
|
|
50
|
+
repeat_penalty: float = 1.0,
|
|
51
|
+
seed: int = 0,
|
|
52
|
+
on_token: Optional[Callable[[str], None]] = None,
|
|
53
|
+
) -> LlmGenerationResult:
|
|
54
|
+
if not hasattr(self._native_llm, "chat"):
|
|
55
|
+
raise RuntimeError("Native LLM backend does not support chat().")
|
|
56
|
+
|
|
57
|
+
return self._native_llm.chat(
|
|
58
|
+
messages,
|
|
59
|
+
max_tokens=max_tokens,
|
|
60
|
+
temperature=temperature,
|
|
61
|
+
top_p=top_p,
|
|
62
|
+
top_k=top_k,
|
|
63
|
+
repeat_penalty=repeat_penalty,
|
|
64
|
+
seed=seed,
|
|
65
|
+
on_token=on_token,
|
|
66
|
+
)
|
|
67
|
+
|
|
68
|
+
def format_chat(
|
|
69
|
+
self,
|
|
70
|
+
messages: Sequence[Dict[str, str]],
|
|
71
|
+
*,
|
|
72
|
+
add_generation_prompt: bool = True,
|
|
73
|
+
) -> str:
|
|
74
|
+
if not hasattr(self._native_llm, "format_chat"):
|
|
75
|
+
raise RuntimeError("Native LLM backend does not support format_chat().")
|
|
76
|
+
|
|
77
|
+
return self._native_llm.format_chat(
|
|
78
|
+
messages,
|
|
79
|
+
add_generation_prompt=add_generation_prompt,
|
|
80
|
+
)
|
|
81
|
+
|
|
82
|
+
def close(self) -> None:
|
|
83
|
+
if hasattr(self._native_llm, "close"):
|
|
84
|
+
self._native_llm.close()
|
wfloat/_llm_assets.py
ADDED
|
@@ -0,0 +1,142 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import shutil
|
|
4
|
+
from dataclasses import dataclass
|
|
5
|
+
from pathlib import Path
|
|
6
|
+
from typing import Mapping, Optional
|
|
7
|
+
from urllib.parse import urlparse
|
|
8
|
+
|
|
9
|
+
from ._assets import LlmModelAssets
|
|
10
|
+
from ._download import download_file, verify_checksum
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
@dataclass(frozen=True)
|
|
14
|
+
class CachedLlmAssets:
|
|
15
|
+
model_name: str
|
|
16
|
+
family: str
|
|
17
|
+
cache_dir: Path
|
|
18
|
+
files: Mapping[str, Path]
|
|
19
|
+
context_size: Optional[int] = None
|
|
20
|
+
chat_template: Optional[str] = None
|
|
21
|
+
chat_template_format: Optional[str] = None
|
|
22
|
+
|
|
23
|
+
def require(self, key: str) -> Path:
|
|
24
|
+
path = self.files.get(key)
|
|
25
|
+
if path is None:
|
|
26
|
+
raise ValueError(f"Missing required LLM asset: {key}")
|
|
27
|
+
return path
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
def _normalize_model_dir_name(model_name: str) -> str:
|
|
31
|
+
return model_name.replace("/", "--").replace(" ", "-")
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
def _is_url(value: str) -> bool:
|
|
35
|
+
parsed = urlparse(value)
|
|
36
|
+
return parsed.scheme in {"http", "https", "file"}
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
def _copy_local_file(source: Path, destination: Path) -> None:
|
|
40
|
+
destination.parent.mkdir(parents=True, exist_ok=True)
|
|
41
|
+
shutil.copyfile(str(source), str(destination))
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
def _materialize_asset(
|
|
45
|
+
source: str | Path,
|
|
46
|
+
destination: Path,
|
|
47
|
+
*,
|
|
48
|
+
expected_checksum: Optional[str],
|
|
49
|
+
force_download: bool,
|
|
50
|
+
) -> Path:
|
|
51
|
+
destination.parent.mkdir(parents=True, exist_ok=True)
|
|
52
|
+
|
|
53
|
+
if not force_download and destination.is_file():
|
|
54
|
+
if expected_checksum is None or verify_checksum(destination, expected_checksum):
|
|
55
|
+
return destination
|
|
56
|
+
|
|
57
|
+
if isinstance(source, Path):
|
|
58
|
+
_copy_local_file(source, destination)
|
|
59
|
+
else:
|
|
60
|
+
source_str = str(source)
|
|
61
|
+
if _is_url(source_str):
|
|
62
|
+
download_file(source_str, destination, expected_checksum=expected_checksum)
|
|
63
|
+
else:
|
|
64
|
+
_copy_local_file(Path(source_str), destination)
|
|
65
|
+
|
|
66
|
+
if expected_checksum is not None and not verify_checksum(destination, expected_checksum):
|
|
67
|
+
raise RuntimeError(f"Cached LLM asset checksum mismatch for {destination}.")
|
|
68
|
+
|
|
69
|
+
return destination
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
def cache_llm_assets(
|
|
73
|
+
model_name: str,
|
|
74
|
+
*,
|
|
75
|
+
family: str,
|
|
76
|
+
sources: Mapping[str, str | Path | None],
|
|
77
|
+
checksums: Optional[Mapping[str, str]] = None,
|
|
78
|
+
cache_dir: Path,
|
|
79
|
+
force_download: bool = False,
|
|
80
|
+
context_size: Optional[int] = None,
|
|
81
|
+
chat_template: Optional[str] = None,
|
|
82
|
+
chat_template_format: Optional[str] = None,
|
|
83
|
+
) -> CachedLlmAssets:
|
|
84
|
+
model_dir = cache_dir / "models" / _normalize_model_dir_name(model_name)
|
|
85
|
+
model_dir.mkdir(parents=True, exist_ok=True)
|
|
86
|
+
checksums = checksums or {}
|
|
87
|
+
|
|
88
|
+
files: dict[str, Path] = {}
|
|
89
|
+
for key, source in sources.items():
|
|
90
|
+
if source is None:
|
|
91
|
+
continue
|
|
92
|
+
|
|
93
|
+
source_value = Path(source) if isinstance(source, Path) else str(source)
|
|
94
|
+
filename = Path(urlparse(str(source_value)).path).name or Path(str(source_value)).name
|
|
95
|
+
if not filename:
|
|
96
|
+
raise ValueError(f"Could not derive filename for LLM asset {key}.")
|
|
97
|
+
|
|
98
|
+
destination = model_dir / filename
|
|
99
|
+
files[key] = _materialize_asset(
|
|
100
|
+
source_value,
|
|
101
|
+
destination,
|
|
102
|
+
expected_checksum=checksums.get(key),
|
|
103
|
+
force_download=force_download,
|
|
104
|
+
)
|
|
105
|
+
|
|
106
|
+
return CachedLlmAssets(
|
|
107
|
+
model_name=model_name,
|
|
108
|
+
family=family,
|
|
109
|
+
cache_dir=cache_dir,
|
|
110
|
+
files=files,
|
|
111
|
+
context_size=context_size,
|
|
112
|
+
chat_template=chat_template,
|
|
113
|
+
chat_template_format=chat_template_format,
|
|
114
|
+
)
|
|
115
|
+
|
|
116
|
+
|
|
117
|
+
def cache_llm_model_assets(
|
|
118
|
+
model_name: str,
|
|
119
|
+
assets: LlmModelAssets,
|
|
120
|
+
*,
|
|
121
|
+
cache_dir: Path,
|
|
122
|
+
force_download: bool = False,
|
|
123
|
+
) -> CachedLlmAssets:
|
|
124
|
+
return cache_llm_assets(
|
|
125
|
+
model_name,
|
|
126
|
+
family=assets.family,
|
|
127
|
+
sources={
|
|
128
|
+
"model": assets.model,
|
|
129
|
+
},
|
|
130
|
+
checksums={
|
|
131
|
+
key: value
|
|
132
|
+
for key, value in {
|
|
133
|
+
"model": assets.model_checksum,
|
|
134
|
+
}.items()
|
|
135
|
+
if value is not None
|
|
136
|
+
},
|
|
137
|
+
cache_dir=cache_dir,
|
|
138
|
+
force_download=force_download,
|
|
139
|
+
context_size=assets.context_size,
|
|
140
|
+
chat_template=assets.chat_template,
|
|
141
|
+
chat_template_format=assets.chat_template_format,
|
|
142
|
+
)
|
wfloat/_llm_load.py
ADDED
|
@@ -0,0 +1,61 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from pathlib import Path
|
|
4
|
+
from typing import Optional
|
|
5
|
+
|
|
6
|
+
from ._assets import fetch_llm_assets
|
|
7
|
+
from ._cache import get_default_cache_dir
|
|
8
|
+
from ._core import create_core_llm
|
|
9
|
+
from ._llm import LlmModel
|
|
10
|
+
from ._llm_assets import cache_llm_model_assets
|
|
11
|
+
|
|
12
|
+
DEFAULT_LLM_CONTEXT_SIZE = 2048
|
|
13
|
+
DEFAULT_LLM_NUM_THREADS = 4
|
|
14
|
+
DEFAULT_LLM_GPU_LAYER_COUNT = 0
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def load_llm_model(
|
|
18
|
+
model_name: str,
|
|
19
|
+
*,
|
|
20
|
+
cache_dir: Optional[Path] = None,
|
|
21
|
+
force_download: bool = False,
|
|
22
|
+
context_size: Optional[int] = None,
|
|
23
|
+
num_threads: int = DEFAULT_LLM_NUM_THREADS,
|
|
24
|
+
gpu_layer_count: int = DEFAULT_LLM_GPU_LAYER_COUNT,
|
|
25
|
+
chat_template: Optional[str] = None,
|
|
26
|
+
) -> LlmModel:
|
|
27
|
+
resolved_cache_dir = Path(cache_dir) if cache_dir is not None else get_default_cache_dir()
|
|
28
|
+
assets = fetch_llm_assets(model_name)
|
|
29
|
+
cached = cache_llm_model_assets(
|
|
30
|
+
model_name,
|
|
31
|
+
assets,
|
|
32
|
+
cache_dir=resolved_cache_dir,
|
|
33
|
+
force_download=force_download,
|
|
34
|
+
)
|
|
35
|
+
family = assets.family
|
|
36
|
+
|
|
37
|
+
resolved_context_size = (
|
|
38
|
+
int(context_size)
|
|
39
|
+
if context_size is not None
|
|
40
|
+
else int(cached.context_size or DEFAULT_LLM_CONTEXT_SIZE)
|
|
41
|
+
)
|
|
42
|
+
resolved_chat_template = chat_template if chat_template is not None else cached.chat_template
|
|
43
|
+
if resolved_chat_template is None and getattr(cached, "chat_template_format", None) == "chatml":
|
|
44
|
+
resolved_chat_template = "chatml"
|
|
45
|
+
|
|
46
|
+
native_llm = create_core_llm(
|
|
47
|
+
model_name=model_name,
|
|
48
|
+
family=family,
|
|
49
|
+
model_path=cached.require("model"),
|
|
50
|
+
context_size=resolved_context_size,
|
|
51
|
+
num_threads=num_threads,
|
|
52
|
+
gpu_layer_count=gpu_layer_count,
|
|
53
|
+
chat_template=resolved_chat_template,
|
|
54
|
+
)
|
|
55
|
+
|
|
56
|
+
return LlmModel(
|
|
57
|
+
model_id=model_name,
|
|
58
|
+
family=family,
|
|
59
|
+
_native_llm=native_llm,
|
|
60
|
+
context_size=resolved_context_size,
|
|
61
|
+
)
|