whisperspeech2 0.9.0__tar.gz
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.
- whisperspeech2-0.9.0/LICENSE +22 -0
- whisperspeech2-0.9.0/MANIFEST.in +5 -0
- whisperspeech2-0.9.0/PKG-INFO +122 -0
- whisperspeech2-0.9.0/README.md +77 -0
- whisperspeech2-0.9.0/settings.ini +26 -0
- whisperspeech2-0.9.0/setup.cfg +4 -0
- whisperspeech2-0.9.0/setup.py +49 -0
- whisperspeech2-0.9.0/whisperspeech2/__init__.py +1 -0
- whisperspeech2-0.9.0/whisperspeech2/a2wav.py +72 -0
- whisperspeech2-0.9.0/whisperspeech2/inference.py +56 -0
- whisperspeech2-0.9.0/whisperspeech2/languages.py +131 -0
- whisperspeech2-0.9.0/whisperspeech2/modules.py +303 -0
- whisperspeech2-0.9.0/whisperspeech2/pipeline.py +107 -0
- whisperspeech2-0.9.0/whisperspeech2/quick_test.py +8 -0
- whisperspeech2-0.9.0/whisperspeech2/s2a_delar_mup_wds_mlang.py +378 -0
- whisperspeech2-0.9.0/whisperspeech2/s2a_delar_mup_wds_mlang_cond.py +434 -0
- whisperspeech2-0.9.0/whisperspeech2/t2s_up_wds_mlang_enclm.py +336 -0
- whisperspeech2-0.9.0/whisperspeech2.egg-info/PKG-INFO +122 -0
- whisperspeech2-0.9.0/whisperspeech2.egg-info/SOURCES.txt +21 -0
- whisperspeech2-0.9.0/whisperspeech2.egg-info/dependency_links.txt +1 -0
- whisperspeech2-0.9.0/whisperspeech2.egg-info/not-zip-safe +1 -0
- whisperspeech2-0.9.0/whisperspeech2.egg-info/requires.txt +11 -0
- whisperspeech2-0.9.0/whisperspeech2.egg-info/top_level.txt +1 -0
|
@@ -0,0 +1,22 @@
|
|
|
1
|
+
MIT License
|
|
2
|
+
|
|
3
|
+
Copyright (c) 2023 Jakub Piotr Cłapa, Collabora Ltd.
|
|
4
|
+
Copyright (c) 2025 Blair Chintella
|
|
5
|
+
|
|
6
|
+
Permission is hereby granted, free of charge, to any person obtaining a copy
|
|
7
|
+
of this software and associated documentation files (the "Software"), to deal
|
|
8
|
+
in the Software without restriction, including without limitation the rights
|
|
9
|
+
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
|
10
|
+
copies of the Software, and to permit persons to whom the Software is
|
|
11
|
+
furnished to do so, subject to the following conditions:
|
|
12
|
+
|
|
13
|
+
The above copyright notice and this permission notice shall be included in all
|
|
14
|
+
copies or substantial portions of the Software.
|
|
15
|
+
|
|
16
|
+
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
|
17
|
+
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
|
18
|
+
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
|
19
|
+
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
|
20
|
+
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
|
21
|
+
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
|
22
|
+
SOFTWARE.
|
|
@@ -0,0 +1,122 @@
|
|
|
1
|
+
Metadata-Version: 2.4
|
|
2
|
+
Name: whisperspeech2
|
|
3
|
+
Version: 0.9.0
|
|
4
|
+
Summary: An Open Source text-to-speech system built by inverting Whisper (fork of WhisperSpeech)
|
|
5
|
+
Home-page: https://github.com/BBC-Esq/whisperspeech2
|
|
6
|
+
Author: Blair Chintella
|
|
7
|
+
Author-email: vici0549@gmail.com
|
|
8
|
+
License: MIT License
|
|
9
|
+
Keywords: tts text-to-speech whisper speech-synthesis
|
|
10
|
+
Classifier: Development Status :: 4 - Beta
|
|
11
|
+
Classifier: Intended Audience :: Developers
|
|
12
|
+
Classifier: Natural Language :: English
|
|
13
|
+
Classifier: Programming Language :: Python :: 3.8
|
|
14
|
+
Classifier: Programming Language :: Python :: 3.9
|
|
15
|
+
Classifier: Programming Language :: Python :: 3.10
|
|
16
|
+
Classifier: Programming Language :: Python :: 3.11
|
|
17
|
+
Classifier: Programming Language :: Python :: 3.12
|
|
18
|
+
Classifier: License :: OSI Approved :: MIT License
|
|
19
|
+
Requires-Python: >=3.8
|
|
20
|
+
Description-Content-Type: text/markdown
|
|
21
|
+
License-File: LICENSE
|
|
22
|
+
Requires-Dist: vocos
|
|
23
|
+
Requires-Dist: torch>=2
|
|
24
|
+
Requires-Dist: torchaudio
|
|
25
|
+
Requires-Dist: soundfile
|
|
26
|
+
Requires-Dist: huggingface_hub
|
|
27
|
+
Requires-Dist: fastprogress
|
|
28
|
+
Requires-Dist: fastcore
|
|
29
|
+
Requires-Dist: numpy
|
|
30
|
+
Provides-Extra: speaker
|
|
31
|
+
Requires-Dist: speechbrain<1.0; extra == "speaker"
|
|
32
|
+
Dynamic: author
|
|
33
|
+
Dynamic: author-email
|
|
34
|
+
Dynamic: classifier
|
|
35
|
+
Dynamic: description
|
|
36
|
+
Dynamic: description-content-type
|
|
37
|
+
Dynamic: home-page
|
|
38
|
+
Dynamic: keywords
|
|
39
|
+
Dynamic: license
|
|
40
|
+
Dynamic: license-file
|
|
41
|
+
Dynamic: provides-extra
|
|
42
|
+
Dynamic: requires-dist
|
|
43
|
+
Dynamic: requires-python
|
|
44
|
+
Dynamic: summary
|
|
45
|
+
|
|
46
|
+
# WhisperSpeech2
|
|
47
|
+
|
|
48
|
+
An Open Source text-to-speech system built by inverting Whisper. This is a fork of [WhisperSpeech](https://github.com/collabora/WhisperSpeech) optimized for inference.
|
|
49
|
+
|
|
50
|
+
## Installation
|
|
51
|
+
```bash
|
|
52
|
+
pip install whisperspeech2
|
|
53
|
+
```
|
|
54
|
+
|
|
55
|
+
**Note:** You must also have PyTorch installed. Visit [pytorch.org](https://pytorch.org/get-started/locally/) for installation instructions.
|
|
56
|
+
|
|
57
|
+
## Quick Start
|
|
58
|
+
```python
|
|
59
|
+
from whisperspeech2.pipeline import Pipeline
|
|
60
|
+
|
|
61
|
+
# Initialize the pipeline
|
|
62
|
+
pipe = Pipeline(s2a_ref='collabora/whisperspeech:s2a-q4-tiny-en+pl.model')
|
|
63
|
+
|
|
64
|
+
# Generate audio and save to file
|
|
65
|
+
pipe.generate_to_file('output.wav', "Hello, world!")
|
|
66
|
+
|
|
67
|
+
# Or get the audio tensor directly
|
|
68
|
+
audio = pipe.generate("Hello, world!")
|
|
69
|
+
```
|
|
70
|
+
|
|
71
|
+
## Available Models
|
|
72
|
+
|
|
73
|
+
| Model | Reference |
|
|
74
|
+
|-------|-----------|
|
|
75
|
+
| Tiny | `collabora/whisperspeech:s2a-q4-tiny-en+pl.model` |
|
|
76
|
+
| Base | `collabora/whisperspeech:s2a-q4-base-en+pl.model` |
|
|
77
|
+
| Small | `collabora/whisperspeech:s2a-q4-small-en+pl.model` |
|
|
78
|
+
|
|
79
|
+
## Speaker Embedding (Optional)
|
|
80
|
+
|
|
81
|
+
To use custom speaker embeddings, install the optional dependency:
|
|
82
|
+
```bash
|
|
83
|
+
pip install whisperspeech2[speaker]
|
|
84
|
+
```
|
|
85
|
+
|
|
86
|
+
Then pass an audio file path to clone a voice:
|
|
87
|
+
```python
|
|
88
|
+
pipe.generate_to_file('output.wav', "Hello!", speaker='reference.wav')
|
|
89
|
+
```
|
|
90
|
+
|
|
91
|
+
## Examples
|
|
92
|
+
|
|
93
|
+
See the `examples/` directory for more usage examples including GUI applications and streaming playback.
|
|
94
|
+
|
|
95
|
+
## License
|
|
96
|
+
|
|
97
|
+
MIT License
|
|
98
|
+
```
|
|
99
|
+
|
|
100
|
+
### 3. `LICENSE`
|
|
101
|
+
```
|
|
102
|
+
MIT License
|
|
103
|
+
|
|
104
|
+
Copyright (c) 2025 Blair Chintella
|
|
105
|
+
|
|
106
|
+
Permission is hereby granted, free of charge, to any person obtaining a copy
|
|
107
|
+
of this software and associated documentation files (the "Software"), to deal
|
|
108
|
+
in the Software without restriction, including without limitation the rights
|
|
109
|
+
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
|
110
|
+
copies of the Software, and to permit persons to whom the Software is
|
|
111
|
+
furnished to do so, subject to the following conditions:
|
|
112
|
+
|
|
113
|
+
The above copyright notice and this permission notice shall be included in all
|
|
114
|
+
copies or substantial portions of the Software.
|
|
115
|
+
|
|
116
|
+
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
|
117
|
+
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
|
118
|
+
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
|
119
|
+
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
|
120
|
+
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
|
121
|
+
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
|
122
|
+
SOFTWARE.
|
|
@@ -0,0 +1,77 @@
|
|
|
1
|
+
# WhisperSpeech2
|
|
2
|
+
|
|
3
|
+
An Open Source text-to-speech system built by inverting Whisper. This is a fork of [WhisperSpeech](https://github.com/collabora/WhisperSpeech) optimized for inference.
|
|
4
|
+
|
|
5
|
+
## Installation
|
|
6
|
+
```bash
|
|
7
|
+
pip install whisperspeech2
|
|
8
|
+
```
|
|
9
|
+
|
|
10
|
+
**Note:** You must also have PyTorch installed. Visit [pytorch.org](https://pytorch.org/get-started/locally/) for installation instructions.
|
|
11
|
+
|
|
12
|
+
## Quick Start
|
|
13
|
+
```python
|
|
14
|
+
from whisperspeech2.pipeline import Pipeline
|
|
15
|
+
|
|
16
|
+
# Initialize the pipeline
|
|
17
|
+
pipe = Pipeline(s2a_ref='collabora/whisperspeech:s2a-q4-tiny-en+pl.model')
|
|
18
|
+
|
|
19
|
+
# Generate audio and save to file
|
|
20
|
+
pipe.generate_to_file('output.wav', "Hello, world!")
|
|
21
|
+
|
|
22
|
+
# Or get the audio tensor directly
|
|
23
|
+
audio = pipe.generate("Hello, world!")
|
|
24
|
+
```
|
|
25
|
+
|
|
26
|
+
## Available Models
|
|
27
|
+
|
|
28
|
+
| Model | Reference |
|
|
29
|
+
|-------|-----------|
|
|
30
|
+
| Tiny | `collabora/whisperspeech:s2a-q4-tiny-en+pl.model` |
|
|
31
|
+
| Base | `collabora/whisperspeech:s2a-q4-base-en+pl.model` |
|
|
32
|
+
| Small | `collabora/whisperspeech:s2a-q4-small-en+pl.model` |
|
|
33
|
+
|
|
34
|
+
## Speaker Embedding (Optional)
|
|
35
|
+
|
|
36
|
+
To use custom speaker embeddings, install the optional dependency:
|
|
37
|
+
```bash
|
|
38
|
+
pip install whisperspeech2[speaker]
|
|
39
|
+
```
|
|
40
|
+
|
|
41
|
+
Then pass an audio file path to clone a voice:
|
|
42
|
+
```python
|
|
43
|
+
pipe.generate_to_file('output.wav', "Hello!", speaker='reference.wav')
|
|
44
|
+
```
|
|
45
|
+
|
|
46
|
+
## Examples
|
|
47
|
+
|
|
48
|
+
See the `examples/` directory for more usage examples including GUI applications and streaming playback.
|
|
49
|
+
|
|
50
|
+
## License
|
|
51
|
+
|
|
52
|
+
MIT License
|
|
53
|
+
```
|
|
54
|
+
|
|
55
|
+
### 3. `LICENSE`
|
|
56
|
+
```
|
|
57
|
+
MIT License
|
|
58
|
+
|
|
59
|
+
Copyright (c) 2025 Blair Chintella
|
|
60
|
+
|
|
61
|
+
Permission is hereby granted, free of charge, to any person obtaining a copy
|
|
62
|
+
of this software and associated documentation files (the "Software"), to deal
|
|
63
|
+
in the Software without restriction, including without limitation the rights
|
|
64
|
+
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
|
65
|
+
copies of the Software, and to permit persons to whom the Software is
|
|
66
|
+
furnished to do so, subject to the following conditions:
|
|
67
|
+
|
|
68
|
+
The above copyright notice and this permission notice shall be included in all
|
|
69
|
+
copies or substantial portions of the Software.
|
|
70
|
+
|
|
71
|
+
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
|
72
|
+
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
|
73
|
+
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
|
74
|
+
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
|
75
|
+
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
|
76
|
+
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
|
77
|
+
SOFTWARE.
|
|
@@ -0,0 +1,26 @@
|
|
|
1
|
+
[DEFAULT]
|
|
2
|
+
repo = whisperspeech2
|
|
3
|
+
lib_name = %(repo)s
|
|
4
|
+
version = 0.9.0
|
|
5
|
+
min_python = 3.8
|
|
6
|
+
license = MIT
|
|
7
|
+
|
|
8
|
+
lib_path = whisperspeech2
|
|
9
|
+
nbs_path = nbs
|
|
10
|
+
recursive = False
|
|
11
|
+
|
|
12
|
+
branch = master
|
|
13
|
+
git_url = https://github.com/%(user)s/%(repo)s
|
|
14
|
+
title = %(lib_name)s
|
|
15
|
+
|
|
16
|
+
audience = Developers
|
|
17
|
+
author = Blair Chintella
|
|
18
|
+
author_email = vici0549@gmail.com
|
|
19
|
+
copyright = 2025 onwards, %(author)s
|
|
20
|
+
description = An Open Source text-to-speech system built by inverting Whisper (fork of WhisperSpeech)
|
|
21
|
+
keywords = tts text-to-speech whisper speech-synthesis
|
|
22
|
+
language = English
|
|
23
|
+
status = 3
|
|
24
|
+
user = BBC-Esq
|
|
25
|
+
|
|
26
|
+
requirements = vocos torch>=2 torchaudio soundfile huggingface_hub fastprogress fastcore numpy
|
|
@@ -0,0 +1,49 @@
|
|
|
1
|
+
from pkg_resources import parse_version
|
|
2
|
+
from configparser import ConfigParser
|
|
3
|
+
import setuptools, shlex
|
|
4
|
+
assert parse_version(setuptools.__version__)>=parse_version('36.2')
|
|
5
|
+
|
|
6
|
+
config = ConfigParser(delimiters=['='])
|
|
7
|
+
config.read('settings.ini')
|
|
8
|
+
cfg = config['DEFAULT']
|
|
9
|
+
|
|
10
|
+
cfg_keys = 'version description keywords author author_email'.split()
|
|
11
|
+
expected = cfg_keys + "lib_name user branch license status min_python audience language".split()
|
|
12
|
+
for o in expected: assert o in cfg, "missing expected setting: {}".format(o)
|
|
13
|
+
setup_cfg = {o:cfg[o] for o in cfg_keys}
|
|
14
|
+
|
|
15
|
+
licenses = {
|
|
16
|
+
'apache2': ('Apache Software License 2.0','OSI Approved :: Apache Software License'),
|
|
17
|
+
'mit': ('MIT License', 'OSI Approved :: MIT License'),
|
|
18
|
+
'gpl2': ('GNU General Public License v2', 'OSI Approved :: GNU General Public License v2 (GPLv2)'),
|
|
19
|
+
'gpl3': ('GNU General Public License v3', 'OSI Approved :: GNU General Public License v3 (GPLv3)'),
|
|
20
|
+
'bsd3': ('BSD License', 'OSI Approved :: BSD License'),
|
|
21
|
+
}
|
|
22
|
+
statuses = [ '1 - Planning', '2 - Pre-Alpha', '3 - Alpha',
|
|
23
|
+
'4 - Beta', '5 - Production/Stable', '6 - Mature', '7 - Inactive' ]
|
|
24
|
+
py_versions = '3.8 3.9 3.10 3.11 3.12'.split()
|
|
25
|
+
|
|
26
|
+
requirements = shlex.split(cfg.get('requirements', ''))
|
|
27
|
+
min_python = cfg['min_python']
|
|
28
|
+
lic = licenses.get(cfg['license'].lower(), (cfg['license'], None))
|
|
29
|
+
|
|
30
|
+
setuptools.setup(
|
|
31
|
+
name = cfg['lib_name'],
|
|
32
|
+
license = lic[0],
|
|
33
|
+
classifiers = [
|
|
34
|
+
'Development Status :: ' + statuses[int(cfg['status'])],
|
|
35
|
+
'Intended Audience :: ' + cfg['audience'].title(),
|
|
36
|
+
'Natural Language :: ' + cfg['language'].title(),
|
|
37
|
+
] + ['Programming Language :: Python :: '+o for o in py_versions[py_versions.index(min_python):]] + (['License :: ' + lic[1] ] if lic[1] else []),
|
|
38
|
+
url = cfg['git_url'],
|
|
39
|
+
packages = setuptools.find_packages(),
|
|
40
|
+
include_package_data = True,
|
|
41
|
+
install_requires = requirements,
|
|
42
|
+
extras_require={
|
|
43
|
+
'speaker': ['speechbrain<1.0'],
|
|
44
|
+
},
|
|
45
|
+
python_requires = '>=' + cfg['min_python'],
|
|
46
|
+
long_description = open('README.md').read(),
|
|
47
|
+
long_description_content_type = 'text/markdown',
|
|
48
|
+
zip_safe = False,
|
|
49
|
+
**setup_cfg)
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
__version__ = "0.9.0"
|
|
@@ -0,0 +1,72 @@
|
|
|
1
|
+
__all__ = ['Vocoder']
|
|
2
|
+
|
|
3
|
+
from vocos import Vocos
|
|
4
|
+
from whisperspeech2 import inference
|
|
5
|
+
import torch
|
|
6
|
+
|
|
7
|
+
class Vocoder:
|
|
8
|
+
def __init__(self, repo_id="charactr/vocos-encodec-24khz", device=None, cache_dir=None):
|
|
9
|
+
if device is None: device = inference.get_compute_device()
|
|
10
|
+
if device == 'mps': device = 'cpu'
|
|
11
|
+
self.device = device
|
|
12
|
+
self.vocos = Vocos.from_pretrained(repo_id).to(device)
|
|
13
|
+
|
|
14
|
+
def is_notebook(self):
|
|
15
|
+
try:
|
|
16
|
+
return get_ipython().__class__.__name__ == "ZMQInteractiveShell"
|
|
17
|
+
except:
|
|
18
|
+
return False
|
|
19
|
+
|
|
20
|
+
@torch.no_grad()
|
|
21
|
+
def decode(self, atoks):
|
|
22
|
+
if len(atoks.shape) == 3:
|
|
23
|
+
b,q,t = atoks.shape
|
|
24
|
+
atoks = atoks.permute(1,0,2)
|
|
25
|
+
else:
|
|
26
|
+
q,t = atoks.shape
|
|
27
|
+
atoks = atoks.to(self.device)
|
|
28
|
+
features = self.vocos.codes_to_features(atoks)
|
|
29
|
+
bandwidth_id = torch.tensor({2: 0, 4: 1, 8: 2}[q]).to(self.device)
|
|
30
|
+
return self.vocos.decode(features, bandwidth_id=bandwidth_id)
|
|
31
|
+
|
|
32
|
+
def _save_audio(self, fname, audio_tensor, sample_rate=24000):
|
|
33
|
+
try:
|
|
34
|
+
import torchaudio
|
|
35
|
+
torchaudio.save(fname, audio_tensor, sample_rate, backend="soundfile")
|
|
36
|
+
return
|
|
37
|
+
except (ImportError, RuntimeError, TypeError):
|
|
38
|
+
pass
|
|
39
|
+
|
|
40
|
+
try:
|
|
41
|
+
import torchaudio
|
|
42
|
+
torchaudio.save(fname, audio_tensor, sample_rate)
|
|
43
|
+
return
|
|
44
|
+
except (ImportError, RuntimeError):
|
|
45
|
+
pass
|
|
46
|
+
|
|
47
|
+
try:
|
|
48
|
+
import soundfile as sf
|
|
49
|
+
audio_np = audio_tensor.numpy().T
|
|
50
|
+
sf.write(fname, audio_np, sample_rate)
|
|
51
|
+
return
|
|
52
|
+
except ImportError:
|
|
53
|
+
pass
|
|
54
|
+
|
|
55
|
+
raise ImportError(
|
|
56
|
+
"No audio backend available. Please install either torchaudio or soundfile:\n"
|
|
57
|
+
" pip install torchaudio\n"
|
|
58
|
+
"or\n"
|
|
59
|
+
" pip install soundfile"
|
|
60
|
+
)
|
|
61
|
+
|
|
62
|
+
def decode_to_file(self, fname, atoks):
|
|
63
|
+
audio = self.decode(atoks)
|
|
64
|
+
self._save_audio(fname, audio.cpu(), 24000)
|
|
65
|
+
if self.is_notebook():
|
|
66
|
+
from IPython.display import display, HTML, Audio
|
|
67
|
+
display(HTML(f'<a href="{fname}" target="_blank">Listen to {fname}</a>'))
|
|
68
|
+
|
|
69
|
+
def decode_to_notebook(self, atoks):
|
|
70
|
+
from IPython.display import display, HTML, Audio
|
|
71
|
+
audio = self.decode(atoks)
|
|
72
|
+
display(Audio(audio.cpu().numpy(), rate=24000))
|
|
@@ -0,0 +1,56 @@
|
|
|
1
|
+
__all__ = ['get_compute_device']
|
|
2
|
+
|
|
3
|
+
import torch
|
|
4
|
+
import torch.nn.functional as F
|
|
5
|
+
from huggingface_hub import hf_hub_download
|
|
6
|
+
from contextlib import nullcontext
|
|
7
|
+
|
|
8
|
+
def get_default_compute_device():
|
|
9
|
+
if torch.cuda.is_available() and (torch.version.cuda or torch.version.hip):
|
|
10
|
+
return 'cuda'
|
|
11
|
+
elif torch.backends.mps.is_available():
|
|
12
|
+
return 'mps'
|
|
13
|
+
else:
|
|
14
|
+
return 'cpu'
|
|
15
|
+
|
|
16
|
+
preferred_device = None
|
|
17
|
+
|
|
18
|
+
def get_compute_device():
|
|
19
|
+
global preferred_device
|
|
20
|
+
if preferred_device is None: preferred_device = get_default_compute_device()
|
|
21
|
+
return preferred_device
|
|
22
|
+
|
|
23
|
+
def load_model(ref=None, spec=None, device='cpu', cache_dir=None):
|
|
24
|
+
if spec is not None: return spec
|
|
25
|
+
if ":" in ref:
|
|
26
|
+
repo_id, filename = ref.split(":", 1)
|
|
27
|
+
local_filename = hf_hub_download(repo_id=repo_id, filename=filename, cache_dir=cache_dir)
|
|
28
|
+
else:
|
|
29
|
+
local_filename = ref
|
|
30
|
+
return torch.load(local_filename, map_location=device)
|
|
31
|
+
|
|
32
|
+
def inference_context():
|
|
33
|
+
if torch.cuda.is_available():
|
|
34
|
+
return torch.backends.cuda.sdp_kernel(enable_flash=False, enable_mem_efficient=False, enable_math=True)
|
|
35
|
+
else:
|
|
36
|
+
return nullcontext()
|
|
37
|
+
|
|
38
|
+
def multinomial_sample_one_no_sync(probs_sort):
|
|
39
|
+
q = torch.empty_like(probs_sort).exponential_(1)
|
|
40
|
+
return torch.argmax(probs_sort / q, dim=-1, keepdim=True).to(dtype=torch.int)
|
|
41
|
+
|
|
42
|
+
def logits_to_probs(logits, T=1.0, top_k=None):
|
|
43
|
+
logits = logits / max(T, 1e-5)
|
|
44
|
+
|
|
45
|
+
if top_k is not None:
|
|
46
|
+
v, _ = torch.topk(logits, min(top_k, logits.size(-1)))
|
|
47
|
+
pivot = v.select(-1, -1).unsqueeze(-1)
|
|
48
|
+
logits = torch.where(logits < pivot, -float("Inf"), logits)
|
|
49
|
+
|
|
50
|
+
probs = torch.nn.functional.softmax(logits, dim=-1)
|
|
51
|
+
return probs
|
|
52
|
+
|
|
53
|
+
def sample(logits, T=1.0, top_k=None):
|
|
54
|
+
probs = logits_to_probs(logits, T, top_k)
|
|
55
|
+
idx_next = multinomial_sample_one_no_sync(probs)
|
|
56
|
+
return idx_next
|
|
@@ -0,0 +1,131 @@
|
|
|
1
|
+
# AUTOGENERATED! DO NOT EDIT! File to edit: ../nbs/B. Languages.ipynb.
|
|
2
|
+
|
|
3
|
+
# %% auto 0
|
|
4
|
+
__all__ = ['to_id']
|
|
5
|
+
|
|
6
|
+
# %% ../nbs/B. Languages.ipynb 3
|
|
7
|
+
LANGUAGES = {
|
|
8
|
+
"en": "english",
|
|
9
|
+
"zh": "chinese",
|
|
10
|
+
"de": "german",
|
|
11
|
+
"es": "spanish",
|
|
12
|
+
"ru": "russian",
|
|
13
|
+
"ko": "korean",
|
|
14
|
+
"fr": "french",
|
|
15
|
+
"ja": "japanese",
|
|
16
|
+
"pt": "portuguese",
|
|
17
|
+
"tr": "turkish",
|
|
18
|
+
"pl": "polish",
|
|
19
|
+
"ca": "catalan",
|
|
20
|
+
"nl": "dutch",
|
|
21
|
+
"ar": "arabic",
|
|
22
|
+
"sv": "swedish",
|
|
23
|
+
"it": "italian",
|
|
24
|
+
"id": "indonesian",
|
|
25
|
+
"hi": "hindi",
|
|
26
|
+
"fi": "finnish",
|
|
27
|
+
"vi": "vietnamese",
|
|
28
|
+
"he": "hebrew",
|
|
29
|
+
"uk": "ukrainian",
|
|
30
|
+
"el": "greek",
|
|
31
|
+
"ms": "malay",
|
|
32
|
+
"cs": "czech",
|
|
33
|
+
"ro": "romanian",
|
|
34
|
+
"da": "danish",
|
|
35
|
+
"hu": "hungarian",
|
|
36
|
+
"ta": "tamil",
|
|
37
|
+
"no": "norwegian",
|
|
38
|
+
"th": "thai",
|
|
39
|
+
"ur": "urdu",
|
|
40
|
+
"hr": "croatian",
|
|
41
|
+
"bg": "bulgarian",
|
|
42
|
+
"lt": "lithuanian",
|
|
43
|
+
"la": "latin",
|
|
44
|
+
"mi": "maori",
|
|
45
|
+
"ml": "malayalam",
|
|
46
|
+
"cy": "welsh",
|
|
47
|
+
"sk": "slovak",
|
|
48
|
+
"te": "telugu",
|
|
49
|
+
"fa": "persian",
|
|
50
|
+
"lv": "latvian",
|
|
51
|
+
"bn": "bengali",
|
|
52
|
+
"sr": "serbian",
|
|
53
|
+
"az": "azerbaijani",
|
|
54
|
+
"sl": "slovenian",
|
|
55
|
+
"kn": "kannada",
|
|
56
|
+
"et": "estonian",
|
|
57
|
+
"mk": "macedonian",
|
|
58
|
+
"br": "breton",
|
|
59
|
+
"eu": "basque",
|
|
60
|
+
"is": "icelandic",
|
|
61
|
+
"hy": "armenian",
|
|
62
|
+
"ne": "nepali",
|
|
63
|
+
"mn": "mongolian",
|
|
64
|
+
"bs": "bosnian",
|
|
65
|
+
"kk": "kazakh",
|
|
66
|
+
"sq": "albanian",
|
|
67
|
+
"sw": "swahili",
|
|
68
|
+
"gl": "galician",
|
|
69
|
+
"mr": "marathi",
|
|
70
|
+
"pa": "punjabi",
|
|
71
|
+
"si": "sinhala",
|
|
72
|
+
"km": "khmer",
|
|
73
|
+
"sn": "shona",
|
|
74
|
+
"yo": "yoruba",
|
|
75
|
+
"so": "somali",
|
|
76
|
+
"af": "afrikaans",
|
|
77
|
+
"oc": "occitan",
|
|
78
|
+
"ka": "georgian",
|
|
79
|
+
"be": "belarusian",
|
|
80
|
+
"tg": "tajik",
|
|
81
|
+
"sd": "sindhi",
|
|
82
|
+
"gu": "gujarati",
|
|
83
|
+
"am": "amharic",
|
|
84
|
+
"yi": "yiddish",
|
|
85
|
+
"lo": "lao",
|
|
86
|
+
"uz": "uzbek",
|
|
87
|
+
"fo": "faroese",
|
|
88
|
+
"ht": "haitian creole",
|
|
89
|
+
"ps": "pashto",
|
|
90
|
+
"tk": "turkmen",
|
|
91
|
+
"nn": "nynorsk",
|
|
92
|
+
"mt": "maltese",
|
|
93
|
+
"sa": "sanskrit",
|
|
94
|
+
"lb": "luxembourgish",
|
|
95
|
+
"my": "myanmar",
|
|
96
|
+
"bo": "tibetan",
|
|
97
|
+
"tl": "tagalog",
|
|
98
|
+
"mg": "malagasy",
|
|
99
|
+
"as": "assamese",
|
|
100
|
+
"tt": "tatar",
|
|
101
|
+
"haw": "hawaiian",
|
|
102
|
+
"ln": "lingala",
|
|
103
|
+
"ha": "hausa",
|
|
104
|
+
"ba": "bashkir",
|
|
105
|
+
"jw": "javanese",
|
|
106
|
+
"su": "sundanese",
|
|
107
|
+
}
|
|
108
|
+
|
|
109
|
+
# %% ../nbs/B. Languages.ipynb 4
|
|
110
|
+
# language code lookup by name, with a few language aliases
|
|
111
|
+
TO_LANGUAGE_CODE = {
|
|
112
|
+
**{language: code for code, language in LANGUAGES.items()},
|
|
113
|
+
"burmese": "my",
|
|
114
|
+
"valencian": "ca",
|
|
115
|
+
"flemish": "nl",
|
|
116
|
+
"haitian": "ht",
|
|
117
|
+
"letzeburgesch": "lb",
|
|
118
|
+
"pushto": "ps",
|
|
119
|
+
"panjabi": "pa",
|
|
120
|
+
"moldavian": "ro",
|
|
121
|
+
"moldovan": "ro",
|
|
122
|
+
"sinhalese": "si",
|
|
123
|
+
"castilian": "es",
|
|
124
|
+
}
|
|
125
|
+
|
|
126
|
+
# %% ../nbs/B. Languages.ipynb 5
|
|
127
|
+
languages = tuple(LANGUAGES.keys())
|
|
128
|
+
|
|
129
|
+
# %% ../nbs/B. Languages.ipynb 6
|
|
130
|
+
def to_id(lang):
|
|
131
|
+
return languages.index(TO_LANGUAGE_CODE.get(lang, lang))
|