wangli commited on
Upload folder using huggingface_hub
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- .gitattributes +1 -0
- .gitignore +2 -1
- MANIFEST.in +5 -0
- README.md +55 -7
- ax_meeting/__init__.py +5 -0
- ax_meeting/ax_model/.gitattributes +2 -0
- ax_meeting/ax_model/auto.npy +3 -0
- ax_meeting/ax_model/campplus.axmodel +3 -0
- ax_meeting/ax_model/chn_jpn_yue_eng_ko_spectok.bpe.model +3 -0
- ax_meeting/ax_model/en.npy +3 -0
- ax_meeting/ax_model/event_emo.npy +3 -0
- ax_meeting/ax_model/ja.npy +3 -0
- ax_meeting/ax_model/ko.npy +3 -0
- ax_meeting/ax_model/sensevoice.axmodel +3 -0
- ax_meeting/ax_model/sensevoice/am.mvn +8 -0
- ax_meeting/ax_model/sensevoice/config.yaml +97 -0
- ax_meeting/ax_model/vad.axmodel +3 -0
- ax_meeting/ax_model/vad/am.mvn +8 -0
- ax_meeting/ax_model/vad/config.yaml +56 -0
- ax_meeting/ax_model/withitn.npy +3 -0
- ax_meeting/ax_model/yue.npy +3 -0
- ax_meeting/ax_model/zh.npy +3 -0
- ax_meeting/axengine_loader.py +34 -0
- ax_meeting/certs/cert.pem +19 -0
- ax_meeting/certs/key.pem +28 -0
- ax_meeting/config.py +19 -0
- ax_meeting/diar_asr_cli.py +93 -0
- ax_meeting/diar_utils.py +13 -0
- ax_meeting/engines.py +207 -0
- ax_meeting/model_bundle.py +87 -0
- ax_meeting/pipeline.py +163 -0
- ax_meeting/positional.py +18 -0
- ax_meeting/server.py +135 -0
- ax_meeting/static/app.js +187 -0
- ax_meeting/static/index.html +50 -0
- ax_meeting/static/style.css +147 -0
- ax_meeting/summarize_cli.py +33 -0
- ax_meeting/summarizer.py +75 -0
- ax_meeting/text_cleaner.py +41 -0
- ax_meeting/utils/__init__.py +0 -0
- ax_meeting/utils/ax_cam_bin.py +231 -0
- ax_meeting/utils/ax_model_bin.py +307 -0
- ax_meeting/utils/ax_vad_bin.py +158 -0
- ax_meeting/utils/cluster_utils.py +241 -0
- ax_meeting/utils/ctc_alignment.py +76 -0
- ax_meeting/utils/frontend.py +433 -0
- ax_meeting/utils/infer_func.py +273 -0
- ax_meeting/utils/infer_utils.py +312 -0
- ax_meeting/utils/sentencepiece_tokenizer.py +46 -0
- ax_meeting/utils/speaker_fbank.py +22 -0
.gitattributes
CHANGED
|
@@ -40,3 +40,4 @@ wav/vad_example.wav filter=lfs diff=lfs merge=lfs -text
|
|
| 40 |
*.json filter=lfs diff=lfs merge=lfs -text
|
| 41 |
assert/gradio_demo.JPG filter=lfs diff=lfs merge=lfs -text
|
| 42 |
wav/20200327_2P.wav filter=lfs diff=lfs merge=lfs -text
|
|
|
|
|
|
| 40 |
*.json filter=lfs diff=lfs merge=lfs -text
|
| 41 |
assert/gradio_demo.JPG filter=lfs diff=lfs merge=lfs -text
|
| 42 |
wav/20200327_2P.wav filter=lfs diff=lfs merge=lfs -text
|
| 43 |
+
dist/ax_meeting-0.1.0-py3-none-any.whl filter=lfs diff=lfs merge=lfs -text
|
.gitignore
CHANGED
|
@@ -1 +1,2 @@
|
|
| 1 |
-
__pycache__
|
|
|
|
|
|
| 1 |
+
__pycache__
|
| 2 |
+
ax_meeting.egg-info
|
MANIFEST.in
ADDED
|
@@ -0,0 +1,5 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
recursive-include ax_meeting/static *
|
| 2 |
+
recursive-include ax_meeting/ax_model *
|
| 3 |
+
recursive-include ax_meeting/certs *
|
| 4 |
+
recursive-include ax_meeting/vendor *
|
| 5 |
+
include README.md
|
README.md
CHANGED
|
@@ -60,7 +60,7 @@ pip3 install -r requirements.txt
|
|
| 60 |
启动:
|
| 61 |
|
| 62 |
```bash
|
| 63 |
-
python -m
|
| 64 |
```
|
| 65 |
|
| 66 |
浏览器访问:
|
|
@@ -79,6 +79,7 @@ HOST=0.0.0.0
|
|
| 79 |
PORT=8000
|
| 80 |
SSL_CERT=cert.pem
|
| 81 |
SSL_KEY=key.pem
|
|
|
|
| 82 |
```
|
| 83 |
|
| 84 |
依赖提示(WebSocket):
|
|
@@ -87,6 +88,9 @@ SSL_KEY=key.pem
|
|
| 87 |
设备权限提示:
|
| 88 |
- 如果遇到 `/dev/axcl_host` 权限错误,请用有权限的账号或 `sudo` 运行
|
| 89 |
|
|
|
|
|
|
|
|
|
|
| 90 |
HTTPS(推荐,便于浏览器麦克风权限):
|
| 91 |
|
| 92 |
```bash
|
|
@@ -96,9 +100,23 @@ openssl req -x509 -newkey rsa:2048 -nodes \\
|
|
| 96 |
```
|
| 97 |
|
| 98 |
```bash
|
| 99 |
-
SSL_CERT=cert.pem SSL_KEY=key.pem python -m
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 100 |
```
|
| 101 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 102 |

|
| 103 |
|
| 104 |
|
|
@@ -107,21 +125,51 @@ SSL_CERT=cert.pem SSL_KEY=key.pem python -m app.server
|
|
| 107 |
对单个会议音频文件执行说话人聚类 + ASR,并导出文本,可选会议总结(LLM 通过参数配置):
|
| 108 |
|
| 109 |
```bash
|
| 110 |
-
python
|
| 111 |
```
|
| 112 |
|
| 113 |
-
|
| 114 |
|
| 115 |
```bash
|
| 116 |
-
python
|
| 117 |
--wav_file wav/vad_example.wav \\
|
| 118 |
-
--output_dir output_dir
|
| 119 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 120 |
--openai_base_url http://127.0.0.1:8001/v1 \\
|
| 121 |
--openai_model AXERA-TECH/Qwen3-1.7B \\
|
| 122 |
--openai_api_key xxx
|
| 123 |
```
|
| 124 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 125 |
## Latency
|
| 126 |
|
| 127 |
AX650N
|
|
|
|
| 60 |
启动:
|
| 61 |
|
| 62 |
```bash
|
| 63 |
+
python -m ax_meeting.server
|
| 64 |
```
|
| 65 |
|
| 66 |
浏览器访问:
|
|
|
|
| 79 |
PORT=8000
|
| 80 |
SSL_CERT=cert.pem
|
| 81 |
SSL_KEY=key.pem
|
| 82 |
+
AX_MODEL_DIR=/path/to/ax_model
|
| 83 |
```
|
| 84 |
|
| 85 |
依赖提示(WebSocket):
|
|
|
|
| 88 |
设备权限提示:
|
| 89 |
- 如果遇到 `/dev/axcl_host` 权限错误,请用有权限的账号或 `sudo` 运行
|
| 90 |
|
| 91 |
+
axengine 依赖提示:
|
| 92 |
+
- 如果无法从 pip 获取 `pyaxengine`,请将本地 wheel 放到 `ax_meeting/vendor/`,或设置 `AXENGINE_WHEEL=/path/to/pyaxengine.whl`
|
| 93 |
+
|
| 94 |
HTTPS(推荐,便于浏览器麦克风权限):
|
| 95 |
|
| 96 |
```bash
|
|
|
|
| 100 |
```
|
| 101 |
|
| 102 |
```bash
|
| 103 |
+
SSL_CERT=cert.pem SSL_KEY=key.pem python -m ax_meeting.server
|
| 104 |
+
```
|
| 105 |
+
|
| 106 |
+
使用包内自签证书(默认打包在 `ax_meeting/certs/`):
|
| 107 |
+
|
| 108 |
+
```bash
|
| 109 |
+
SSL_CERT=ax_meeting/certs/cert.pem SSL_KEY=ax_meeting/certs/key.pem python -m ax_meeting.server
|
| 110 |
```
|
| 111 |
|
| 112 |
+
## 生成 wheel 包
|
| 113 |
+
|
| 114 |
+
```bash
|
| 115 |
+
./build_wheel.sh
|
| 116 |
+
```
|
| 117 |
+
|
| 118 |
+
生成结果在 `dist/` 目录。
|
| 119 |
+
|
| 120 |

|
| 121 |
|
| 122 |
|
|
|
|
| 125 |
对单个会议音频文件执行说话人聚类 + ASR,并导出文本,可选会议总结(LLM 通过参数配置):
|
| 126 |
|
| 127 |
```bash
|
| 128 |
+
python -m ax_meeting.vad_asr_cli --input wav/vad_example.wav --output_dir output_dir
|
| 129 |
```
|
| 130 |
|
| 131 |
+
说话人 + ASR(离线):
|
| 132 |
|
| 133 |
```bash
|
| 134 |
+
python -m ax_meeting.diar_asr_cli \\
|
| 135 |
--wav_file wav/vad_example.wav \\
|
| 136 |
+
--output_dir output_dir
|
| 137 |
+
```
|
| 138 |
+
|
| 139 |
+
会议总结:
|
| 140 |
+
|
| 141 |
+
```bash
|
| 142 |
+
python -m ax_meeting.summarize_cli \\
|
| 143 |
+
--input output_dir/vad_example.txt \\
|
| 144 |
--openai_base_url http://127.0.0.1:8001/v1 \\
|
| 145 |
--openai_model AXERA-TECH/Qwen3-1.7B \\
|
| 146 |
--openai_api_key xxx
|
| 147 |
```
|
| 148 |
|
| 149 |
+
## Python API(轻量)
|
| 150 |
+
|
| 151 |
+
```python
|
| 152 |
+
from ax_meeting import VadAsrEngine, DiarAsrEngine, IncrementalSummarizer
|
| 153 |
+
|
| 154 |
+
# VAD + ASR(流式)
|
| 155 |
+
vad_asr = VadAsrEngine(stream=True)
|
| 156 |
+
vad_asr.feed(audio_chunk) # numpy / bytes / path / list
|
| 157 |
+
segments = vad_asr.poll() # 可能为空
|
| 158 |
+
|
| 159 |
+
# 说话人 + ASR(离线)
|
| 160 |
+
diar = DiarAsrEngine()
|
| 161 |
+
text = diar.transcribe("wav/vad_example.wav")
|
| 162 |
+
|
| 163 |
+
# 会议总结
|
| 164 |
+
summarizer = IncrementalSummarizer()
|
| 165 |
+
summary = summarizer.summarize_incrementally(text)
|
| 166 |
+
```
|
| 167 |
+
|
| 168 |
+
示例脚本:
|
| 169 |
+
- `examples/vad_asr_stream.py`
|
| 170 |
+
- `examples/diar_asr_offline.py`
|
| 171 |
+
- `examples/summarize_text.py`
|
| 172 |
+
|
| 173 |
## Latency
|
| 174 |
|
| 175 |
AX650N
|
ax_meeting/__init__.py
ADDED
|
@@ -0,0 +1,5 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# -*- coding: utf-8 -*-
|
| 2 |
+
__version__ = "0.1.0"
|
| 3 |
+
|
| 4 |
+
from ax_meeting.engines import VadAsrEngine, DiarAsrEngine # noqa: F401
|
| 5 |
+
from ax_meeting.summarizer import IncrementalSummarizer # noqa: F401
|
ax_meeting/ax_model/.gitattributes
ADDED
|
@@ -0,0 +1,2 @@
|
|
|
|
|
|
|
|
|
|
| 1 |
+
*.axmodel filter=lfs diff=lfs merge=lfs -text
|
| 2 |
+
*.npy filter=lfs diff=lfs merge=lfs -text
|
ax_meeting/ax_model/auto.npy
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:8d0997706b30274f7ff3b157ca90df50b7ed8ced35091a0231700355d5ee1374
|
| 3 |
+
size 2368
|
ax_meeting/ax_model/campplus.axmodel
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:3e85b4d59a94488caa727ca76c7fe2ba408de02587cd8c9a6680b83488338b16
|
| 3 |
+
size 10744358
|
ax_meeting/ax_model/chn_jpn_yue_eng_ko_spectok.bpe.model
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:aa87f86064c3730d799ddf7af3c04659151102cba548bce325cf06ba4da4e6a8
|
| 3 |
+
size 377341
|
ax_meeting/ax_model/en.npy
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:1f21c6d5a0b876c696b15ee80317ee558b6a19a084c1147a13a2871900f6c72e
|
| 3 |
+
size 2368
|
ax_meeting/ax_model/event_emo.npy
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:1d22e3df5d192fdc3e73e368a2cb576975a5a43a114a8432a91c036adf8e2263
|
| 3 |
+
size 4608
|
ax_meeting/ax_model/ja.npy
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:75a83ded838de7d02a56de32b774a50d2f02b5ec69880e7acdce600666f77517
|
| 3 |
+
size 2368
|
ax_meeting/ax_model/ko.npy
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:f371fffcecd75c1268e2ef8043c276ec9b231177b5138b39e5732c65401749db
|
| 3 |
+
size 2368
|
ax_meeting/ax_model/sensevoice.axmodel
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:7b64a36fa15e75ab5e3b75f18ae87a058970cff76219407e503b54fb53dd8e38
|
| 3 |
+
size 262170623
|
ax_meeting/ax_model/sensevoice/am.mvn
ADDED
|
@@ -0,0 +1,8 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
<Nnet>
|
| 2 |
+
<Splice> 560 560
|
| 3 |
+
[ 0 ]
|
| 4 |
+
<AddShift> 560 560
|
| 5 |
+
<LearnRateCoef> 0 [ -8.311879 -8.600912 -9.615928 -10.43595 -11.21292 -11.88333 -12.36243 -12.63706 -12.8818 -12.83066 -12.89103 -12.95666 -13.19763 -13.40598 -13.49113 -13.5546 -13.55639 -13.51915 -13.68284 -13.53289 -13.42107 -13.65519 -13.50713 -13.75251 -13.76715 -13.87408 -13.73109 -13.70412 -13.56073 -13.53488 -13.54895 -13.56228 -13.59408 -13.62047 -13.64198 -13.66109 -13.62669 -13.58297 -13.57387 -13.4739 -13.53063 -13.48348 -13.61047 -13.64716 -13.71546 -13.79184 -13.90614 -14.03098 -14.18205 -14.35881 -14.48419 -14.60172 -14.70591 -14.83362 -14.92122 -15.00622 -15.05122 -15.03119 -14.99028 -14.92302 -14.86927 -14.82691 -14.7972 -14.76909 -14.71356 -14.61277 -14.51696 -14.42252 -14.36405 -14.30451 -14.23161 -14.19851 -14.16633 -14.15649 -14.10504 -13.99518 -13.79562 -13.3996 -12.7767 -11.71208 -8.311879 -8.600912 -9.615928 -10.43595 -11.21292 -11.88333 -12.36243 -12.63706 -12.8818 -12.83066 -12.89103 -12.95666 -13.19763 -13.40598 -13.49113 -13.5546 -13.55639 -13.51915 -13.68284 -13.53289 -13.42107 -13.65519 -13.50713 -13.75251 -13.76715 -13.87408 -13.73109 -13.70412 -13.56073 -13.53488 -13.54895 -13.56228 -13.59408 -13.62047 -13.64198 -13.66109 -13.62669 -13.58297 -13.57387 -13.4739 -13.53063 -13.48348 -13.61047 -13.64716 -13.71546 -13.79184 -13.90614 -14.03098 -14.18205 -14.35881 -14.48419 -14.60172 -14.70591 -14.83362 -14.92122 -15.00622 -15.05122 -15.03119 -14.99028 -14.92302 -14.86927 -14.82691 -14.7972 -14.76909 -14.71356 -14.61277 -14.51696 -14.42252 -14.36405 -14.30451 -14.23161 -14.19851 -14.16633 -14.15649 -14.10504 -13.99518 -13.79562 -13.3996 -12.7767 -11.71208 -8.311879 -8.600912 -9.615928 -10.43595 -11.21292 -11.88333 -12.36243 -12.63706 -12.8818 -12.83066 -12.89103 -12.95666 -13.19763 -13.40598 -13.49113 -13.5546 -13.55639 -13.51915 -13.68284 -13.53289 -13.42107 -13.65519 -13.50713 -13.75251 -13.76715 -13.87408 -13.73109 -13.70412 -13.56073 -13.53488 -13.54895 -13.56228 -13.59408 -13.62047 -13.64198 -13.66109 -13.62669 -13.58297 -13.57387 -13.4739 -13.53063 -13.48348 -13.61047 -13.64716 -13.71546 -13.79184 -13.90614 -14.03098 -14.18205 -14.35881 -14.48419 -14.60172 -14.70591 -14.83362 -14.92122 -15.00622 -15.05122 -15.03119 -14.99028 -14.92302 -14.86927 -14.82691 -14.7972 -14.76909 -14.71356 -14.61277 -14.51696 -14.42252 -14.36405 -14.30451 -14.23161 -14.19851 -14.16633 -14.15649 -14.10504 -13.99518 -13.79562 -13.3996 -12.7767 -11.71208 -8.311879 -8.600912 -9.615928 -10.43595 -11.21292 -11.88333 -12.36243 -12.63706 -12.8818 -12.83066 -12.89103 -12.95666 -13.19763 -13.40598 -13.49113 -13.5546 -13.55639 -13.51915 -13.68284 -13.53289 -13.42107 -13.65519 -13.50713 -13.75251 -13.76715 -13.87408 -13.73109 -13.70412 -13.56073 -13.53488 -13.54895 -13.56228 -13.59408 -13.62047 -13.64198 -13.66109 -13.62669 -13.58297 -13.57387 -13.4739 -13.53063 -13.48348 -13.61047 -13.64716 -13.71546 -13.79184 -13.90614 -14.03098 -14.18205 -14.35881 -14.48419 -14.60172 -14.70591 -14.83362 -14.92122 -15.00622 -15.05122 -15.03119 -14.99028 -14.92302 -14.86927 -14.82691 -14.7972 -14.76909 -14.71356 -14.61277 -14.51696 -14.42252 -14.36405 -14.30451 -14.23161 -14.19851 -14.16633 -14.15649 -14.10504 -13.99518 -13.79562 -13.3996 -12.7767 -11.71208 -8.311879 -8.600912 -9.615928 -10.43595 -11.21292 -11.88333 -12.36243 -12.63706 -12.8818 -12.83066 -12.89103 -12.95666 -13.19763 -13.40598 -13.49113 -13.5546 -13.55639 -13.51915 -13.68284 -13.53289 -13.42107 -13.65519 -13.50713 -13.75251 -13.76715 -13.87408 -13.73109 -13.70412 -13.56073 -13.53488 -13.54895 -13.56228 -13.59408 -13.62047 -13.64198 -13.66109 -13.62669 -13.58297 -13.57387 -13.4739 -13.53063 -13.48348 -13.61047 -13.64716 -13.71546 -13.79184 -13.90614 -14.03098 -14.18205 -14.35881 -14.48419 -14.60172 -14.70591 -14.83362 -14.92122 -15.00622 -15.05122 -15.03119 -14.99028 -14.92302 -14.86927 -14.82691 -14.7972 -14.76909 -14.71356 -14.61277 -14.51696 -14.42252 -14.36405 -14.30451 -14.23161 -14.19851 -14.16633 -14.15649 -14.10504 -13.99518 -13.79562 -13.3996 -12.7767 -11.71208 -8.311879 -8.600912 -9.615928 -10.43595 -11.21292 -11.88333 -12.36243 -12.63706 -12.8818 -12.83066 -12.89103 -12.95666 -13.19763 -13.40598 -13.49113 -13.5546 -13.55639 -13.51915 -13.68284 -13.53289 -13.42107 -13.65519 -13.50713 -13.75251 -13.76715 -13.87408 -13.73109 -13.70412 -13.56073 -13.53488 -13.54895 -13.56228 -13.59408 -13.62047 -13.64198 -13.66109 -13.62669 -13.58297 -13.57387 -13.4739 -13.53063 -13.48348 -13.61047 -13.64716 -13.71546 -13.79184 -13.90614 -14.03098 -14.18205 -14.35881 -14.48419 -14.60172 -14.70591 -14.83362 -14.92122 -15.00622 -15.05122 -15.03119 -14.99028 -14.92302 -14.86927 -14.82691 -14.7972 -14.76909 -14.71356 -14.61277 -14.51696 -14.42252 -14.36405 -14.30451 -14.23161 -14.19851 -14.16633 -14.15649 -14.10504 -13.99518 -13.79562 -13.3996 -12.7767 -11.71208 -8.311879 -8.600912 -9.615928 -10.43595 -11.21292 -11.88333 -12.36243 -12.63706 -12.8818 -12.83066 -12.89103 -12.95666 -13.19763 -13.40598 -13.49113 -13.5546 -13.55639 -13.51915 -13.68284 -13.53289 -13.42107 -13.65519 -13.50713 -13.75251 -13.76715 -13.87408 -13.73109 -13.70412 -13.56073 -13.53488 -13.54895 -13.56228 -13.59408 -13.62047 -13.64198 -13.66109 -13.62669 -13.58297 -13.57387 -13.4739 -13.53063 -13.48348 -13.61047 -13.64716 -13.71546 -13.79184 -13.90614 -14.03098 -14.18205 -14.35881 -14.48419 -14.60172 -14.70591 -14.83362 -14.92122 -15.00622 -15.05122 -15.03119 -14.99028 -14.92302 -14.86927 -14.82691 -14.7972 -14.76909 -14.71356 -14.61277 -14.51696 -14.42252 -14.36405 -14.30451 -14.23161 -14.19851 -14.16633 -14.15649 -14.10504 -13.99518 -13.79562 -13.3996 -12.7767 -11.71208 ]
|
| 6 |
+
<Rescale> 560 560
|
| 7 |
+
<LearnRateCoef> 0 [ 0.155775 0.154484 0.1527379 0.1518718 0.1506028 0.1489256 0.147067 0.1447061 0.1436307 0.1443568 0.1451849 0.1455157 0.1452821 0.1445717 0.1439195 0.1435867 0.1436018 0.1438781 0.1442086 0.1448844 0.1454756 0.145663 0.146268 0.1467386 0.1472724 0.147664 0.1480913 0.1483739 0.1488841 0.1493636 0.1497088 0.1500379 0.1502916 0.1505389 0.1506787 0.1507102 0.1505992 0.1505445 0.1505938 0.1508133 0.1509569 0.1512396 0.1514625 0.1516195 0.1516156 0.1515561 0.1514966 0.1513976 0.1512612 0.151076 0.1510596 0.1510431 0.151077 0.1511168 0.1511917 0.151023 0.1508045 0.1505885 0.1503493 0.1502373 0.1501726 0.1500762 0.1500065 0.1499782 0.150057 0.1502658 0.150469 0.1505335 0.1505505 0.1505328 0.1504275 0.1502438 0.1499674 0.1497118 0.1494661 0.1493102 0.1493681 0.1495501 0.1499738 0.1509654 0.155775 0.154484 0.1527379 0.1518718 0.1506028 0.1489256 0.147067 0.1447061 0.1436307 0.1443568 0.1451849 0.1455157 0.1452821 0.1445717 0.1439195 0.1435867 0.1436018 0.1438781 0.1442086 0.1448844 0.1454756 0.145663 0.146268 0.1467386 0.1472724 0.147664 0.1480913 0.1483739 0.1488841 0.1493636 0.1497088 0.1500379 0.1502916 0.1505389 0.1506787 0.1507102 0.1505992 0.1505445 0.1505938 0.1508133 0.1509569 0.1512396 0.1514625 0.1516195 0.1516156 0.1515561 0.1514966 0.1513976 0.1512612 0.151076 0.1510596 0.1510431 0.151077 0.1511168 0.1511917 0.151023 0.1508045 0.1505885 0.1503493 0.1502373 0.1501726 0.1500762 0.1500065 0.1499782 0.150057 0.1502658 0.150469 0.1505335 0.1505505 0.1505328 0.1504275 0.1502438 0.1499674 0.1497118 0.1494661 0.1493102 0.1493681 0.1495501 0.1499738 0.1509654 0.155775 0.154484 0.1527379 0.1518718 0.1506028 0.1489256 0.147067 0.1447061 0.1436307 0.1443568 0.1451849 0.1455157 0.1452821 0.1445717 0.1439195 0.1435867 0.1436018 0.1438781 0.1442086 0.1448844 0.1454756 0.145663 0.146268 0.1467386 0.1472724 0.147664 0.1480913 0.1483739 0.1488841 0.1493636 0.1497088 0.1500379 0.1502916 0.1505389 0.1506787 0.1507102 0.1505992 0.1505445 0.1505938 0.1508133 0.1509569 0.1512396 0.1514625 0.1516195 0.1516156 0.1515561 0.1514966 0.1513976 0.1512612 0.151076 0.1510596 0.1510431 0.151077 0.1511168 0.1511917 0.151023 0.1508045 0.1505885 0.1503493 0.1502373 0.1501726 0.1500762 0.1500065 0.1499782 0.150057 0.1502658 0.150469 0.1505335 0.1505505 0.1505328 0.1504275 0.1502438 0.1499674 0.1497118 0.1494661 0.1493102 0.1493681 0.1495501 0.1499738 0.1509654 0.155775 0.154484 0.1527379 0.1518718 0.1506028 0.1489256 0.147067 0.1447061 0.1436307 0.1443568 0.1451849 0.1455157 0.1452821 0.1445717 0.1439195 0.1435867 0.1436018 0.1438781 0.1442086 0.1448844 0.1454756 0.145663 0.146268 0.1467386 0.1472724 0.147664 0.1480913 0.1483739 0.1488841 0.1493636 0.1497088 0.1500379 0.1502916 0.1505389 0.1506787 0.1507102 0.1505992 0.1505445 0.1505938 0.1508133 0.1509569 0.1512396 0.1514625 0.1516195 0.1516156 0.1515561 0.1514966 0.1513976 0.1512612 0.151076 0.1510596 0.1510431 0.151077 0.1511168 0.1511917 0.151023 0.1508045 0.1505885 0.1503493 0.1502373 0.1501726 0.1500762 0.1500065 0.1499782 0.150057 0.1502658 0.150469 0.1505335 0.1505505 0.1505328 0.1504275 0.1502438 0.1499674 0.1497118 0.1494661 0.1493102 0.1493681 0.1495501 0.1499738 0.1509654 0.155775 0.154484 0.1527379 0.1518718 0.1506028 0.1489256 0.147067 0.1447061 0.1436307 0.1443568 0.1451849 0.1455157 0.1452821 0.1445717 0.1439195 0.1435867 0.1436018 0.1438781 0.1442086 0.1448844 0.1454756 0.145663 0.146268 0.1467386 0.1472724 0.147664 0.1480913 0.1483739 0.1488841 0.1493636 0.1497088 0.1500379 0.1502916 0.1505389 0.1506787 0.1507102 0.1505992 0.1505445 0.1505938 0.1508133 0.1509569 0.1512396 0.1514625 0.1516195 0.1516156 0.1515561 0.1514966 0.1513976 0.1512612 0.151076 0.1510596 0.1510431 0.151077 0.1511168 0.1511917 0.151023 0.1508045 0.1505885 0.1503493 0.1502373 0.1501726 0.1500762 0.1500065 0.1499782 0.150057 0.1502658 0.150469 0.1505335 0.1505505 0.1505328 0.1504275 0.1502438 0.1499674 0.1497118 0.1494661 0.1493102 0.1493681 0.1495501 0.1499738 0.1509654 0.155775 0.154484 0.1527379 0.1518718 0.1506028 0.1489256 0.147067 0.1447061 0.1436307 0.1443568 0.1451849 0.1455157 0.1452821 0.1445717 0.1439195 0.1435867 0.1436018 0.1438781 0.1442086 0.1448844 0.1454756 0.145663 0.146268 0.1467386 0.1472724 0.147664 0.1480913 0.1483739 0.1488841 0.1493636 0.1497088 0.1500379 0.1502916 0.1505389 0.1506787 0.1507102 0.1505992 0.1505445 0.1505938 0.1508133 0.1509569 0.1512396 0.1514625 0.1516195 0.1516156 0.1515561 0.1514966 0.1513976 0.1512612 0.151076 0.1510596 0.1510431 0.151077 0.1511168 0.1511917 0.151023 0.1508045 0.1505885 0.1503493 0.1502373 0.1501726 0.1500762 0.1500065 0.1499782 0.150057 0.1502658 0.150469 0.1505335 0.1505505 0.1505328 0.1504275 0.1502438 0.1499674 0.1497118 0.1494661 0.1493102 0.1493681 0.1495501 0.1499738 0.1509654 0.155775 0.154484 0.1527379 0.1518718 0.1506028 0.1489256 0.147067 0.1447061 0.1436307 0.1443568 0.1451849 0.1455157 0.1452821 0.1445717 0.1439195 0.1435867 0.1436018 0.1438781 0.1442086 0.1448844 0.1454756 0.145663 0.146268 0.1467386 0.1472724 0.147664 0.1480913 0.1483739 0.1488841 0.1493636 0.1497088 0.1500379 0.1502916 0.1505389 0.1506787 0.1507102 0.1505992 0.1505445 0.1505938 0.1508133 0.1509569 0.1512396 0.1514625 0.1516195 0.1516156 0.1515561 0.1514966 0.1513976 0.1512612 0.151076 0.1510596 0.1510431 0.151077 0.1511168 0.1511917 0.151023 0.1508045 0.1505885 0.1503493 0.1502373 0.1501726 0.1500762 0.1500065 0.1499782 0.150057 0.1502658 0.150469 0.1505335 0.1505505 0.1505328 0.1504275 0.1502438 0.1499674 0.1497118 0.1494661 0.1493102 0.1493681 0.1495501 0.1499738 0.1509654 ]
|
| 8 |
+
</Nnet>
|
ax_meeting/ax_model/sensevoice/config.yaml
ADDED
|
@@ -0,0 +1,97 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
encoder: SenseVoiceEncoderSmall
|
| 2 |
+
encoder_conf:
|
| 3 |
+
output_size: 512
|
| 4 |
+
attention_heads: 4
|
| 5 |
+
linear_units: 2048
|
| 6 |
+
num_blocks: 50
|
| 7 |
+
tp_blocks: 20
|
| 8 |
+
dropout_rate: 0.1
|
| 9 |
+
positional_dropout_rate: 0.1
|
| 10 |
+
attention_dropout_rate: 0.1
|
| 11 |
+
input_layer: pe
|
| 12 |
+
pos_enc_class: SinusoidalPositionEncoder
|
| 13 |
+
normalize_before: true
|
| 14 |
+
kernel_size: 11
|
| 15 |
+
sanm_shfit: 0
|
| 16 |
+
selfattention_layer_type: sanm
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
model: SenseVoiceSmall
|
| 20 |
+
model_conf:
|
| 21 |
+
length_normalized_loss: true
|
| 22 |
+
sos: 1
|
| 23 |
+
eos: 2
|
| 24 |
+
ignore_id: -1
|
| 25 |
+
|
| 26 |
+
tokenizer: SentencepiecesTokenizer
|
| 27 |
+
tokenizer_conf:
|
| 28 |
+
bpemodel: null
|
| 29 |
+
unk_symbol: <unk>
|
| 30 |
+
split_with_space: true
|
| 31 |
+
|
| 32 |
+
frontend: WavFrontend
|
| 33 |
+
frontend_conf:
|
| 34 |
+
fs: 16000
|
| 35 |
+
window: hamming
|
| 36 |
+
n_mels: 80
|
| 37 |
+
frame_length: 25
|
| 38 |
+
frame_shift: 10
|
| 39 |
+
lfr_m: 7
|
| 40 |
+
lfr_n: 6
|
| 41 |
+
cmvn_file: null
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
dataset: SenseVoiceCTCDataset
|
| 45 |
+
dataset_conf:
|
| 46 |
+
index_ds: IndexDSJsonl
|
| 47 |
+
batch_sampler: EspnetStyleBatchSampler
|
| 48 |
+
data_split_num: 32
|
| 49 |
+
batch_type: token
|
| 50 |
+
batch_size: 14000
|
| 51 |
+
max_token_length: 2000
|
| 52 |
+
min_token_length: 60
|
| 53 |
+
max_source_length: 2000
|
| 54 |
+
min_source_length: 60
|
| 55 |
+
max_target_length: 200
|
| 56 |
+
min_target_length: 0
|
| 57 |
+
shuffle: true
|
| 58 |
+
num_workers: 4
|
| 59 |
+
sos: ${model_conf.sos}
|
| 60 |
+
eos: ${model_conf.eos}
|
| 61 |
+
IndexDSJsonl: IndexDSJsonl
|
| 62 |
+
retry: 20
|
| 63 |
+
|
| 64 |
+
train_conf:
|
| 65 |
+
accum_grad: 1
|
| 66 |
+
grad_clip: 5
|
| 67 |
+
max_epoch: 20
|
| 68 |
+
keep_nbest_models: 10
|
| 69 |
+
avg_nbest_model: 10
|
| 70 |
+
log_interval: 100
|
| 71 |
+
resume: true
|
| 72 |
+
validate_interval: 10000
|
| 73 |
+
save_checkpoint_interval: 10000
|
| 74 |
+
|
| 75 |
+
optim: adamw
|
| 76 |
+
optim_conf:
|
| 77 |
+
lr: 0.00002
|
| 78 |
+
scheduler: warmuplr
|
| 79 |
+
scheduler_conf:
|
| 80 |
+
warmup_steps: 25000
|
| 81 |
+
|
| 82 |
+
specaug: SpecAugLFR
|
| 83 |
+
specaug_conf:
|
| 84 |
+
apply_time_warp: false
|
| 85 |
+
time_warp_window: 5
|
| 86 |
+
time_warp_mode: bicubic
|
| 87 |
+
apply_freq_mask: true
|
| 88 |
+
freq_mask_width_range:
|
| 89 |
+
- 0
|
| 90 |
+
- 30
|
| 91 |
+
lfr_rate: 6
|
| 92 |
+
num_freq_mask: 1
|
| 93 |
+
apply_time_mask: true
|
| 94 |
+
time_mask_width_range:
|
| 95 |
+
- 0
|
| 96 |
+
- 12
|
| 97 |
+
num_time_mask: 1
|
ax_meeting/ax_model/vad.axmodel
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:7d89d1087f14597bc93ae5cf89e1c60d8ee76d7bdffbe679ef6dc002635e6f6b
|
| 3 |
+
size 1150874
|
ax_meeting/ax_model/vad/am.mvn
ADDED
|
@@ -0,0 +1,8 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
<Nnet>
|
| 2 |
+
<Splice> 400 400
|
| 3 |
+
[ 0 ]
|
| 4 |
+
<AddShift> 400 400
|
| 5 |
+
<LearnRateCoef> 0 [ -8.311879 -8.600912 -9.615928 -10.43595 -11.21292 -11.88333 -12.36243 -12.63706 -12.8818 -12.83066 -12.89103 -12.95666 -13.19763 -13.40598 -13.49113 -13.5546 -13.55639 -13.51915 -13.68284 -13.53289 -13.42107 -13.65519 -13.50713 -13.75251 -13.76715 -13.87408 -13.73109 -13.70412 -13.56073 -13.53488 -13.54895 -13.56228 -13.59408 -13.62047 -13.64198 -13.66109 -13.62669 -13.58297 -13.57387 -13.4739 -13.53063 -13.48348 -13.61047 -13.64716 -13.71546 -13.79184 -13.90614 -14.03098 -14.18205 -14.35881 -14.48419 -14.60172 -14.70591 -14.83362 -14.92122 -15.00622 -15.05122 -15.03119 -14.99028 -14.92302 -14.86927 -14.82691 -14.7972 -14.76909 -14.71356 -14.61277 -14.51696 -14.42252 -14.36405 -14.30451 -14.23161 -14.19851 -14.16633 -14.15649 -14.10504 -13.99518 -13.79562 -13.3996 -12.7767 -11.71208 -8.311879 -8.600912 -9.615928 -10.43595 -11.21292 -11.88333 -12.36243 -12.63706 -12.8818 -12.83066 -12.89103 -12.95666 -13.19763 -13.40598 -13.49113 -13.5546 -13.55639 -13.51915 -13.68284 -13.53289 -13.42107 -13.65519 -13.50713 -13.75251 -13.76715 -13.87408 -13.73109 -13.70412 -13.56073 -13.53488 -13.54895 -13.56228 -13.59408 -13.62047 -13.64198 -13.66109 -13.62669 -13.58297 -13.57387 -13.4739 -13.53063 -13.48348 -13.61047 -13.64716 -13.71546 -13.79184 -13.90614 -14.03098 -14.18205 -14.35881 -14.48419 -14.60172 -14.70591 -14.83362 -14.92122 -15.00622 -15.05122 -15.03119 -14.99028 -14.92302 -14.86927 -14.82691 -14.7972 -14.76909 -14.71356 -14.61277 -14.51696 -14.42252 -14.36405 -14.30451 -14.23161 -14.19851 -14.16633 -14.15649 -14.10504 -13.99518 -13.79562 -13.3996 -12.7767 -11.71208 -8.311879 -8.600912 -9.615928 -10.43595 -11.21292 -11.88333 -12.36243 -12.63706 -12.8818 -12.83066 -12.89103 -12.95666 -13.19763 -13.40598 -13.49113 -13.5546 -13.55639 -13.51915 -13.68284 -13.53289 -13.42107 -13.65519 -13.50713 -13.75251 -13.76715 -13.87408 -13.73109 -13.70412 -13.56073 -13.53488 -13.54895 -13.56228 -13.59408 -13.62047 -13.64198 -13.66109 -13.62669 -13.58297 -13.57387 -13.4739 -13.53063 -13.48348 -13.61047 -13.64716 -13.71546 -13.79184 -13.90614 -14.03098 -14.18205 -14.35881 -14.48419 -14.60172 -14.70591 -14.83362 -14.92122 -15.00622 -15.05122 -15.03119 -14.99028 -14.92302 -14.86927 -14.82691 -14.7972 -14.76909 -14.71356 -14.61277 -14.51696 -14.42252 -14.36405 -14.30451 -14.23161 -14.19851 -14.16633 -14.15649 -14.10504 -13.99518 -13.79562 -13.3996 -12.7767 -11.71208 -8.311879 -8.600912 -9.615928 -10.43595 -11.21292 -11.88333 -12.36243 -12.63706 -12.8818 -12.83066 -12.89103 -12.95666 -13.19763 -13.40598 -13.49113 -13.5546 -13.55639 -13.51915 -13.68284 -13.53289 -13.42107 -13.65519 -13.50713 -13.75251 -13.76715 -13.87408 -13.73109 -13.70412 -13.56073 -13.53488 -13.54895 -13.56228 -13.59408 -13.62047 -13.64198 -13.66109 -13.62669 -13.58297 -13.57387 -13.4739 -13.53063 -13.48348 -13.61047 -13.64716 -13.71546 -13.79184 -13.90614 -14.03098 -14.18205 -14.35881 -14.48419 -14.60172 -14.70591 -14.83362 -14.92122 -15.00622 -15.05122 -15.03119 -14.99028 -14.92302 -14.86927 -14.82691 -14.7972 -14.76909 -14.71356 -14.61277 -14.51696 -14.42252 -14.36405 -14.30451 -14.23161 -14.19851 -14.16633 -14.15649 -14.10504 -13.99518 -13.79562 -13.3996 -12.7767 -11.71208 -8.311879 -8.600912 -9.615928 -10.43595 -11.21292 -11.88333 -12.36243 -12.63706 -12.8818 -12.83066 -12.89103 -12.95666 -13.19763 -13.40598 -13.49113 -13.5546 -13.55639 -13.51915 -13.68284 -13.53289 -13.42107 -13.65519 -13.50713 -13.75251 -13.76715 -13.87408 -13.73109 -13.70412 -13.56073 -13.53488 -13.54895 -13.56228 -13.59408 -13.62047 -13.64198 -13.66109 -13.62669 -13.58297 -13.57387 -13.4739 -13.53063 -13.48348 -13.61047 -13.64716 -13.71546 -13.79184 -13.90614 -14.03098 -14.18205 -14.35881 -14.48419 -14.60172 -14.70591 -14.83362 -14.92122 -15.00622 -15.05122 -15.03119 -14.99028 -14.92302 -14.86927 -14.82691 -14.7972 -14.76909 -14.71356 -14.61277 -14.51696 -14.42252 -14.36405 -14.30451 -14.23161 -14.19851 -14.16633 -14.15649 -14.10504 -13.99518 -13.79562 -13.3996 -12.7767 -11.71208 ]
|
| 6 |
+
<Rescale> 400 400
|
| 7 |
+
<LearnRateCoef> 0 [ 0.155775 0.154484 0.1527379 0.1518718 0.1506028 0.1489256 0.147067 0.1447061 0.1436307 0.1443568 0.1451849 0.1455157 0.1452821 0.1445717 0.1439195 0.1435867 0.1436018 0.1438781 0.1442086 0.1448844 0.1454756 0.145663 0.146268 0.1467386 0.1472724 0.147664 0.1480913 0.1483739 0.1488841 0.1493636 0.1497088 0.1500379 0.1502916 0.1505389 0.1506787 0.1507102 0.1505992 0.1505445 0.1505938 0.1508133 0.1509569 0.1512396 0.1514625 0.1516195 0.1516156 0.1515561 0.1514966 0.1513976 0.1512612 0.151076 0.1510596 0.1510431 0.151077 0.1511168 0.1511917 0.151023 0.1508045 0.1505885 0.1503493 0.1502373 0.1501726 0.1500762 0.1500065 0.1499782 0.150057 0.1502658 0.150469 0.1505335 0.1505505 0.1505328 0.1504275 0.1502438 0.1499674 0.1497118 0.1494661 0.1493102 0.1493681 0.1495501 0.1499738 0.1509654 0.155775 0.154484 0.1527379 0.1518718 0.1506028 0.1489256 0.147067 0.1447061 0.1436307 0.1443568 0.1451849 0.1455157 0.1452821 0.1445717 0.1439195 0.1435867 0.1436018 0.1438781 0.1442086 0.1448844 0.1454756 0.145663 0.146268 0.1467386 0.1472724 0.147664 0.1480913 0.1483739 0.1488841 0.1493636 0.1497088 0.1500379 0.1502916 0.1505389 0.1506787 0.1507102 0.1505992 0.1505445 0.1505938 0.1508133 0.1509569 0.1512396 0.1514625 0.1516195 0.1516156 0.1515561 0.1514966 0.1513976 0.1512612 0.151076 0.1510596 0.1510431 0.151077 0.1511168 0.1511917 0.151023 0.1508045 0.1505885 0.1503493 0.1502373 0.1501726 0.1500762 0.1500065 0.1499782 0.150057 0.1502658 0.150469 0.1505335 0.1505505 0.1505328 0.1504275 0.1502438 0.1499674 0.1497118 0.1494661 0.1493102 0.1493681 0.1495501 0.1499738 0.1509654 0.155775 0.154484 0.1527379 0.1518718 0.1506028 0.1489256 0.147067 0.1447061 0.1436307 0.1443568 0.1451849 0.1455157 0.1452821 0.1445717 0.1439195 0.1435867 0.1436018 0.1438781 0.1442086 0.1448844 0.1454756 0.145663 0.146268 0.1467386 0.1472724 0.147664 0.1480913 0.1483739 0.1488841 0.1493636 0.1497088 0.1500379 0.1502916 0.1505389 0.1506787 0.1507102 0.1505992 0.1505445 0.1505938 0.1508133 0.1509569 0.1512396 0.1514625 0.1516195 0.1516156 0.1515561 0.1514966 0.1513976 0.1512612 0.151076 0.1510596 0.1510431 0.151077 0.1511168 0.1511917 0.151023 0.1508045 0.1505885 0.1503493 0.1502373 0.1501726 0.1500762 0.1500065 0.1499782 0.150057 0.1502658 0.150469 0.1505335 0.1505505 0.1505328 0.1504275 0.1502438 0.1499674 0.1497118 0.1494661 0.1493102 0.1493681 0.1495501 0.1499738 0.1509654 0.155775 0.154484 0.1527379 0.1518718 0.1506028 0.1489256 0.147067 0.1447061 0.1436307 0.1443568 0.1451849 0.1455157 0.1452821 0.1445717 0.1439195 0.1435867 0.1436018 0.1438781 0.1442086 0.1448844 0.1454756 0.145663 0.146268 0.1467386 0.1472724 0.147664 0.1480913 0.1483739 0.1488841 0.1493636 0.1497088 0.1500379 0.1502916 0.1505389 0.1506787 0.1507102 0.1505992 0.1505445 0.1505938 0.1508133 0.1509569 0.1512396 0.1514625 0.1516195 0.1516156 0.1515561 0.1514966 0.1513976 0.1512612 0.151076 0.1510596 0.1510431 0.151077 0.1511168 0.1511917 0.151023 0.1508045 0.1505885 0.1503493 0.1502373 0.1501726 0.1500762 0.1500065 0.1499782 0.150057 0.1502658 0.150469 0.1505335 0.1505505 0.1505328 0.1504275 0.1502438 0.1499674 0.1497118 0.1494661 0.1493102 0.1493681 0.1495501 0.1499738 0.1509654 0.155775 0.154484 0.1527379 0.1518718 0.1506028 0.1489256 0.147067 0.1447061 0.1436307 0.1443568 0.1451849 0.1455157 0.1452821 0.1445717 0.1439195 0.1435867 0.1436018 0.1438781 0.1442086 0.1448844 0.1454756 0.145663 0.146268 0.1467386 0.1472724 0.147664 0.1480913 0.1483739 0.1488841 0.1493636 0.1497088 0.1500379 0.1502916 0.1505389 0.1506787 0.1507102 0.1505992 0.1505445 0.1505938 0.1508133 0.1509569 0.1512396 0.1514625 0.1516195 0.1516156 0.1515561 0.1514966 0.1513976 0.1512612 0.151076 0.1510596 0.1510431 0.151077 0.1511168 0.1511917 0.151023 0.1508045 0.1505885 0.1503493 0.1502373 0.1501726 0.1500762 0.1500065 0.1499782 0.150057 0.1502658 0.150469 0.1505335 0.1505505 0.1505328 0.1504275 0.1502438 0.1499674 0.1497118 0.1494661 0.1493102 0.1493681 0.1495501 0.1499738 0.1509654 ]
|
| 8 |
+
</Nnet>
|
ax_meeting/ax_model/vad/config.yaml
ADDED
|
@@ -0,0 +1,56 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
frontend: WavFrontendOnline
|
| 2 |
+
frontend_conf:
|
| 3 |
+
fs: 16000
|
| 4 |
+
window: hamming
|
| 5 |
+
n_mels: 80
|
| 6 |
+
frame_length: 25
|
| 7 |
+
frame_shift: 10
|
| 8 |
+
dither: 0.0
|
| 9 |
+
lfr_m: 5
|
| 10 |
+
lfr_n: 1
|
| 11 |
+
|
| 12 |
+
model: FsmnVADStreaming
|
| 13 |
+
model_conf:
|
| 14 |
+
sample_rate: 16000
|
| 15 |
+
detect_mode: 1
|
| 16 |
+
snr_mode: 0
|
| 17 |
+
max_end_silence_time: 800
|
| 18 |
+
max_start_silence_time: 3000
|
| 19 |
+
do_start_point_detection: True
|
| 20 |
+
do_end_point_detection: True
|
| 21 |
+
window_size_ms: 200
|
| 22 |
+
sil_to_speech_time_thres: 150
|
| 23 |
+
speech_to_sil_time_thres: 150
|
| 24 |
+
speech_2_noise_ratio: 1.0
|
| 25 |
+
do_extend: 1
|
| 26 |
+
lookback_time_start_point: 200
|
| 27 |
+
lookahead_time_end_point: 100
|
| 28 |
+
max_single_segment_time: 60000
|
| 29 |
+
snr_thres: -100.0
|
| 30 |
+
noise_frame_num_used_for_snr: 100
|
| 31 |
+
decibel_thres: -100.0
|
| 32 |
+
speech_noise_thres: 0.6
|
| 33 |
+
fe_prior_thres: 0.0001
|
| 34 |
+
silence_pdf_num: 1
|
| 35 |
+
sil_pdf_ids: [0]
|
| 36 |
+
speech_noise_thresh_low: -0.1
|
| 37 |
+
speech_noise_thresh_high: 0.3
|
| 38 |
+
output_frame_probs: False
|
| 39 |
+
frame_in_ms: 10
|
| 40 |
+
frame_length_ms: 25
|
| 41 |
+
|
| 42 |
+
encoder: FSMN
|
| 43 |
+
encoder_conf:
|
| 44 |
+
input_dim: 400
|
| 45 |
+
input_affine_dim: 140
|
| 46 |
+
fsmn_layers: 4
|
| 47 |
+
linear_dim: 250
|
| 48 |
+
proj_dim: 128
|
| 49 |
+
lorder: 20
|
| 50 |
+
rorder: 0
|
| 51 |
+
lstride: 1
|
| 52 |
+
rstride: 0
|
| 53 |
+
output_affine_dim: 140
|
| 54 |
+
output_dim: 248
|
| 55 |
+
|
| 56 |
+
|
ax_meeting/ax_model/withitn.npy
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:39bf02586f59237894fc2918ab2db4f12ec3c084c41465718832fbd7646ea729
|
| 3 |
+
size 2368
|
ax_meeting/ax_model/yue.npy
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:319c3d30e6b176338a6d8871f81ec57af5c3c112cecbdce4a0bf36bec51f3ceb
|
| 3 |
+
size 2368
|
ax_meeting/ax_model/zh.npy
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:0c952f976ae70642c8f682c88711abce3a2b64ba20a6d8ca0b706f227de816eb
|
| 3 |
+
size 2368
|
ax_meeting/axengine_loader.py
ADDED
|
@@ -0,0 +1,34 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# -*- coding: utf-8 -*-
|
| 2 |
+
import os
|
| 3 |
+
import glob
|
| 4 |
+
import subprocess
|
| 5 |
+
import sys
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
def ensure_axengine():
|
| 9 |
+
try:
|
| 10 |
+
import axengine # noqa: F401
|
| 11 |
+
return
|
| 12 |
+
except Exception:
|
| 13 |
+
pass
|
| 14 |
+
|
| 15 |
+
wheel_path = os.getenv("AXENGINE_WHEEL")
|
| 16 |
+
if wheel_path and os.path.exists(wheel_path):
|
| 17 |
+
_install_wheel(wheel_path)
|
| 18 |
+
return
|
| 19 |
+
|
| 20 |
+
# Try bundled wheel
|
| 21 |
+
here = os.path.dirname(__file__)
|
| 22 |
+
candidates = glob.glob(os.path.join(here, "vendor", "pyaxengine*.whl"))
|
| 23 |
+
candidates += glob.glob(os.path.join(here, "vendor", "axengine*.whl"))
|
| 24 |
+
if candidates:
|
| 25 |
+
_install_wheel(candidates[0])
|
| 26 |
+
return
|
| 27 |
+
|
| 28 |
+
raise RuntimeError(
|
| 29 |
+
"axengine not found. Please install pyaxengine wheel or set AXENGINE_WHEEL."
|
| 30 |
+
)
|
| 31 |
+
|
| 32 |
+
|
| 33 |
+
def _install_wheel(path: str):
|
| 34 |
+
subprocess.check_call([sys.executable, "-m", "pip", "install", "--no-deps", path])
|
ax_meeting/certs/cert.pem
ADDED
|
@@ -0,0 +1,19 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
-----BEGIN CERTIFICATE-----
|
| 2 |
+
MIIDETCCAfmgAwIBAgIUPI3rAC24IoJ7FUydYijEOp7GeygwDQYJKoZIhvcNAQEL
|
| 3 |
+
BQAwGDEWMBQGA1UEAwwNMTAuMTI2LjMzLjE0MDAeFw0yNjAyMjcwMzE3MzRaFw0y
|
| 4 |
+
NzAyMjcwMzE3MzRaMBgxFjAUBgNVBAMMDTEwLjEyNi4zMy4xNDAwggEiMA0GCSqG
|
| 5 |
+
SIb3DQEBAQUAA4IBDwAwggEKAoIBAQChzhKivPsGUSjfzFWwocMg1FT56iDFo8yy
|
| 6 |
+
tba/LvbP2i2BpujTqK5/6iInLh6N9ZptJg4PsLSEQ2HWfLdYupEYvMrXDy4nYwsX
|
| 7 |
+
gAbrjAz3uhsJV2+LSVF/0g8PNhwidDN4WNWLQoVf5g9FCxl0SneCoyKpQdpIB10r
|
| 8 |
+
J0ZXtKe5SY9Ydq0EdjS+5898U83XgOIFQfKDRdakPuxKLX00DFd6S5Xz4Yw148As
|
| 9 |
+
ypAOfNSXCqX0+2wtSyfAednwlPea+VxQPExQBpx7yOYe0eDrMDDPP7z6UUdzMdoD
|
| 10 |
+
ZOlWTgtIarBCthD9HnhZhqK6iom9sxm9hvHOV+kM/pje1iMeaFbpAgMBAAGjUzBR
|
| 11 |
+
MB0GA1UdDgQWBBRhENjhWth7MneI5GxGFzCu2RjUbTAfBgNVHSMEGDAWgBRhENjh
|
| 12 |
+
Wth7MneI5GxGFzCu2RjUbTAPBgNVHRMBAf8EBTADAQH/MA0GCSqGSIb3DQEBCwUA
|
| 13 |
+
A4IBAQAXi8HfQF1LyACZSkilE6emnB0zh31NMR7oimhc9rlKgX5yPD4RQp+OKr8l
|
| 14 |
+
APeIFeWLiHadvxKpjnT+MQwzODs5i+Qai0vDaRoy/OLPiaI3rH8yaAMxMI0hrh4c
|
| 15 |
+
Rfyf6bWBvcTpuUTFO/hC1KxfVVCjCwMcV/Jmv0xukCMCXCHNqDdkG2m24RSmKUHl
|
| 16 |
+
Af85bduX8PabtZxEDESkLc2Z8G8w5VvEWJizQ7OF0EcYJIA3JEDcr50o5kFpu7OQ
|
| 17 |
+
ZkRfWopQ3I1DOUOfSrc4d5gMy3b1rvDBbeLTRsNnec9WciwP4U/+Hbxs2P+BALo+
|
| 18 |
+
WSq6D7FC0WTTXtYdGFv2oY5s38Dr
|
| 19 |
+
-----END CERTIFICATE-----
|
ax_meeting/certs/key.pem
ADDED
|
@@ -0,0 +1,28 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
-----BEGIN PRIVATE KEY-----
|
| 2 |
+
MIIEvQIBADANBgkqhkiG9w0BAQEFAASCBKcwggSjAgEAAoIBAQChzhKivPsGUSjf
|
| 3 |
+
zFWwocMg1FT56iDFo8yytba/LvbP2i2BpujTqK5/6iInLh6N9ZptJg4PsLSEQ2HW
|
| 4 |
+
fLdYupEYvMrXDy4nYwsXgAbrjAz3uhsJV2+LSVF/0g8PNhwidDN4WNWLQoVf5g9F
|
| 5 |
+
Cxl0SneCoyKpQdpIB10rJ0ZXtKe5SY9Ydq0EdjS+5898U83XgOIFQfKDRdakPuxK
|
| 6 |
+
LX00DFd6S5Xz4Yw148AsypAOfNSXCqX0+2wtSyfAednwlPea+VxQPExQBpx7yOYe
|
| 7 |
+
0eDrMDDPP7z6UUdzMdoDZOlWTgtIarBCthD9HnhZhqK6iom9sxm9hvHOV+kM/pje
|
| 8 |
+
1iMeaFbpAgMBAAECggEAQjovPYX9dy3z/XpM3pGvZPoT2AEBLfQn/kPLS4CFDDlg
|
| 9 |
+
o+814CBsYDXsib3iSreq4B8R5VEt6e8MljaQ8xPV/NqVaaYwfXWYHiPMcU/vJNx7
|
| 10 |
+
YXz0zn2RirBncpHyvRVz1cACk9AD+GcZe+iZoBQ0y3dLYhzuo8nD1Dxsmcx7VCaF
|
| 11 |
+
VTAVKGsCSW3ZHfXXxMDkB/3tiRknb4KNpvVzLo2GNgj5fiygkkgf9hM1hRfqt7/P
|
| 12 |
+
nejF3A0laWQc0N0OluM6SF+2A4uNnWQwIeyLK/AdM14Me5rvYjiOG/xE1hNVBvup
|
| 13 |
+
utkjVyVYP1pCoK0eCqaApd61DuDcagFPi6gSmKeY4wKBgQDXF0cGE6tbohNUa/GU
|
| 14 |
+
r8jCb51F1gWuLJNl4y6SMDr9/iku1OKkJtbK0InMmfeyeKeg/xeuXSc5QJY5hcrd
|
| 15 |
+
A/Y4mTRWlpfzpYCj5qos7f1HVl/mtwBqN2VUVi2Tid1OrfDgTup0ZZlN1IjszH3/
|
| 16 |
+
37IQnJTs/sfVxwwh+ij+VKvP6wKBgQDAlFDTq8lRtbc/nVAqLPnnc2HnLtrAzeLg
|
| 17 |
+
N1qD7QQvcfWP9gxSg0ZBYJf52uYRjLDB4b6LBPV5Ys9SIcXz8G6VYwUESkJcnQ+s
|
| 18 |
+
yTNUuJolzAFq0NoMLAAN0vKV2yyIdGP/oupTYEfSRz1bRtIGchhY/XMdVjj5Ppxc
|
| 19 |
+
w3CfAAcTewKBgCI4nuE1oe7bU437+py4dw2Qaopg6dhzWSQ9x/wUVl5w4KaF0mVh
|
| 20 |
+
lI0CLtpxqLopfiocS+0+/u2Z/Ay837DYX4VTwsMABL8MFvJ80ZiCaOi/slRny1Ya
|
| 21 |
+
6DFJ4Mh3h9Fr1UYq6ByKyaBbb0mVo3phYdhIwV0PkEXP/HsvbPRCDm/vAoGBALng
|
| 22 |
+
bgNgk/giBLWKCY4ryyny3FRfjRT7pDf2NY+QfbGttO829b3Op0kDCq1G8zmNKi54
|
| 23 |
+
zYkxSB3ZmXIU1xQUxSe7Y2Q4qMTrc+26ZakoZOCGf/exjkShU4wER9EMs3choENl
|
| 24 |
+
4/aFv8zepgIr4RwHlCiQuUNfra4lGJcQrOtLA4lxAoGAWn8ovOGTEo5Seuq5CoCQ
|
| 25 |
+
HbDdMmd6z9kLInKcqaGsuclFaEZjG6CHznD7B+SRtJFBN3+L0K3UbwNeAka6yEqg
|
| 26 |
+
3TUXNBXMypOKYPYDAbGE8t7wrLL7WEtgHiqpkN/zZoxX61f3nLOIoku817dW3KtW
|
| 27 |
+
PTUCWpm7u5IQWl6T+ThOJmU=
|
| 28 |
+
-----END PRIVATE KEY-----
|
ax_meeting/config.py
ADDED
|
@@ -0,0 +1,19 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# -*- coding: utf-8 -*-
|
| 2 |
+
|
| 3 |
+
SAMPLE_RATE = 16000
|
| 4 |
+
|
| 5 |
+
# How often to run VAD checks on the rolling audio buffer
|
| 6 |
+
VAD_CHECK_INTERVAL_SEC = 0.8
|
| 7 |
+
|
| 8 |
+
# If end time is earlier than (now - PAUSE_MS), treat as a completed segment
|
| 9 |
+
PAUSE_MS = 800
|
| 10 |
+
|
| 11 |
+
# Ignore too-short segments
|
| 12 |
+
MIN_SEGMENT_MS = 300
|
| 13 |
+
|
| 14 |
+
# Merge VAD segments shorter than this in offline diarization
|
| 15 |
+
MERGE_VAD_MAX_LEN_MS = 15 * 1000
|
| 16 |
+
|
| 17 |
+
# LLM summarization chunking
|
| 18 |
+
SUMMARY_CHUNK_CHARS = 1000
|
| 19 |
+
SUMMARY_TARGET_CHARS = 100
|
ax_meeting/diar_asr_cli.py
ADDED
|
@@ -0,0 +1,93 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# -*- coding: utf-8 -*-
|
| 2 |
+
import argparse
|
| 3 |
+
from pathlib import Path
|
| 4 |
+
|
| 5 |
+
import numpy as np
|
| 6 |
+
import soundfile as sf
|
| 7 |
+
|
| 8 |
+
from ax_meeting.model_bundle import ModelBundle
|
| 9 |
+
from ax_meeting.utils.vad_utils import merge_vad
|
| 10 |
+
from ax_meeting.utils.ax_cam_bin import do_clustering, chunk
|
| 11 |
+
from ax_meeting.diar_utils import pick_speaker
|
| 12 |
+
from ax_meeting.config import MERGE_VAD_MAX_LEN_MS
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
def load_audio(path: str, target_sr: int = 16000) -> np.ndarray:
|
| 16 |
+
audio, sr = sf.read(path, dtype="float32")
|
| 17 |
+
if audio.ndim > 1:
|
| 18 |
+
audio = audio.mean(axis=1)
|
| 19 |
+
if sr != target_sr:
|
| 20 |
+
ratio = target_sr / float(sr)
|
| 21 |
+
new_len = int(len(audio) * ratio)
|
| 22 |
+
if new_len <= 0:
|
| 23 |
+
return np.zeros(0, dtype=np.float32)
|
| 24 |
+
x_old = np.linspace(0, 1, num=len(audio), endpoint=False)
|
| 25 |
+
x_new = np.linspace(0, 1, num=new_len, endpoint=False)
|
| 26 |
+
audio = np.interp(x_new, x_old, audio).astype(np.float32)
|
| 27 |
+
return audio
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
def diar_asr(bundle: ModelBundle, speech: np.ndarray, fs: int = 16000) -> str:
|
| 31 |
+
if speech.size == 0:
|
| 32 |
+
return ""
|
| 33 |
+
|
| 34 |
+
res_vad = bundle.vad_infer(speech)
|
| 35 |
+
if not res_vad:
|
| 36 |
+
return ""
|
| 37 |
+
|
| 38 |
+
vad_segments = merge_vad(res_vad, MERGE_VAD_MAX_LEN_MS)
|
| 39 |
+
vad_time = [[vad_t[0] / 1000, vad_t[1] / 1000] for vad_t in res_vad]
|
| 40 |
+
chunks = [c for (st, ed) in vad_time for c in chunk(st, ed)]
|
| 41 |
+
|
| 42 |
+
if not chunks:
|
| 43 |
+
return ""
|
| 44 |
+
|
| 45 |
+
embeddings = bundle.speaker_infer(speech, fs, chunks=chunks)
|
| 46 |
+
_, diar_results = do_clustering(chunks, embeddings, speaker_num=None)
|
| 47 |
+
|
| 48 |
+
lines = []
|
| 49 |
+
for i, segment in enumerate(vad_segments):
|
| 50 |
+
segment_start, segment_end = segment
|
| 51 |
+
start_sample = int(segment_start / 1000 * fs)
|
| 52 |
+
end_sample = min(int(segment_end / 1000 * fs), speech.shape[0])
|
| 53 |
+
segment_speech = speech[start_sample:end_sample]
|
| 54 |
+
|
| 55 |
+
text, _ = bundle.asr_infer(
|
| 56 |
+
segment_speech,
|
| 57 |
+
output_timestamp=False,
|
| 58 |
+
key=f"segment_{i}",
|
| 59 |
+
)
|
| 60 |
+
if not text or not text.strip():
|
| 61 |
+
continue
|
| 62 |
+
|
| 63 |
+
spk = pick_speaker(segment_start / 1000.0, segment_end / 1000.0, diar_results)
|
| 64 |
+
lines.append(
|
| 65 |
+
f"Speaker_{spk}: [{segment_start/1000.0:.3f} {segment_end/1000.0:.3f}] {text.strip()}"
|
| 66 |
+
)
|
| 67 |
+
return "\n".join(lines)
|
| 68 |
+
|
| 69 |
+
|
| 70 |
+
def main():
|
| 71 |
+
parser = argparse.ArgumentParser(description="Offline diarization + ASR")
|
| 72 |
+
parser.add_argument("--wav_file", required=True, help="Input audio file")
|
| 73 |
+
parser.add_argument("--output_dir", default="output_dir", help="Output directory")
|
| 74 |
+
args = parser.parse_args()
|
| 75 |
+
|
| 76 |
+
out_dir = Path(args.output_dir)
|
| 77 |
+
out_dir.mkdir(parents=True, exist_ok=True)
|
| 78 |
+
|
| 79 |
+
speech = load_audio(args.wav_file, target_sr=16000)
|
| 80 |
+
|
| 81 |
+
bundle = ModelBundle()
|
| 82 |
+
bundle.ensure_loaded()
|
| 83 |
+
|
| 84 |
+
transcript = diar_asr(bundle, speech, fs=16000)
|
| 85 |
+
transcript_path = out_dir / (Path(args.wav_file).stem + ".txt")
|
| 86 |
+
transcript_path.write_text(transcript, encoding="utf-8")
|
| 87 |
+
print(f"Transcript saved: {transcript_path}")
|
| 88 |
+
|
| 89 |
+
# Summary is handled by ax_meeting.summarize_cli
|
| 90 |
+
|
| 91 |
+
|
| 92 |
+
if __name__ == "__main__":
|
| 93 |
+
main()
|
ax_meeting/diar_utils.py
ADDED
|
@@ -0,0 +1,13 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# -*- coding: utf-8 -*-
|
| 2 |
+
|
| 3 |
+
def pick_speaker(seg_st: float, seg_ed: float, diar_results) -> int:
|
| 4 |
+
if not diar_results:
|
| 5 |
+
return 0
|
| 6 |
+
best_spk = diar_results[0][2]
|
| 7 |
+
best_overlap = 0.0
|
| 8 |
+
for st_spk, ed_spk, spk in diar_results:
|
| 9 |
+
overlap = min(seg_ed, ed_spk) - max(seg_st, st_spk)
|
| 10 |
+
if overlap > best_overlap:
|
| 11 |
+
best_overlap = overlap
|
| 12 |
+
best_spk = spk
|
| 13 |
+
return int(best_spk)
|
ax_meeting/engines.py
ADDED
|
@@ -0,0 +1,207 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# -*- coding: utf-8 -*-
|
| 2 |
+
import os
|
| 3 |
+
import shutil
|
| 4 |
+
import subprocess
|
| 5 |
+
from dataclasses import dataclass
|
| 6 |
+
from typing import List, Iterable, Optional, Union, Tuple
|
| 7 |
+
|
| 8 |
+
import numpy as np
|
| 9 |
+
import soundfile as sf
|
| 10 |
+
|
| 11 |
+
from ax_meeting.model_bundle import ModelBundle
|
| 12 |
+
from ax_meeting.config import SAMPLE_RATE, PAUSE_MS, MIN_SEGMENT_MS, MERGE_VAD_MAX_LEN_MS
|
| 13 |
+
from ax_meeting.diar_utils import pick_speaker
|
| 14 |
+
from ax_meeting.utils.vad_utils import merge_vad
|
| 15 |
+
from ax_meeting.utils.ax_cam_bin import do_clustering, chunk
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
@dataclass
|
| 19 |
+
class VadAsrSegment:
|
| 20 |
+
start_ms: int
|
| 21 |
+
end_ms: int
|
| 22 |
+
text: str
|
| 23 |
+
audio: np.ndarray
|
| 24 |
+
|
| 25 |
+
|
| 26 |
+
def _resample(audio: np.ndarray, src_sr: int, tgt_sr: int) -> np.ndarray:
|
| 27 |
+
if src_sr == tgt_sr:
|
| 28 |
+
return audio
|
| 29 |
+
ratio = tgt_sr / float(src_sr)
|
| 30 |
+
new_len = int(len(audio) * ratio)
|
| 31 |
+
if new_len <= 0:
|
| 32 |
+
return np.zeros(0, dtype=np.float32)
|
| 33 |
+
x_old = np.linspace(0, 1, num=len(audio), endpoint=False)
|
| 34 |
+
x_new = np.linspace(0, 1, num=new_len, endpoint=False)
|
| 35 |
+
return np.interp(x_new, x_old, audio).astype(np.float32)
|
| 36 |
+
|
| 37 |
+
|
| 38 |
+
def _load_audio_file(path: str, target_sr: int = SAMPLE_RATE) -> np.ndarray:
|
| 39 |
+
try:
|
| 40 |
+
audio, sr = sf.read(path, dtype="float32")
|
| 41 |
+
if audio.ndim > 1:
|
| 42 |
+
audio = audio.mean(axis=1)
|
| 43 |
+
return _resample(audio, sr, target_sr)
|
| 44 |
+
except Exception:
|
| 45 |
+
# Try ffmpeg for formats like mp4/m4a
|
| 46 |
+
if shutil.which("ffmpeg") is None:
|
| 47 |
+
raise
|
| 48 |
+
cmd = [
|
| 49 |
+
"ffmpeg",
|
| 50 |
+
"-i",
|
| 51 |
+
path,
|
| 52 |
+
"-f",
|
| 53 |
+
"f32le",
|
| 54 |
+
"-ac",
|
| 55 |
+
"1",
|
| 56 |
+
"-ar",
|
| 57 |
+
str(target_sr),
|
| 58 |
+
"-",
|
| 59 |
+
]
|
| 60 |
+
proc = subprocess.run(cmd, stdout=subprocess.PIPE, stderr=subprocess.DEVNULL, check=True)
|
| 61 |
+
return np.frombuffer(proc.stdout, dtype=np.float32)
|
| 62 |
+
|
| 63 |
+
|
| 64 |
+
def _normalize_inputs(inputs: Union[np.ndarray, bytes, str, Iterable]) -> List[np.ndarray]:
|
| 65 |
+
if isinstance(inputs, np.ndarray):
|
| 66 |
+
return [inputs.astype(np.float32)]
|
| 67 |
+
if isinstance(inputs, (bytes, bytearray)):
|
| 68 |
+
audio = np.frombuffer(inputs, dtype=np.int16).astype(np.float32) / 32768.0
|
| 69 |
+
return [audio]
|
| 70 |
+
if isinstance(inputs, str):
|
| 71 |
+
return [_load_audio_file(inputs)]
|
| 72 |
+
if isinstance(inputs, Iterable):
|
| 73 |
+
out = []
|
| 74 |
+
for item in inputs:
|
| 75 |
+
out.extend(_normalize_inputs(item))
|
| 76 |
+
return out
|
| 77 |
+
raise TypeError(f"Unsupported input type: {type(inputs)}")
|
| 78 |
+
|
| 79 |
+
|
| 80 |
+
class VadAsrEngine:
|
| 81 |
+
def __init__(
|
| 82 |
+
self,
|
| 83 |
+
model_bundle: Optional[ModelBundle] = None,
|
| 84 |
+
stream: bool = False,
|
| 85 |
+
sample_rate: int = SAMPLE_RATE,
|
| 86 |
+
pause_ms: int = PAUSE_MS,
|
| 87 |
+
min_segment_ms: int = MIN_SEGMENT_MS,
|
| 88 |
+
):
|
| 89 |
+
self.bundle = model_bundle or ModelBundle()
|
| 90 |
+
self.bundle.ensure_loaded()
|
| 91 |
+
self.stream = stream
|
| 92 |
+
self.sample_rate = sample_rate
|
| 93 |
+
self.pause_ms = pause_ms
|
| 94 |
+
self.min_segment_ms = min_segment_ms
|
| 95 |
+
self.audio_chunks: List[np.ndarray] = []
|
| 96 |
+
self.total_samples: int = 0
|
| 97 |
+
self.last_processed_ms: int = 0
|
| 98 |
+
|
| 99 |
+
def feed(self, audio: Union[np.ndarray, bytes, str, Iterable]) -> None:
|
| 100 |
+
chunks = _normalize_inputs(audio)
|
| 101 |
+
for ch in chunks:
|
| 102 |
+
if ch.size == 0:
|
| 103 |
+
continue
|
| 104 |
+
self.audio_chunks.append(ch)
|
| 105 |
+
self.total_samples += ch.size
|
| 106 |
+
|
| 107 |
+
def _audio_all(self) -> np.ndarray:
|
| 108 |
+
if not self.audio_chunks:
|
| 109 |
+
return np.zeros(0, dtype=np.float32)
|
| 110 |
+
return np.concatenate(self.audio_chunks, axis=0)
|
| 111 |
+
|
| 112 |
+
def _current_ms(self) -> int:
|
| 113 |
+
return int(self.total_samples / self.sample_rate * 1000)
|
| 114 |
+
|
| 115 |
+
def poll(self) -> List[VadAsrSegment]:
|
| 116 |
+
if not self.stream:
|
| 117 |
+
raise RuntimeError("poll() only available in stream mode")
|
| 118 |
+
audio = self._audio_all()
|
| 119 |
+
if audio.size == 0:
|
| 120 |
+
return []
|
| 121 |
+
current_ms = self._current_ms()
|
| 122 |
+
vad_segments = self.bundle.vad_infer(audio)
|
| 123 |
+
out: List[VadAsrSegment] = []
|
| 124 |
+
for start_ms, end_ms in vad_segments:
|
| 125 |
+
if end_ms <= self.last_processed_ms:
|
| 126 |
+
continue
|
| 127 |
+
if (end_ms - start_ms) < self.min_segment_ms:
|
| 128 |
+
continue
|
| 129 |
+
if end_ms > current_ms - self.pause_ms:
|
| 130 |
+
continue
|
| 131 |
+
start_sample = int(start_ms / 1000 * self.sample_rate)
|
| 132 |
+
end_sample = int(end_ms / 1000 * self.sample_rate)
|
| 133 |
+
seg_audio = audio[start_sample:end_sample]
|
| 134 |
+
text, _ = self.bundle.asr_infer(seg_audio, output_timestamp=False, key="stream")
|
| 135 |
+
if text.strip():
|
| 136 |
+
out.append(VadAsrSegment(start_ms, end_ms, text.strip(), seg_audio))
|
| 137 |
+
self.last_processed_ms = max(self.last_processed_ms, int(end_ms))
|
| 138 |
+
return out
|
| 139 |
+
|
| 140 |
+
def transcribe(self, audio: Union[np.ndarray, bytes, str, Iterable]) -> List[VadAsrSegment]:
|
| 141 |
+
chunks = _normalize_inputs(audio)
|
| 142 |
+
if not chunks:
|
| 143 |
+
return []
|
| 144 |
+
speech = np.concatenate(chunks, axis=0)
|
| 145 |
+
vad_segments = self.bundle.vad_infer(speech)
|
| 146 |
+
vad_segments = merge_vad(vad_segments, MERGE_VAD_MAX_LEN_MS)
|
| 147 |
+
out: List[VadAsrSegment] = []
|
| 148 |
+
for start_ms, end_ms in vad_segments:
|
| 149 |
+
if (end_ms - start_ms) < self.min_segment_ms:
|
| 150 |
+
continue
|
| 151 |
+
start_sample = int(start_ms / 1000 * self.sample_rate)
|
| 152 |
+
end_sample = int(end_ms / 1000 * self.sample_rate)
|
| 153 |
+
seg_audio = speech[start_sample:end_sample]
|
| 154 |
+
text, _ = self.bundle.asr_infer(seg_audio, output_timestamp=False, key="offline")
|
| 155 |
+
if text.strip():
|
| 156 |
+
out.append(VadAsrSegment(start_ms, end_ms, text.strip(), seg_audio))
|
| 157 |
+
return out
|
| 158 |
+
|
| 159 |
+
|
| 160 |
+
class DiarAsrEngine:
|
| 161 |
+
def __init__(self, model_bundle: Optional[ModelBundle] = None, sample_rate: int = SAMPLE_RATE):
|
| 162 |
+
self.bundle = model_bundle or ModelBundle()
|
| 163 |
+
self.bundle.ensure_loaded()
|
| 164 |
+
self.sample_rate = sample_rate
|
| 165 |
+
self.audio_chunks: List[np.ndarray] = []
|
| 166 |
+
|
| 167 |
+
def feed(self, audio: Union[np.ndarray, bytes, str, Iterable]) -> None:
|
| 168 |
+
chunks = _normalize_inputs(audio)
|
| 169 |
+
for ch in chunks:
|
| 170 |
+
if ch.size == 0:
|
| 171 |
+
continue
|
| 172 |
+
self.audio_chunks.append(ch)
|
| 173 |
+
|
| 174 |
+
def finalize(self) -> str:
|
| 175 |
+
if not self.audio_chunks:
|
| 176 |
+
return ""
|
| 177 |
+
speech = np.concatenate(self.audio_chunks, axis=0)
|
| 178 |
+
return self.transcribe(speech)
|
| 179 |
+
|
| 180 |
+
def transcribe(self, audio: Union[np.ndarray, bytes, str, Iterable]) -> str:
|
| 181 |
+
chunks = _normalize_inputs(audio)
|
| 182 |
+
if not chunks:
|
| 183 |
+
return ""
|
| 184 |
+
speech = np.concatenate(chunks, axis=0)
|
| 185 |
+
res_vad = self.bundle.vad_infer(speech)
|
| 186 |
+
vad_segments = merge_vad(res_vad, MERGE_VAD_MAX_LEN_MS)
|
| 187 |
+
vad_time = [[vad_t[0] / 1000, vad_t[1] / 1000] for vad_t in res_vad]
|
| 188 |
+
chunks = [c for (st, ed) in vad_time for c in chunk(st, ed)]
|
| 189 |
+
if not chunks:
|
| 190 |
+
return ""
|
| 191 |
+
embeddings = self.bundle.speaker_infer(speech, self.sample_rate, chunks=chunks)
|
| 192 |
+
_, diar_results = do_clustering(chunks, embeddings, speaker_num=None)
|
| 193 |
+
|
| 194 |
+
lines = []
|
| 195 |
+
for i, segment in enumerate(vad_segments):
|
| 196 |
+
segment_start, segment_end = segment
|
| 197 |
+
start_sample = int(segment_start / 1000 * self.sample_rate)
|
| 198 |
+
end_sample = min(int(segment_end / 1000 * self.sample_rate), speech.shape[0])
|
| 199 |
+
segment_speech = speech[start_sample:end_sample]
|
| 200 |
+
text, _ = self.bundle.asr_infer(segment_speech, output_timestamp=False, key=f"segment_{i}")
|
| 201 |
+
if not text or not text.strip():
|
| 202 |
+
continue
|
| 203 |
+
spk = pick_speaker(segment_start / 1000.0, segment_end / 1000.0, diar_results)
|
| 204 |
+
lines.append(
|
| 205 |
+
f"Speaker_{spk}: [{segment_start/1000.0:.3f} {segment_end/1000.0:.3f}] {text.strip()}"
|
| 206 |
+
)
|
| 207 |
+
return "\n".join(lines)
|
ax_meeting/model_bundle.py
ADDED
|
@@ -0,0 +1,87 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# -*- coding: utf-8 -*-
|
| 2 |
+
import os
|
| 3 |
+
from pathlib import Path
|
| 4 |
+
import threading
|
| 5 |
+
import numpy as np
|
| 6 |
+
|
| 7 |
+
from ax_meeting.positional import SinusoidalPositionEncoder
|
| 8 |
+
from ax_meeting.utils.ax_model_bin import AX_SenseVoiceSmall
|
| 9 |
+
from ax_meeting.utils.ax_vad_bin import AX_Fsmn_vad
|
| 10 |
+
from ax_meeting.utils.ax_cam_bin import AX_SpeakerEmbeddingInference
|
| 11 |
+
from ax_meeting.utils.sentencepiece_tokenizer import SentencepiecesTokenizer
|
| 12 |
+
from ax_meeting.text_cleaner import clean_asr_text
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
class ModelBundle:
|
| 16 |
+
def __init__(self, ax_model_dir: str | None = None, seq_len: int = 132):
|
| 17 |
+
if ax_model_dir is None:
|
| 18 |
+
env_dir = os.getenv("AX_MODEL_DIR")
|
| 19 |
+
if env_dir:
|
| 20 |
+
ax_model_dir = env_dir
|
| 21 |
+
else:
|
| 22 |
+
ax_model_dir = str(Path(__file__).resolve().parent / "ax_model")
|
| 23 |
+
self.ax_model_dir = ax_model_dir
|
| 24 |
+
self.seq_len = seq_len
|
| 25 |
+
|
| 26 |
+
self._lock = threading.Lock()
|
| 27 |
+
self._loaded = False
|
| 28 |
+
self._load_error = None
|
| 29 |
+
|
| 30 |
+
self.vad = None
|
| 31 |
+
self.speaker = None
|
| 32 |
+
self.position_encoding = None
|
| 33 |
+
self.asr = None
|
| 34 |
+
self.tokenizer = None
|
| 35 |
+
self.language = "zh"
|
| 36 |
+
self.withitn = True
|
| 37 |
+
|
| 38 |
+
def load(self):
|
| 39 |
+
if self._loaded:
|
| 40 |
+
return
|
| 41 |
+
try:
|
| 42 |
+
self.vad = AX_Fsmn_vad(self.ax_model_dir)
|
| 43 |
+
self.speaker = AX_SpeakerEmbeddingInference(model_dir=self.ax_model_dir)
|
| 44 |
+
|
| 45 |
+
embed = SinusoidalPositionEncoder()
|
| 46 |
+
dummy = np.zeros((1, self.seq_len, 560), dtype=np.float32)
|
| 47 |
+
self.position_encoding = embed.get_position_encoding(dummy)
|
| 48 |
+
self.asr = AX_SenseVoiceSmall(self.ax_model_dir, seq_len=self.seq_len)
|
| 49 |
+
|
| 50 |
+
tokenizer_path = os.path.join(self.ax_model_dir, "chn_jpn_yue_eng_ko_spectok.bpe.model")
|
| 51 |
+
self.tokenizer = SentencepiecesTokenizer(bpemodel=tokenizer_path)
|
| 52 |
+
|
| 53 |
+
self._loaded = True
|
| 54 |
+
except Exception as e:
|
| 55 |
+
self._load_error = e
|
| 56 |
+
raise
|
| 57 |
+
|
| 58 |
+
def ensure_loaded(self):
|
| 59 |
+
if not self._loaded:
|
| 60 |
+
self.load()
|
| 61 |
+
|
| 62 |
+
def vad_infer(self, audio: np.ndarray):
|
| 63 |
+
self.ensure_loaded()
|
| 64 |
+
with self._lock:
|
| 65 |
+
return self.vad(audio)[0]
|
| 66 |
+
|
| 67 |
+
def speaker_infer(self, audio: np.ndarray, fs: int, chunks):
|
| 68 |
+
self.ensure_loaded()
|
| 69 |
+
with self._lock:
|
| 70 |
+
return self.speaker(audio, fs, chunks=chunks)
|
| 71 |
+
|
| 72 |
+
def asr_infer(self, audio: np.ndarray, output_timestamp: bool = False, key: str = "segment"):
|
| 73 |
+
self.ensure_loaded()
|
| 74 |
+
with self._lock:
|
| 75 |
+
results, meta = self.asr(
|
| 76 |
+
audio,
|
| 77 |
+
self.language,
|
| 78 |
+
self.withitn,
|
| 79 |
+
self.position_encoding,
|
| 80 |
+
tokenizer=self.tokenizer,
|
| 81 |
+
output_timestamp=output_timestamp,
|
| 82 |
+
ban_emo_unk=False,
|
| 83 |
+
key=[key],
|
| 84 |
+
)
|
| 85 |
+
text = "".join([r.get("text", "") for r in results])
|
| 86 |
+
text = clean_asr_text(text)
|
| 87 |
+
return text, meta
|
ax_meeting/pipeline.py
ADDED
|
@@ -0,0 +1,163 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# -*- coding: utf-8 -*-
|
| 2 |
+
import asyncio
|
| 3 |
+
import time
|
| 4 |
+
from dataclasses import dataclass, field
|
| 5 |
+
from typing import List, Dict, Any, Optional
|
| 6 |
+
|
| 7 |
+
import numpy as np
|
| 8 |
+
|
| 9 |
+
from ax_meeting.config import (
|
| 10 |
+
SAMPLE_RATE,
|
| 11 |
+
VAD_CHECK_INTERVAL_SEC,
|
| 12 |
+
PAUSE_MS,
|
| 13 |
+
MIN_SEGMENT_MS,
|
| 14 |
+
MERGE_VAD_MAX_LEN_MS,
|
| 15 |
+
)
|
| 16 |
+
from ax_meeting.utils.vad_utils import merge_vad
|
| 17 |
+
from ax_meeting.utils.ax_cam_bin import do_clustering, chunk
|
| 18 |
+
from ax_meeting.diar_utils import pick_speaker
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
@dataclass
|
| 22 |
+
class TranscriptSegment:
|
| 23 |
+
start_ms: int
|
| 24 |
+
end_ms: int
|
| 25 |
+
text: str
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
@dataclass
|
| 29 |
+
class StreamingMeetingSession:
|
| 30 |
+
model_bundle: Any
|
| 31 |
+
session_id: str
|
| 32 |
+
sample_rate: int = SAMPLE_RATE
|
| 33 |
+
pause_ms: int = PAUSE_MS
|
| 34 |
+
min_segment_ms: int = MIN_SEGMENT_MS
|
| 35 |
+
vad_check_interval_sec: float = VAD_CHECK_INTERVAL_SEC
|
| 36 |
+
|
| 37 |
+
audio_chunks: List[np.ndarray] = field(default_factory=list)
|
| 38 |
+
total_samples: int = 0
|
| 39 |
+
last_processed_ms: int = 0
|
| 40 |
+
last_vad_check_ts: float = field(default_factory=lambda: 0.0)
|
| 41 |
+
processing: bool = False
|
| 42 |
+
|
| 43 |
+
async def add_audio(self, pcm_bytes: bytes, websocket):
|
| 44 |
+
if not pcm_bytes:
|
| 45 |
+
return
|
| 46 |
+
audio = np.frombuffer(pcm_bytes, dtype=np.int16).astype(np.float32) / 32768.0
|
| 47 |
+
if audio.size == 0:
|
| 48 |
+
return
|
| 49 |
+
self.audio_chunks.append(audio)
|
| 50 |
+
self.total_samples += audio.size
|
| 51 |
+
|
| 52 |
+
now = time.time()
|
| 53 |
+
if (now - self.last_vad_check_ts) >= self.vad_check_interval_sec and not self.processing:
|
| 54 |
+
self.last_vad_check_ts = now
|
| 55 |
+
self.processing = True
|
| 56 |
+
asyncio.create_task(self._run_vad_and_asr(websocket))
|
| 57 |
+
|
| 58 |
+
def _current_ms(self) -> int:
|
| 59 |
+
return int(self.total_samples / self.sample_rate * 1000)
|
| 60 |
+
|
| 61 |
+
def _audio_all(self) -> np.ndarray:
|
| 62 |
+
if not self.audio_chunks:
|
| 63 |
+
return np.zeros(0, dtype=np.float32)
|
| 64 |
+
return np.concatenate(self.audio_chunks, axis=0)
|
| 65 |
+
|
| 66 |
+
async def _run_vad_and_asr(self, websocket):
|
| 67 |
+
try:
|
| 68 |
+
segments = await asyncio.to_thread(self._detect_ready_segments)
|
| 69 |
+
for seg in segments:
|
| 70 |
+
await websocket.send_json(
|
| 71 |
+
{
|
| 72 |
+
"type": "transcript",
|
| 73 |
+
"start_ms": seg.start_ms,
|
| 74 |
+
"end_ms": seg.end_ms,
|
| 75 |
+
"text": seg.text,
|
| 76 |
+
}
|
| 77 |
+
)
|
| 78 |
+
finally:
|
| 79 |
+
self.processing = False
|
| 80 |
+
|
| 81 |
+
def _detect_ready_segments(self) -> List[TranscriptSegment]:
|
| 82 |
+
audio = self._audio_all()
|
| 83 |
+
if audio.size == 0:
|
| 84 |
+
return []
|
| 85 |
+
|
| 86 |
+
current_ms = self._current_ms()
|
| 87 |
+
vad_segments = self.model_bundle.vad_infer(audio)
|
| 88 |
+
ready_segments: List[TranscriptSegment] = []
|
| 89 |
+
|
| 90 |
+
for start_ms, end_ms in vad_segments:
|
| 91 |
+
if end_ms <= self.last_processed_ms:
|
| 92 |
+
continue
|
| 93 |
+
if (end_ms - start_ms) < self.min_segment_ms:
|
| 94 |
+
continue
|
| 95 |
+
if end_ms > current_ms - self.pause_ms:
|
| 96 |
+
continue
|
| 97 |
+
|
| 98 |
+
start_sample = int(start_ms / 1000 * self.sample_rate)
|
| 99 |
+
end_sample = int(end_ms / 1000 * self.sample_rate)
|
| 100 |
+
seg_audio = audio[start_sample:end_sample]
|
| 101 |
+
text, _ = self.model_bundle.asr_infer(seg_audio, output_timestamp=False, key="live")
|
| 102 |
+
if text.strip():
|
| 103 |
+
ready_segments.append(TranscriptSegment(start_ms, end_ms, text.strip()))
|
| 104 |
+
self.last_processed_ms = max(self.last_processed_ms, int(end_ms))
|
| 105 |
+
|
| 106 |
+
return ready_segments
|
| 107 |
+
|
| 108 |
+
async def finalize(self) -> Dict[str, Any]:
|
| 109 |
+
# Flush remaining segments
|
| 110 |
+
await asyncio.to_thread(self._flush_remaining_segments)
|
| 111 |
+
|
| 112 |
+
# Full diarization + ASR
|
| 113 |
+
transcript = await asyncio.to_thread(self._full_diarization_asr)
|
| 114 |
+
return {"transcript": transcript}
|
| 115 |
+
|
| 116 |
+
def _flush_remaining_segments(self):
|
| 117 |
+
audio = self._audio_all()
|
| 118 |
+
if audio.size == 0:
|
| 119 |
+
return
|
| 120 |
+
vad_segments = self.model_bundle.vad_infer(audio)
|
| 121 |
+
for start_ms, end_ms in vad_segments:
|
| 122 |
+
if end_ms <= self.last_processed_ms:
|
| 123 |
+
continue
|
| 124 |
+
start_sample = int(start_ms / 1000 * self.sample_rate)
|
| 125 |
+
end_sample = int(end_ms / 1000 * self.sample_rate)
|
| 126 |
+
seg_audio = audio[start_sample:end_sample]
|
| 127 |
+
text, _ = self.model_bundle.asr_infer(seg_audio, output_timestamp=False, key="final")
|
| 128 |
+
self.last_processed_ms = max(self.last_processed_ms, int(end_ms))
|
| 129 |
+
|
| 130 |
+
def _full_diarization_asr(self) -> str:
|
| 131 |
+
speech = self._audio_all()
|
| 132 |
+
if speech.size == 0:
|
| 133 |
+
return ""
|
| 134 |
+
|
| 135 |
+
res_vad = self.model_bundle.vad_infer(speech)
|
| 136 |
+
vad_segments = merge_vad(res_vad, MERGE_VAD_MAX_LEN_MS)
|
| 137 |
+
|
| 138 |
+
vad_time = [[vad_t[0] / 1000, vad_t[1] / 1000] for vad_t in res_vad]
|
| 139 |
+
chunks = [c for (st, ed) in vad_time for c in chunk(st, ed)]
|
| 140 |
+
|
| 141 |
+
embeddings = self.model_bundle.speaker_infer(speech, self.sample_rate, chunks=chunks)
|
| 142 |
+
_, diar_results = do_clustering(chunks, embeddings, speaker_num=None)
|
| 143 |
+
|
| 144 |
+
lines = []
|
| 145 |
+
for i, segment in enumerate(vad_segments):
|
| 146 |
+
segment_start, segment_end = segment
|
| 147 |
+
start_sample = int(segment_start / 1000 * self.sample_rate)
|
| 148 |
+
end_sample = min(int(segment_end / 1000 * self.sample_rate), speech.shape[0])
|
| 149 |
+
segment_speech = speech[start_sample:end_sample]
|
| 150 |
+
|
| 151 |
+
text, _ = self.model_bundle.asr_infer(
|
| 152 |
+
segment_speech,
|
| 153 |
+
output_timestamp=False,
|
| 154 |
+
key=f"segment_{i}",
|
| 155 |
+
)
|
| 156 |
+
if not text or not text.strip():
|
| 157 |
+
continue
|
| 158 |
+
|
| 159 |
+
spk = pick_speaker(segment_start / 1000.0, segment_end / 1000.0, diar_results)
|
| 160 |
+
lines.append(
|
| 161 |
+
f"Speaker_{spk}: [{segment_start/1000.0:.3f} {segment_end/1000.0:.3f}] {text.strip()}"
|
| 162 |
+
)
|
| 163 |
+
return "\n".join(lines)
|
ax_meeting/positional.py
ADDED
|
@@ -0,0 +1,18 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# -*- coding: utf-8 -*-
|
| 2 |
+
import numpy as np
|
| 3 |
+
|
| 4 |
+
|
| 5 |
+
class SinusoidalPositionEncoder:
|
| 6 |
+
def __init__(self, base: float = 10000.0):
|
| 7 |
+
self.base = base
|
| 8 |
+
|
| 9 |
+
def get_position_encoding(self, dummy: np.ndarray) -> np.ndarray:
|
| 10 |
+
# dummy shape: [1, seq_len, d_model]
|
| 11 |
+
seq_len = int(dummy.shape[1])
|
| 12 |
+
d_model = int(dummy.shape[2])
|
| 13 |
+
position = np.arange(seq_len)[:, None]
|
| 14 |
+
div_term = np.exp(np.arange(0, d_model, 2) * -(np.log(self.base) / d_model))
|
| 15 |
+
pe = np.zeros((seq_len, d_model), dtype=np.float32)
|
| 16 |
+
pe[:, 0::2] = np.sin(position * div_term)
|
| 17 |
+
pe[:, 1::2] = np.cos(position * div_term)
|
| 18 |
+
return pe[None, :, :]
|
ax_meeting/server.py
ADDED
|
@@ -0,0 +1,135 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# -*- coding: utf-8 -*-
|
| 2 |
+
import asyncio
|
| 3 |
+
import json
|
| 4 |
+
import os
|
| 5 |
+
import uuid
|
| 6 |
+
from pathlib import Path
|
| 7 |
+
|
| 8 |
+
from fastapi import FastAPI, WebSocket, WebSocketDisconnect, UploadFile, File, Form
|
| 9 |
+
from fastapi.staticfiles import StaticFiles
|
| 10 |
+
from fastapi.responses import HTMLResponse, PlainTextResponse
|
| 11 |
+
|
| 12 |
+
from ax_meeting.model_bundle import ModelBundle
|
| 13 |
+
from ax_meeting.engines import DiarAsrEngine
|
| 14 |
+
from ax_meeting.pipeline import StreamingMeetingSession
|
| 15 |
+
from ax_meeting.summarizer import IncrementalSummarizer
|
| 16 |
+
|
| 17 |
+
APP_DIR = Path(__file__).parent
|
| 18 |
+
STATIC_DIR = APP_DIR / "static"
|
| 19 |
+
|
| 20 |
+
app = FastAPI()
|
| 21 |
+
models = ModelBundle()
|
| 22 |
+
app.mount("/static", StaticFiles(directory=STATIC_DIR), name="static")
|
| 23 |
+
TRANSCRIPT_STORE = {}
|
| 24 |
+
|
| 25 |
+
|
| 26 |
+
@app.get("/")
|
| 27 |
+
def index():
|
| 28 |
+
html_path = STATIC_DIR / "index.html"
|
| 29 |
+
return HTMLResponse(html_path.read_text(encoding="utf-8"))
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
@app.websocket("/ws")
|
| 33 |
+
async def ws_endpoint(ws: WebSocket):
|
| 34 |
+
await ws.accept()
|
| 35 |
+
session_id = str(uuid.uuid4())
|
| 36 |
+
# Lazy model init to avoid failing import on machines without device access
|
| 37 |
+
try:
|
| 38 |
+
await asyncio.to_thread(models.ensure_loaded)
|
| 39 |
+
except Exception as e:
|
| 40 |
+
await ws.send_json({"type": "error", "message": f"模型初始化失败: {e}"})
|
| 41 |
+
await ws.close()
|
| 42 |
+
return
|
| 43 |
+
|
| 44 |
+
session = StreamingMeetingSession(models, session_id=session_id)
|
| 45 |
+
await ws.send_json({"type": "ready", "session_id": session_id})
|
| 46 |
+
|
| 47 |
+
try:
|
| 48 |
+
while True:
|
| 49 |
+
message = await ws.receive()
|
| 50 |
+
|
| 51 |
+
if "text" in message and message["text"]:
|
| 52 |
+
try:
|
| 53 |
+
data = json.loads(message["text"])
|
| 54 |
+
except json.JSONDecodeError:
|
| 55 |
+
continue
|
| 56 |
+
|
| 57 |
+
msg_type = data.get("type")
|
| 58 |
+
if msg_type == "end":
|
| 59 |
+
result = await session.finalize()
|
| 60 |
+
|
| 61 |
+
# Summarize
|
| 62 |
+
try:
|
| 63 |
+
summarizer = IncrementalSummarizer()
|
| 64 |
+
summary = await asyncio.to_thread(
|
| 65 |
+
summarizer.summarize_incrementally, result.get("transcript", "")
|
| 66 |
+
)
|
| 67 |
+
except Exception as e:
|
| 68 |
+
summary = f"LLM 总结失败: {e}"
|
| 69 |
+
|
| 70 |
+
final_text = result.get("transcript", "")
|
| 71 |
+
TRANSCRIPT_STORE[session_id] = final_text
|
| 72 |
+
await ws.send_json({"type": "final_transcript", "text": final_text})
|
| 73 |
+
await ws.send_json({"type": "summary", "text": summary})
|
| 74 |
+
await ws.send_json({"type": "end_ack"})
|
| 75 |
+
|
| 76 |
+
elif msg_type == "ping":
|
| 77 |
+
await ws.send_json({"type": "pong"})
|
| 78 |
+
|
| 79 |
+
if "bytes" in message and message["bytes"]:
|
| 80 |
+
await session.add_audio(message["bytes"], ws)
|
| 81 |
+
|
| 82 |
+
except WebSocketDisconnect:
|
| 83 |
+
return
|
| 84 |
+
|
| 85 |
+
|
| 86 |
+
@app.post("/diar_asr")
|
| 87 |
+
async def diar_asr_api(file: UploadFile = File(...), session_id: str = Form(default="")):
|
| 88 |
+
if not session_id:
|
| 89 |
+
session_id = str(uuid.uuid4())
|
| 90 |
+
temp_path = APP_DIR / f"_upload_{session_id}_{file.filename}"
|
| 91 |
+
content = await file.read()
|
| 92 |
+
temp_path.write_bytes(content)
|
| 93 |
+
try:
|
| 94 |
+
engine = DiarAsrEngine(models)
|
| 95 |
+
text = engine.transcribe(str(temp_path))
|
| 96 |
+
finally:
|
| 97 |
+
try:
|
| 98 |
+
temp_path.unlink(missing_ok=True)
|
| 99 |
+
except Exception:
|
| 100 |
+
pass
|
| 101 |
+
TRANSCRIPT_STORE[session_id] = text
|
| 102 |
+
return {"session_id": session_id, "text": text}
|
| 103 |
+
|
| 104 |
+
|
| 105 |
+
@app.get("/export/{session_id}")
|
| 106 |
+
async def export_transcript(session_id: str):
|
| 107 |
+
text = TRANSCRIPT_STORE.get(session_id, "")
|
| 108 |
+
return PlainTextResponse(text, media_type="text/plain; charset=utf-8")
|
| 109 |
+
|
| 110 |
+
|
| 111 |
+
if __name__ == "__main__":
|
| 112 |
+
import uvicorn
|
| 113 |
+
import socket
|
| 114 |
+
|
| 115 |
+
host = os.getenv("HOST", "0.0.0.0")
|
| 116 |
+
port = int(os.getenv("PORT", "8000"))
|
| 117 |
+
ssl_cert = os.getenv("SSL_CERT")
|
| 118 |
+
ssl_key = os.getenv("SSL_KEY")
|
| 119 |
+
try:
|
| 120 |
+
s = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
|
| 121 |
+
s.connect(("8.8.8.8", 80))
|
| 122 |
+
local_ip = s.getsockname()[0]
|
| 123 |
+
s.close()
|
| 124 |
+
except Exception:
|
| 125 |
+
local_ip = "127.0.0.1"
|
| 126 |
+
scheme = "https" if ssl_cert and ssl_key else "http"
|
| 127 |
+
print(f"Local URL: {scheme}://{local_ip}:{port}")
|
| 128 |
+
uvicorn.run(
|
| 129 |
+
"ax_meeting.server:app",
|
| 130 |
+
host=host,
|
| 131 |
+
port=port,
|
| 132 |
+
reload=False,
|
| 133 |
+
ssl_certfile=ssl_cert,
|
| 134 |
+
ssl_keyfile=ssl_key,
|
| 135 |
+
)
|
ax_meeting/static/app.js
ADDED
|
@@ -0,0 +1,187 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
const startBtn = document.getElementById('startBtn');
|
| 2 |
+
const stopBtn = document.getElementById('stopBtn');
|
| 3 |
+
const liveLog = document.getElementById('liveLog');
|
| 4 |
+
const summary = document.getElementById('summary');
|
| 5 |
+
const finalTranscript = document.getElementById('finalTranscript');
|
| 6 |
+
const fileInput = document.getElementById('fileInput');
|
| 7 |
+
const diarBtn = document.getElementById('diarBtn');
|
| 8 |
+
const exportBtn = document.getElementById('exportBtn');
|
| 9 |
+
|
| 10 |
+
let ws = null;
|
| 11 |
+
let audioCtx = null;
|
| 12 |
+
let processor = null;
|
| 13 |
+
let sourceNode = null;
|
| 14 |
+
let mediaStream = null;
|
| 15 |
+
let started = false;
|
| 16 |
+
let sessionId = null;
|
| 17 |
+
|
| 18 |
+
function logLine(text) {
|
| 19 |
+
const line = document.createElement('div');
|
| 20 |
+
line.textContent = text;
|
| 21 |
+
liveLog.appendChild(line);
|
| 22 |
+
liveLog.scrollTop = liveLog.scrollHeight;
|
| 23 |
+
}
|
| 24 |
+
|
| 25 |
+
function formatTime(ms) {
|
| 26 |
+
const totalSec = ms / 1000;
|
| 27 |
+
const hours = Math.floor(totalSec / 3600);
|
| 28 |
+
const minutes = Math.floor((totalSec % 3600) / 60);
|
| 29 |
+
const seconds = totalSec % 60;
|
| 30 |
+
const h = String(hours).padStart(2, '0');
|
| 31 |
+
const m = String(minutes).padStart(2, '0');
|
| 32 |
+
const s = seconds.toFixed(2).padStart(5, '0');
|
| 33 |
+
return `${h}:${m}:${s}`;
|
| 34 |
+
}
|
| 35 |
+
|
| 36 |
+
function downsampleBuffer(buffer, inputRate, outputRate) {
|
| 37 |
+
if (outputRate === inputRate) {
|
| 38 |
+
return buffer;
|
| 39 |
+
}
|
| 40 |
+
const ratio = inputRate / outputRate;
|
| 41 |
+
const newLength = Math.round(buffer.length / ratio);
|
| 42 |
+
const result = new Float32Array(newLength);
|
| 43 |
+
let offset = 0;
|
| 44 |
+
for (let i = 0; i < newLength; i++) {
|
| 45 |
+
const nextOffset = Math.round((i + 1) * ratio);
|
| 46 |
+
let accum = 0;
|
| 47 |
+
let count = 0;
|
| 48 |
+
for (let j = offset; j < nextOffset && j < buffer.length; j++) {
|
| 49 |
+
accum += buffer[j];
|
| 50 |
+
count++;
|
| 51 |
+
}
|
| 52 |
+
result[i] = accum / Math.max(1, count);
|
| 53 |
+
offset = nextOffset;
|
| 54 |
+
}
|
| 55 |
+
return result;
|
| 56 |
+
}
|
| 57 |
+
|
| 58 |
+
function floatTo16BitPCM(float32Array) {
|
| 59 |
+
const out = new Int16Array(float32Array.length);
|
| 60 |
+
for (let i = 0; i < float32Array.length; i++) {
|
| 61 |
+
let s = Math.max(-1, Math.min(1, float32Array[i]));
|
| 62 |
+
out[i] = s < 0 ? s * 0x8000 : s * 0x7FFF;
|
| 63 |
+
}
|
| 64 |
+
return out;
|
| 65 |
+
}
|
| 66 |
+
|
| 67 |
+
async function startMeeting() {
|
| 68 |
+
if (started) return;
|
| 69 |
+
summary.textContent = '会议进行中...';
|
| 70 |
+
finalTranscript.textContent = '';
|
| 71 |
+
liveLog.innerHTML = '';
|
| 72 |
+
|
| 73 |
+
if (!window.isSecureContext && location.hostname !== 'localhost' && location.hostname !== '127.0.0.1') {
|
| 74 |
+
logLine('麦克风权限需要 HTTPS 或 localhost 访问。');
|
| 75 |
+
alert('麦克风权限需要 HTTPS 或 localhost 访问。请使用 https 或在本机用 127.0.0.1 访问。');
|
| 76 |
+
}
|
| 77 |
+
|
| 78 |
+
const wsProtocol = location.protocol === 'https:' ? 'wss' : 'ws';
|
| 79 |
+
ws = new WebSocket(`${wsProtocol}://${location.host}/ws`);
|
| 80 |
+
ws.binaryType = 'arraybuffer';
|
| 81 |
+
|
| 82 |
+
ws.onmessage = (event) => {
|
| 83 |
+
if (typeof event.data === 'string') {
|
| 84 |
+
const data = JSON.parse(event.data);
|
| 85 |
+
if (data.type === 'transcript') {
|
| 86 |
+
const start = formatTime(data.start_ms);
|
| 87 |
+
const end = formatTime(data.end_ms);
|
| 88 |
+
const t = `${start}-${end} ${data.text}`;
|
| 89 |
+
logLine(t);
|
| 90 |
+
} else if (data.type === 'ready') {
|
| 91 |
+
sessionId = data.session_id;
|
| 92 |
+
} else if (data.type === 'final_transcript') {
|
| 93 |
+
finalTranscript.textContent = data.text || '';
|
| 94 |
+
} else if (data.type === 'summary') {
|
| 95 |
+
summary.textContent = data.text || '';
|
| 96 |
+
}
|
| 97 |
+
}
|
| 98 |
+
};
|
| 99 |
+
|
| 100 |
+
ws.onopen = async () => {
|
| 101 |
+
try {
|
| 102 |
+
mediaStream = await navigator.mediaDevices.getUserMedia({ audio: true });
|
| 103 |
+
} catch (err) {
|
| 104 |
+
logLine(`麦克风权限请求失败: ${err && err.name ? err.name : err}`);
|
| 105 |
+
alert('麦克风权限请求失败。请检查浏览器权限或使用 HTTPS/localhost。');
|
| 106 |
+
return;
|
| 107 |
+
}
|
| 108 |
+
audioCtx = new (window.AudioContext || window.webkitAudioContext)();
|
| 109 |
+
sourceNode = audioCtx.createMediaStreamSource(mediaStream);
|
| 110 |
+
|
| 111 |
+
const bufferSize = 4096;
|
| 112 |
+
processor = audioCtx.createScriptProcessor(bufferSize, 1, 1);
|
| 113 |
+
|
| 114 |
+
processor.onaudioprocess = (event) => {
|
| 115 |
+
if (!ws || ws.readyState !== WebSocket.OPEN) return;
|
| 116 |
+
const input = event.inputBuffer.getChannelData(0);
|
| 117 |
+
const downsampled = downsampleBuffer(input, audioCtx.sampleRate, 16000);
|
| 118 |
+
const pcm16 = floatTo16BitPCM(downsampled);
|
| 119 |
+
ws.send(pcm16.buffer);
|
| 120 |
+
};
|
| 121 |
+
|
| 122 |
+
sourceNode.connect(processor);
|
| 123 |
+
processor.connect(audioCtx.destination);
|
| 124 |
+
|
| 125 |
+
started = true;
|
| 126 |
+
startBtn.disabled = true;
|
| 127 |
+
stopBtn.disabled = false;
|
| 128 |
+
};
|
| 129 |
+
}
|
| 130 |
+
|
| 131 |
+
async function stopMeeting() {
|
| 132 |
+
if (!started) return;
|
| 133 |
+
stopBtn.disabled = true;
|
| 134 |
+
if (ws && ws.readyState === WebSocket.OPEN) {
|
| 135 |
+
ws.send(JSON.stringify({ type: 'end' }));
|
| 136 |
+
}
|
| 137 |
+
|
| 138 |
+
if (processor) {
|
| 139 |
+
processor.disconnect();
|
| 140 |
+
processor.onaudioprocess = null;
|
| 141 |
+
}
|
| 142 |
+
if (sourceNode) {
|
| 143 |
+
sourceNode.disconnect();
|
| 144 |
+
}
|
| 145 |
+
if (audioCtx) {
|
| 146 |
+
await audioCtx.close();
|
| 147 |
+
}
|
| 148 |
+
if (mediaStream) {
|
| 149 |
+
mediaStream.getTracks().forEach((t) => t.stop());
|
| 150 |
+
}
|
| 151 |
+
|
| 152 |
+
started = false;
|
| 153 |
+
startBtn.disabled = false;
|
| 154 |
+
}
|
| 155 |
+
|
| 156 |
+
startBtn.addEventListener('click', startMeeting);
|
| 157 |
+
stopBtn.addEventListener('click', stopMeeting);
|
| 158 |
+
|
| 159 |
+
diarBtn.addEventListener('click', async () => {
|
| 160 |
+
if (!fileInput.files || fileInput.files.length === 0) {
|
| 161 |
+
alert('请选择音频文件');
|
| 162 |
+
return;
|
| 163 |
+
}
|
| 164 |
+
const file = fileInput.files[0];
|
| 165 |
+
const form = new FormData();
|
| 166 |
+
form.append('file', file);
|
| 167 |
+
if (sessionId) form.append('session_id', sessionId);
|
| 168 |
+
|
| 169 |
+
diarBtn.disabled = true;
|
| 170 |
+
try {
|
| 171 |
+
const resp = await fetch('/diar_asr', { method: 'POST', body: form });
|
| 172 |
+
const data = await resp.json();
|
| 173 |
+
if (!resp.ok) throw new Error(data.message || '识别失败');
|
| 174 |
+
if (data.session_id) sessionId = data.session_id;
|
| 175 |
+
finalTranscript.textContent = data.text || '';
|
| 176 |
+
exportBtn.disabled = !data.text;
|
| 177 |
+
} catch (err) {
|
| 178 |
+
alert(`识别失败: ${err.message || err}`);
|
| 179 |
+
} finally {
|
| 180 |
+
diarBtn.disabled = false;
|
| 181 |
+
}
|
| 182 |
+
});
|
| 183 |
+
|
| 184 |
+
exportBtn.addEventListener('click', () => {
|
| 185 |
+
if (!sessionId) return;
|
| 186 |
+
window.location.href = `/export/${sessionId}`;
|
| 187 |
+
});
|
ax_meeting/static/index.html
ADDED
|
@@ -0,0 +1,50 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
<!doctype html>
|
| 2 |
+
<html lang="zh">
|
| 3 |
+
<head>
|
| 4 |
+
<meta charset="utf-8" />
|
| 5 |
+
<meta name="viewport" content="width=device-width, initial-scale=1" />
|
| 6 |
+
<title>会议纪要流式 Demo</title>
|
| 7 |
+
<link rel="stylesheet" href="/static/style.css" />
|
| 8 |
+
</head>
|
| 9 |
+
<body>
|
| 10 |
+
<div class="page">
|
| 11 |
+
<header class="header">
|
| 12 |
+
<div>
|
| 13 |
+
<h1>会议纪要流式 Demo</h1>
|
| 14 |
+
<p>实时分段转录 + 会议结束后一键总结</p>
|
| 15 |
+
</div>
|
| 16 |
+
<div class="controls">
|
| 17 |
+
<button id="startBtn" class="btn primary">开始会议</button>
|
| 18 |
+
<button id="stopBtn" class="btn" disabled>结束会议</button>
|
| 19 |
+
</div>
|
| 20 |
+
</header>
|
| 21 |
+
|
| 22 |
+
<section class="panel">
|
| 23 |
+
<h2>实时转录</h2>
|
| 24 |
+
<div id="liveLog" class="log"></div>
|
| 25 |
+
</section>
|
| 26 |
+
|
| 27 |
+
<section class="panel">
|
| 28 |
+
<h2>离线说话人识别 + ASR</h2>
|
| 29 |
+
<div class="upload-row">
|
| 30 |
+
<input id="fileInput" class="file-input" type="file" accept=".wav,.flac,.mp3,.mp4,.m4a" />
|
| 31 |
+
<button id="diarBtn" class="btn">导入音频识别</button>
|
| 32 |
+
<button id="exportBtn" class="btn" disabled>导出 TXT</button>
|
| 33 |
+
</div>
|
| 34 |
+
<div class="hint">支持 wav/flac/mp3/mp4/m4a。处理完成后可导出文本。</div>
|
| 35 |
+
</section>
|
| 36 |
+
|
| 37 |
+
<section class="panel">
|
| 38 |
+
<h2>最终会议纪要</h2>
|
| 39 |
+
<div id="summary" class="summary">等待会议结束...</div>
|
| 40 |
+
</section>
|
| 41 |
+
|
| 42 |
+
<section class="panel">
|
| 43 |
+
<h2>完整转写(含说话人)</h2>
|
| 44 |
+
<pre id="finalTranscript" class="final"></pre>
|
| 45 |
+
</section>
|
| 46 |
+
</div>
|
| 47 |
+
|
| 48 |
+
<script src="/static/app.js"></script>
|
| 49 |
+
</body>
|
| 50 |
+
</html>
|
ax_meeting/static/style.css
ADDED
|
@@ -0,0 +1,147 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
:root {
|
| 2 |
+
--bg: #f5f3ee;
|
| 3 |
+
--ink: #1f1b16;
|
| 4 |
+
--accent: #e3693b;
|
| 5 |
+
--panel: #ffffff;
|
| 6 |
+
--muted: #666056;
|
| 7 |
+
}
|
| 8 |
+
|
| 9 |
+
* {
|
| 10 |
+
box-sizing: border-box;
|
| 11 |
+
}
|
| 12 |
+
|
| 13 |
+
body {
|
| 14 |
+
margin: 0;
|
| 15 |
+
font-family: "Noto Serif SC", "Source Han Serif SC", "Songti SC", serif;
|
| 16 |
+
background: var(--bg);
|
| 17 |
+
color: var(--ink);
|
| 18 |
+
}
|
| 19 |
+
|
| 20 |
+
.page {
|
| 21 |
+
max-width: 1100px;
|
| 22 |
+
margin: 32px auto;
|
| 23 |
+
padding: 0 20px 60px;
|
| 24 |
+
}
|
| 25 |
+
|
| 26 |
+
.header {
|
| 27 |
+
display: flex;
|
| 28 |
+
align-items: center;
|
| 29 |
+
justify-content: space-between;
|
| 30 |
+
gap: 20px;
|
| 31 |
+
margin-bottom: 24px;
|
| 32 |
+
}
|
| 33 |
+
|
| 34 |
+
.header h1 {
|
| 35 |
+
margin: 0 0 6px 0;
|
| 36 |
+
font-size: 28px;
|
| 37 |
+
}
|
| 38 |
+
|
| 39 |
+
.header p {
|
| 40 |
+
margin: 0;
|
| 41 |
+
color: var(--muted);
|
| 42 |
+
}
|
| 43 |
+
|
| 44 |
+
.controls {
|
| 45 |
+
display: flex;
|
| 46 |
+
gap: 10px;
|
| 47 |
+
}
|
| 48 |
+
|
| 49 |
+
.btn {
|
| 50 |
+
border: 1px solid var(--ink);
|
| 51 |
+
background: transparent;
|
| 52 |
+
color: var(--ink);
|
| 53 |
+
padding: 10px 16px;
|
| 54 |
+
border-radius: 8px;
|
| 55 |
+
cursor: pointer;
|
| 56 |
+
}
|
| 57 |
+
|
| 58 |
+
.btn.primary {
|
| 59 |
+
background: var(--accent);
|
| 60 |
+
border-color: var(--accent);
|
| 61 |
+
color: #fff;
|
| 62 |
+
}
|
| 63 |
+
|
| 64 |
+
.btn:disabled {
|
| 65 |
+
opacity: 0.5;
|
| 66 |
+
cursor: not-allowed;
|
| 67 |
+
}
|
| 68 |
+
|
| 69 |
+
.panel {
|
| 70 |
+
background: var(--panel);
|
| 71 |
+
border-radius: 12px;
|
| 72 |
+
padding: 16px;
|
| 73 |
+
box-shadow: 0 1px 6px rgba(0, 0, 0, 0.05);
|
| 74 |
+
margin-bottom: 16px;
|
| 75 |
+
}
|
| 76 |
+
|
| 77 |
+
.panel h2 {
|
| 78 |
+
margin: 0 0 12px 0;
|
| 79 |
+
font-size: 18px;
|
| 80 |
+
}
|
| 81 |
+
|
| 82 |
+
.upload-row {
|
| 83 |
+
display: flex;
|
| 84 |
+
gap: 10px;
|
| 85 |
+
align-items: center;
|
| 86 |
+
flex-wrap: wrap;
|
| 87 |
+
}
|
| 88 |
+
|
| 89 |
+
.file-input {
|
| 90 |
+
flex: 1 1 360px;
|
| 91 |
+
padding: 6px 8px;
|
| 92 |
+
border: 1px solid #d9d2c7;
|
| 93 |
+
border-radius: 8px;
|
| 94 |
+
background: #fffaf2;
|
| 95 |
+
font-family: "Noto Sans SC", "Source Han Sans SC", "PingFang SC", sans-serif;
|
| 96 |
+
color: var(--ink);
|
| 97 |
+
}
|
| 98 |
+
|
| 99 |
+
.file-input::file-selector-button {
|
| 100 |
+
margin-right: 10px;
|
| 101 |
+
padding: 8px 12px;
|
| 102 |
+
border: 1px solid #c9c1b5;
|
| 103 |
+
border-radius: 8px;
|
| 104 |
+
background: #efe7dc;
|
| 105 |
+
color: var(--ink);
|
| 106 |
+
cursor: pointer;
|
| 107 |
+
}
|
| 108 |
+
|
| 109 |
+
.file-input::file-selector-button:hover {
|
| 110 |
+
background: #e6ddcf;
|
| 111 |
+
}
|
| 112 |
+
|
| 113 |
+
.hint {
|
| 114 |
+
margin-top: 8px;
|
| 115 |
+
color: var(--muted);
|
| 116 |
+
font-size: 13px;
|
| 117 |
+
}
|
| 118 |
+
|
| 119 |
+
.log {
|
| 120 |
+
min-height: 160px;
|
| 121 |
+
max-height: 280px;
|
| 122 |
+
overflow: auto;
|
| 123 |
+
font-family: "JetBrains Mono", "SFMono-Regular", "Menlo", "Consolas", monospace;
|
| 124 |
+
white-space: pre-wrap;
|
| 125 |
+
}
|
| 126 |
+
|
| 127 |
+
.summary {
|
| 128 |
+
font-family: "Noto Sans SC", "Source Han Sans SC", "PingFang SC", sans-serif;
|
| 129 |
+
min-height: 120px;
|
| 130 |
+
white-space: pre-wrap;
|
| 131 |
+
}
|
| 132 |
+
|
| 133 |
+
.final {
|
| 134 |
+
min-height: 140px;
|
| 135 |
+
max-height: 320px;
|
| 136 |
+
overflow: auto;
|
| 137 |
+
background: #f6f6f6;
|
| 138 |
+
padding: 12px;
|
| 139 |
+
border-radius: 8px;
|
| 140 |
+
}
|
| 141 |
+
|
| 142 |
+
@media (max-width: 720px) {
|
| 143 |
+
.header {
|
| 144 |
+
flex-direction: column;
|
| 145 |
+
align-items: flex-start;
|
| 146 |
+
}
|
| 147 |
+
}
|
ax_meeting/summarize_cli.py
ADDED
|
@@ -0,0 +1,33 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# -*- coding: utf-8 -*-
|
| 2 |
+
import argparse
|
| 3 |
+
from pathlib import Path
|
| 4 |
+
|
| 5 |
+
from ax_meeting.summarizer import IncrementalSummarizer
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
def main():
|
| 9 |
+
parser = argparse.ArgumentParser(description="Meeting summarization (OpenAI-style API)")
|
| 10 |
+
parser.add_argument("--input", required=True, help="Transcript text file")
|
| 11 |
+
parser.add_argument("--output", default=None, help="Output summary file")
|
| 12 |
+
parser.add_argument("--openai_api_key", default=None)
|
| 13 |
+
parser.add_argument("--openai_base_url", default=None)
|
| 14 |
+
parser.add_argument("--openai_model", default=None)
|
| 15 |
+
args = parser.parse_args()
|
| 16 |
+
|
| 17 |
+
text = Path(args.input).read_text(encoding="utf-8")
|
| 18 |
+
summarizer = IncrementalSummarizer(
|
| 19 |
+
api_key=args.openai_api_key,
|
| 20 |
+
base_url=args.openai_base_url,
|
| 21 |
+
model=args.openai_model,
|
| 22 |
+
)
|
| 23 |
+
summary = summarizer.summarize_incrementally(text)
|
| 24 |
+
|
| 25 |
+
if args.output:
|
| 26 |
+
Path(args.output).write_text(summary, encoding="utf-8")
|
| 27 |
+
print(f"Summary saved: {args.output}")
|
| 28 |
+
else:
|
| 29 |
+
print(summary)
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
if __name__ == "__main__":
|
| 33 |
+
main()
|
ax_meeting/summarizer.py
ADDED
|
@@ -0,0 +1,75 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# -*- coding: utf-8 -*-
|
| 2 |
+
import os
|
| 3 |
+
import re
|
| 4 |
+
from typing import List
|
| 5 |
+
|
| 6 |
+
from ax_meeting.config import SUMMARY_CHUNK_CHARS, SUMMARY_TARGET_CHARS
|
| 7 |
+
|
| 8 |
+
try:
|
| 9 |
+
from openai import OpenAI
|
| 10 |
+
except Exception as e: # pragma: no cover
|
| 11 |
+
OpenAI = None
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
def _split_text(text: str, max_chars: int) -> List[str]:
|
| 15 |
+
text = text.strip()
|
| 16 |
+
if not text:
|
| 17 |
+
return []
|
| 18 |
+
chunks = []
|
| 19 |
+
buf = []
|
| 20 |
+
cur = 0
|
| 21 |
+
for line in text.splitlines():
|
| 22 |
+
if cur + len(line) + 1 > max_chars and buf:
|
| 23 |
+
chunks.append("\n".join(buf))
|
| 24 |
+
buf = []
|
| 25 |
+
cur = 0
|
| 26 |
+
buf.append(line)
|
| 27 |
+
cur += len(line) + 1
|
| 28 |
+
if buf:
|
| 29 |
+
chunks.append("\n".join(buf))
|
| 30 |
+
# Fallback if a single line is too long
|
| 31 |
+
if len(chunks) == 1 and len(chunks[0]) > max_chars:
|
| 32 |
+
chunks = [chunks[0][i:i + max_chars] for i in range(0, len(chunks[0]), max_chars)]
|
| 33 |
+
return chunks
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
class IncrementalSummarizer:
|
| 37 |
+
def __init__(self, api_key: str | None = None, base_url: str | None = None, model: str | None = None):
|
| 38 |
+
if OpenAI is None:
|
| 39 |
+
raise RuntimeError("openai package not installed")
|
| 40 |
+
api_key = api_key if api_key is not None else os.getenv("OPENAI_API_KEY", "")
|
| 41 |
+
base_url = base_url if base_url is not None else os.getenv("OPENAI_BASE_URL")
|
| 42 |
+
if base_url:
|
| 43 |
+
self.client = OpenAI(api_key=api_key, base_url=base_url)
|
| 44 |
+
else:
|
| 45 |
+
self.client = OpenAI(api_key=api_key)
|
| 46 |
+
self.model = model if model is not None else os.getenv("OPENAI_MODEL", "AXERA-TECH/Qwen3-1.7B")
|
| 47 |
+
|
| 48 |
+
def summarize_incrementally(self, transcript: str) -> str:
|
| 49 |
+
chunks = _split_text(transcript, SUMMARY_CHUNK_CHARS)
|
| 50 |
+
if not chunks:
|
| 51 |
+
return ""
|
| 52 |
+
|
| 53 |
+
summary = ""
|
| 54 |
+
for idx, chunk in enumerate(chunks):
|
| 55 |
+
prompt = (
|
| 56 |
+
"你是会议纪要助手。"\
|
| 57 |
+
f"\n前情提要(可为空): {summary}"\
|
| 58 |
+
f"\n本段会议文本(第{idx + 1}段):\n{chunk}"\
|
| 59 |
+
f"\n请将本段内容总结为约{SUMMARY_TARGET_CHARS}字中文摘要。"\
|
| 60 |
+
"输出要求: 只输出摘要正文, 不要标题。/no_think"
|
| 61 |
+
)
|
| 62 |
+
|
| 63 |
+
resp = self.client.chat.completions.create(
|
| 64 |
+
model=self.model,
|
| 65 |
+
messages=[
|
| 66 |
+
{"role": "system", "content": "你擅长从会议记录中抽取关键信息并总结。"},
|
| 67 |
+
{"role": "user", "content": prompt},
|
| 68 |
+
],
|
| 69 |
+
temperature=0.2,
|
| 70 |
+
)
|
| 71 |
+
|
| 72 |
+
summary = (resp.choices[0].message.content or "").strip()
|
| 73 |
+
summary = re.sub(r"<think>.*?</think>", "", summary, flags=re.DOTALL).strip()
|
| 74 |
+
print(f"Summary chunk {idx + 1}: {summary}")
|
| 75 |
+
return summary
|
ax_meeting/text_cleaner.py
ADDED
|
@@ -0,0 +1,41 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# -*- coding: utf-8 -*-
|
| 2 |
+
import re
|
| 3 |
+
|
| 4 |
+
|
| 5 |
+
def clean_asr_text(text: str) -> str:
|
| 6 |
+
if not text:
|
| 7 |
+
return ""
|
| 8 |
+
# Remove token tags like <|zh|>, <|SPEECH|> or broken sequences
|
| 9 |
+
text = re.sub(r"<\|[^>]*?\|>", "", text)
|
| 10 |
+
text = re.sub(r"<\|[^|>]*\|", "", text)
|
| 11 |
+
text = text.replace("<|", "").replace("|>", "")
|
| 12 |
+
|
| 13 |
+
# Remove plain token chains like zh|NEUTRAL|Speech|withitn|
|
| 14 |
+
text = re.sub(
|
| 15 |
+
r"(?:\b(zh|en|yue|ja|ko|speech|happy|sad|angry|neutral|emo_unknown|withitn|woitn)\b\|)+",
|
| 16 |
+
"",
|
| 17 |
+
text,
|
| 18 |
+
flags=re.IGNORECASE,
|
| 19 |
+
)
|
| 20 |
+
|
| 21 |
+
# Remove leftover meta tokens joined by pipes anywhere
|
| 22 |
+
text = re.sub(
|
| 23 |
+
r"\b(zh|en|yue|ja|ko|speech|happy|sad|angry|neutral|emo_unknown|withitn|woitn)\b",
|
| 24 |
+
"",
|
| 25 |
+
text,
|
| 26 |
+
flags=re.IGNORECASE,
|
| 27 |
+
)
|
| 28 |
+
text = text.replace("|", "")
|
| 29 |
+
|
| 30 |
+
# Remove standalone metadata tokens (case-insensitive)
|
| 31 |
+
meta_re = r"\b(zh|en|yue|ja|ko|speech|happy|sad|angry|neutral|emo_unknown|withitn|woitn)\b"
|
| 32 |
+
text = re.sub(meta_re, "", text, flags=re.IGNORECASE)
|
| 33 |
+
|
| 34 |
+
# Remove concatenated metadata tokens like zhEMO_UNKNOWNSpeechwithitn
|
| 35 |
+
meta_tokens = ["zh","en","yue","ja","ko","speech","happy","sad","angry","neutral","emo_unknown","withitn","woitn"]
|
| 36 |
+
meta_concat_re = r"(?:%s)+" % "|".join(meta_tokens)
|
| 37 |
+
text = re.sub(meta_concat_re, "", text, flags=re.IGNORECASE)
|
| 38 |
+
|
| 39 |
+
# Normalize spaces
|
| 40 |
+
text = re.sub(r"\s+", " ", text).strip()
|
| 41 |
+
return text
|
ax_meeting/utils/__init__.py
ADDED
|
File without changes
|
ax_meeting/utils/ax_cam_bin.py
ADDED
|
@@ -0,0 +1,231 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
import sys
|
| 3 |
+
import numpy as np
|
| 4 |
+
|
| 5 |
+
sys.path.append('%s'%os.path.dirname(__file__))
|
| 6 |
+
|
| 7 |
+
from ax_meeting.utils.speaker_fbank import compute_fbank
|
| 8 |
+
from ax_meeting.utils.cluster_utils import CommonClustering
|
| 9 |
+
from ax_meeting.axengine_loader import ensure_axengine
|
| 10 |
+
ensure_axengine()
|
| 11 |
+
import axengine as axe
|
| 12 |
+
|
| 13 |
+
def get_trans_sentence_sensevoice(output_asr):
|
| 14 |
+
"""Get transcription with timestamps from ASR"""
|
| 15 |
+
sentence_info = [[]]
|
| 16 |
+
punc_pattern = r'[,.!?;:"\-—…、,。!?;:""'']'
|
| 17 |
+
|
| 18 |
+
words = output_asr['merged_words']
|
| 19 |
+
#text = asr_result[0]['text']
|
| 20 |
+
timestamp = output_asr['merged_timestamps']
|
| 21 |
+
assert len(timestamp) == len(words)
|
| 22 |
+
text_pt = 0
|
| 23 |
+
|
| 24 |
+
# 遍历每个单词及其时间戳
|
| 25 |
+
for i, wd in enumerate(words):
|
| 26 |
+
# 如果当前单词是标点符号,将其与前一个单词合并
|
| 27 |
+
if wd in punc_pattern and sentence_info and sentence_info[-1]:
|
| 28 |
+
# 合并标点符号到前一个单词
|
| 29 |
+
prev_word, prev_ts = sentence_info[-1][-1]
|
| 30 |
+
sentence_info[-1][-1] = [prev_word + wd, [prev_ts[0], timestamp[i][1]]]
|
| 31 |
+
|
| 32 |
+
# # 如果标点是句子结束标点,开始新句子
|
| 33 |
+
if i < len(words) - 1:
|
| 34 |
+
sentence_info.append([])
|
| 35 |
+
else:
|
| 36 |
+
# 对于非标点单词,直接添加到当前句子
|
| 37 |
+
sentence_info[-1].append([wd, timestamp[i]])
|
| 38 |
+
return sentence_info
|
| 39 |
+
|
| 40 |
+
|
| 41 |
+
def match_spk(sentence, output_field_labels):
|
| 42 |
+
"""Match speaker ID with transcription segments"""
|
| 43 |
+
if len(sentence) == 0:
|
| 44 |
+
return []
|
| 45 |
+
|
| 46 |
+
st_sent = sentence[0][1][0]
|
| 47 |
+
ed_sent = sentence[-1][1][1]
|
| 48 |
+
overlap_per_spk = {}
|
| 49 |
+
|
| 50 |
+
for st_spk, ed_spk, spk in output_field_labels:
|
| 51 |
+
overlap_dur = min(ed_sent, ed_spk) - max(st_sent, st_spk)
|
| 52 |
+
if spk not in overlap_per_spk:
|
| 53 |
+
overlap_per_spk[spk] = 0
|
| 54 |
+
if overlap_dur > 0:
|
| 55 |
+
overlap_per_spk[spk] += overlap_dur
|
| 56 |
+
|
| 57 |
+
overlap_per_spk_list = [[spk, overlap_per_spk[spk]] for spk in overlap_per_spk if overlap_per_spk[spk] > 0]
|
| 58 |
+
overlap_per_spk_list = sorted(overlap_per_spk_list, key=lambda x:x[1], reverse=True)
|
| 59 |
+
overlap_per_spk_list = [i[0] for i in overlap_per_spk_list]
|
| 60 |
+
|
| 61 |
+
return overlap_per_spk_list
|
| 62 |
+
def distribute_spk(sentence_info, output_field_labels):
|
| 63 |
+
"""Distribute speaker IDs to transcription"""
|
| 64 |
+
last_spk = 0
|
| 65 |
+
for sentence in sentence_info:
|
| 66 |
+
main_spks = match_spk(sentence, output_field_labels)
|
| 67 |
+
main_spk = main_spks[0] if len(main_spks) > 0 else last_spk
|
| 68 |
+
|
| 69 |
+
for i, wd in enumerate(sentence):
|
| 70 |
+
wd_spks = match_spk([wd], output_field_labels)
|
| 71 |
+
if main_spk in wd_spks:
|
| 72 |
+
sentence[i].append(main_spk)
|
| 73 |
+
elif len(wd_spks) > 0:
|
| 74 |
+
sentence[i].append(wd_spks[0])
|
| 75 |
+
else:
|
| 76 |
+
sentence[i].append(last_spk)
|
| 77 |
+
last_spk = sentence[-1][2]
|
| 78 |
+
|
| 79 |
+
if len(sentence_info) == 0:
|
| 80 |
+
return []
|
| 81 |
+
|
| 82 |
+
# Merge consecutive segments from same speaker
|
| 83 |
+
sentence_info = [j for i in sentence_info for j in i]
|
| 84 |
+
sentence_info_with_spk_merge = [sentence_info[0]]
|
| 85 |
+
|
| 86 |
+
for i in sentence_info[1:]:
|
| 87 |
+
if (i[2] == sentence_info_with_spk_merge[-1][2] and
|
| 88 |
+
i[1][0] < sentence_info_with_spk_merge[-1][1][1] + 2):
|
| 89 |
+
sentence_info_with_spk_merge[-1][0] += i[0]
|
| 90 |
+
sentence_info_with_spk_merge[-1][1][1] = i[1][1]
|
| 91 |
+
else:
|
| 92 |
+
sentence_info_with_spk_merge.append(i)
|
| 93 |
+
|
| 94 |
+
return sentence_info_with_spk_merge
|
| 95 |
+
|
| 96 |
+
def chunk(st, ed, dur=1.5, step=0.75):
|
| 97 |
+
chunks = []
|
| 98 |
+
subseg_st = st
|
| 99 |
+
while subseg_st + dur < ed + step:
|
| 100 |
+
subseg_ed = min(subseg_st + dur, ed)
|
| 101 |
+
chunks.append([subseg_st, subseg_ed])
|
| 102 |
+
subseg_st += step
|
| 103 |
+
return chunks
|
| 104 |
+
|
| 105 |
+
def compressed_seg(seg_list):
|
| 106 |
+
new_seg_list = []
|
| 107 |
+
for i, seg in enumerate(seg_list):
|
| 108 |
+
seg_st, seg_ed, cluster_id = seg
|
| 109 |
+
if i == 0:
|
| 110 |
+
new_seg_list.append([seg_st, seg_ed, cluster_id])
|
| 111 |
+
elif cluster_id == new_seg_list[-1][2]:
|
| 112 |
+
if seg_st > new_seg_list[-1][1]:
|
| 113 |
+
new_seg_list.append([seg_st, seg_ed, cluster_id])
|
| 114 |
+
else:
|
| 115 |
+
new_seg_list[-1][1] = seg_ed
|
| 116 |
+
else:
|
| 117 |
+
if seg_st < new_seg_list[-1][1]:
|
| 118 |
+
p = (new_seg_list[-1][1]+seg_st) / 2
|
| 119 |
+
new_seg_list[-1][1] = p
|
| 120 |
+
seg_st = p
|
| 121 |
+
new_seg_list.append([seg_st, seg_ed, cluster_id])
|
| 122 |
+
return new_seg_list
|
| 123 |
+
|
| 124 |
+
def do_clustering(chunks, embeddings, speaker_num=None):
|
| 125 |
+
|
| 126 |
+
# kmeans 和 DBSCAN 聚类效果都不太好,pca降维无法提升聚类速度
|
| 127 |
+
|
| 128 |
+
# 对嵌入向量进行降维处理
|
| 129 |
+
# from sklearn.decomposition import PCA
|
| 130 |
+
# if embeddings.shape[1] > 50:
|
| 131 |
+
# pca = PCA(n_components=50)
|
| 132 |
+
# embeddings = pca.fit_transform(embeddings)
|
| 133 |
+
|
| 134 |
+
cluster = CommonClustering(
|
| 135 |
+
cluster_type='spectral',
|
| 136 |
+
mer_cos=0.8,
|
| 137 |
+
min_num_spks=1,
|
| 138 |
+
max_num_spks=15,
|
| 139 |
+
min_cluster_size=4,
|
| 140 |
+
oracle_num=None,
|
| 141 |
+
pval=0.012
|
| 142 |
+
)
|
| 143 |
+
cluster_labels = cluster(
|
| 144 |
+
embeddings,
|
| 145 |
+
speaker_num = speaker_num if speaker_num is not None else speaker_num
|
| 146 |
+
)
|
| 147 |
+
|
| 148 |
+
# from sklearn.cluster import DBSCAN
|
| 149 |
+
# from sklearn.neighbors import NearestNeighbors
|
| 150 |
+
# neigh = NearestNeighbors(n_neighbors=2)
|
| 151 |
+
# nbrs = neigh.fit(embeddings)
|
| 152 |
+
# distances, _ = nbrs.kneighbors(embeddings)
|
| 153 |
+
# distances = np.sort(distances, axis=0)
|
| 154 |
+
# distances = distances[:,1]
|
| 155 |
+
# eps = np.percentile(distances, 90)
|
| 156 |
+
# cluster_labels = DBSCAN(eps=eps, min_samples=5).fit_predict(embeddings)
|
| 157 |
+
|
| 158 |
+
# from sklearn.cluster import KMeans
|
| 159 |
+
# if speaker_num is None:
|
| 160 |
+
# # 如果没有提供说话人数量,使用默认值2
|
| 161 |
+
# speaker_num = 4
|
| 162 |
+
# cluster_labels = KMeans(n_clusters=speaker_num, random_state=0, n_init=10).fit_predict(embeddings)
|
| 163 |
+
|
| 164 |
+
speaker_num = cluster_labels.max()+1
|
| 165 |
+
output_field_labels = [[i[0], i[1], int(j)] for i, j in zip(chunks, cluster_labels)]
|
| 166 |
+
output_field_labels = compressed_seg(output_field_labels)
|
| 167 |
+
return speaker_num, output_field_labels
|
| 168 |
+
|
| 169 |
+
|
| 170 |
+
|
| 171 |
+
class AX_SpeakerEmbeddingInference:
|
| 172 |
+
def __init__(self, model_dir):
|
| 173 |
+
#"Initialize speaker embedding model for inference"
|
| 174 |
+
model_file = os.path.join(model_dir, "campplus.axmodel")
|
| 175 |
+
# model_file = os.path.join(model_dir, "res2netv2.axmodel")
|
| 176 |
+
|
| 177 |
+
self.session = axe.InferenceSession(model_file)
|
| 178 |
+
|
| 179 |
+
def infer(self, feats: np.ndarray) -> np.ndarray:
|
| 180 |
+
# Run inference with ONNX Runtime
|
| 181 |
+
# Run inference
|
| 182 |
+
# feats = np.expand_dims(feats, axis=-1).astype(np.float32)
|
| 183 |
+
outputs = self.session.run(None, {'feature': feats})
|
| 184 |
+
return outputs[0]
|
| 185 |
+
|
| 186 |
+
def __call__(self, wav_file, fs, chunks=None, **kwargs):
|
| 187 |
+
"""Process audio file with chunks
|
| 188 |
+
Args:
|
| 189 |
+
wav_file: path to wav file
|
| 190 |
+
chunks: list of [start_time, end_time] in seconds
|
| 191 |
+
"""
|
| 192 |
+
if chunks is None or len(chunks) == 0:
|
| 193 |
+
return np.zeros((0, 192), dtype=np.float32)
|
| 194 |
+
|
| 195 |
+
wav = wav_file.astype(np.float32)
|
| 196 |
+
if wav.ndim > 1:
|
| 197 |
+
wav = wav.reshape(-1)
|
| 198 |
+
|
| 199 |
+
wavs = [wav[int(st * fs):int(ed * fs)] for st, ed in chunks]
|
| 200 |
+
max_len = max([x.shape[0] for x in wavs])
|
| 201 |
+
max_len = max(max_len, 57900)
|
| 202 |
+
|
| 203 |
+
def circle_pad_np(x: np.ndarray, target_len: int) -> np.ndarray:
|
| 204 |
+
if x.shape[0] >= target_len:
|
| 205 |
+
return x[:target_len]
|
| 206 |
+
n = int(np.ceil(target_len / x.shape[0]))
|
| 207 |
+
xcat = np.tile(x, n)
|
| 208 |
+
return xcat[:target_len]
|
| 209 |
+
|
| 210 |
+
wavs = [circle_pad_np(x, max_len) for x in wavs]
|
| 211 |
+
|
| 212 |
+
batch_size = 1
|
| 213 |
+
embeddings = []
|
| 214 |
+
for i in range(0, len(wavs), batch_size):
|
| 215 |
+
batch_wavs = wavs[i:i+batch_size]
|
| 216 |
+
feats_list = []
|
| 217 |
+
for w in batch_wavs:
|
| 218 |
+
feat = compute_fbank(w, fs, n_mels=80, mean_nor=True)
|
| 219 |
+
if feat.shape[0] >= 360:
|
| 220 |
+
feat = feat[:360]
|
| 221 |
+
else:
|
| 222 |
+
pad = np.zeros((360 - feat.shape[0], 80), dtype=np.float32)
|
| 223 |
+
feat = np.concatenate([feat, pad], axis=0)
|
| 224 |
+
feats_list.append(feat)
|
| 225 |
+
feats_batch = np.stack(feats_list, axis=0).astype(np.float32)
|
| 226 |
+
embeddings_batch = self.infer(feats_batch)
|
| 227 |
+
embeddings.append(embeddings_batch)
|
| 228 |
+
|
| 229 |
+
# Concatenate all embeddings
|
| 230 |
+
embeddings = np.concatenate(embeddings, axis=0)
|
| 231 |
+
return embeddings
|
ax_meeting/utils/ax_model_bin.py
ADDED
|
@@ -0,0 +1,307 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python3
|
| 2 |
+
# -*- encoding: utf-8 -*-
|
| 3 |
+
# Copyright FunASR (https://github.com/FunAudioLLM/SenseVoice). All Rights Reserved.
|
| 4 |
+
# MIT License (https://opensource.org/licenses/MIT)
|
| 5 |
+
|
| 6 |
+
import os.path
|
| 7 |
+
from pathlib import Path
|
| 8 |
+
from typing import List, Union, Tuple
|
| 9 |
+
TORCH_AVAILABLE = False
|
| 10 |
+
import numpy as np
|
| 11 |
+
from ax_meeting.axengine_loader import ensure_axengine
|
| 12 |
+
ensure_axengine()
|
| 13 |
+
import axengine as axe
|
| 14 |
+
|
| 15 |
+
try:
|
| 16 |
+
import librosa
|
| 17 |
+
except ImportError:
|
| 18 |
+
print("Warning: librosa not found. Please install it using 'pip install librosa'.")
|
| 19 |
+
# Provide a fallback implementation if needed
|
| 20 |
+
def load_wav_fallback(path, sr=None):
|
| 21 |
+
import wave
|
| 22 |
+
import numpy as np
|
| 23 |
+
with wave.open(path, 'rb') as wf:
|
| 24 |
+
num_frames = wf.getnframes()
|
| 25 |
+
frames = wf.readframes(num_frames)
|
| 26 |
+
return np.frombuffer(frames, dtype=np.int16).astype(np.float32) / 32768.0, wf.getframerate()
|
| 27 |
+
|
| 28 |
+
from ax_meeting.utils.infer_utils import (
|
| 29 |
+
CharTokenizer,
|
| 30 |
+
get_logger,
|
| 31 |
+
read_yaml,
|
| 32 |
+
)
|
| 33 |
+
from ax_meeting.utils.frontend import WavFrontend
|
| 34 |
+
|
| 35 |
+
logging = get_logger()
|
| 36 |
+
|
| 37 |
+
|
| 38 |
+
def sequence_mask_np(lengths, maxlen=None, dtype=np.float32):
|
| 39 |
+
lengths = np.asarray(lengths).astype(np.int64)
|
| 40 |
+
if maxlen is None:
|
| 41 |
+
maxlen = int(lengths.max())
|
| 42 |
+
row_vector = np.arange(0, maxlen, 1)
|
| 43 |
+
matrix = lengths.reshape(-1, 1)
|
| 44 |
+
mask = row_vector < matrix
|
| 45 |
+
return mask.astype(dtype)
|
| 46 |
+
|
| 47 |
+
|
| 48 |
+
def unique_consecutive_np(arr: np.ndarray) -> np.ndarray:
|
| 49 |
+
if arr.size == 0:
|
| 50 |
+
return arr
|
| 51 |
+
out = [arr[0]]
|
| 52 |
+
for v in arr[1:]:
|
| 53 |
+
if v != out[-1]:
|
| 54 |
+
out.append(v)
|
| 55 |
+
return np.asarray(out, dtype=arr.dtype)
|
| 56 |
+
|
| 57 |
+
|
| 58 |
+
class AX_SenseVoiceSmall:
|
| 59 |
+
"""
|
| 60 |
+
Author: Speech Lab of DAMO Academy, Alibaba Group
|
| 61 |
+
Paraformer: Fast and Accurate Parallel Transformer for Non-autoregressive End-to-End Speech Recognition
|
| 62 |
+
https://arxiv.org/abs/2206.08317
|
| 63 |
+
"""
|
| 64 |
+
|
| 65 |
+
def __init__(
|
| 66 |
+
self,
|
| 67 |
+
model_dir: Union[str, Path] = None,
|
| 68 |
+
batch_size: int = 1,
|
| 69 |
+
seq_len: int = 68
|
| 70 |
+
):
|
| 71 |
+
|
| 72 |
+
model_file = os.path.join(model_dir, "sensevoice.axmodel")
|
| 73 |
+
config_file = os.path.join(model_dir, "sensevoice/config.yaml")
|
| 74 |
+
cmvn_file = os.path.join(model_dir, "sensevoice/am.mvn")
|
| 75 |
+
config = read_yaml(config_file)
|
| 76 |
+
self.model_dir = model_dir
|
| 77 |
+
# token_list = os.path.join(model_dir, "tokens.json")
|
| 78 |
+
# with open(token_list, "r", encoding="utf-8") as f:
|
| 79 |
+
# token_list = json.load(f)
|
| 80 |
+
|
| 81 |
+
# self.converter = TokenIDConverter(token_list)
|
| 82 |
+
self.tokenizer = CharTokenizer()
|
| 83 |
+
config["frontend_conf"]['cmvn_file'] = cmvn_file
|
| 84 |
+
self.frontend = WavFrontend(**config["frontend_conf"])
|
| 85 |
+
# self.ort_infer = OrtInferSession(
|
| 86 |
+
# model_file, device_id, intra_op_num_threads=intra_op_num_threads
|
| 87 |
+
# )
|
| 88 |
+
self.session = axe.InferenceSession(model_file)
|
| 89 |
+
self.batch_size = batch_size
|
| 90 |
+
self.blank_id = 0
|
| 91 |
+
self.seq_len = seq_len
|
| 92 |
+
|
| 93 |
+
self.lid_dict = {"auto": 0, "zh": 3, "en": 4, "yue": 7, "ja": 11, "ko": 12, "nospeech": 13}
|
| 94 |
+
self.lid_int_dict = {24884: 3, 24885: 4, 24888: 7, 24892: 11, 24896: 12, 24992: 13}
|
| 95 |
+
self.textnorm_dict = {"withitn": 14, "woitn": 15}
|
| 96 |
+
self.textnorm_int_dict = {25016: 14, 25017: 15}
|
| 97 |
+
self.emo_dict = {"unk": 25009, "happy": 25001, "sad": 25002, "angry": 25003, "neutral": 25004}
|
| 98 |
+
|
| 99 |
+
def __call__(self,
|
| 100 |
+
wav_content: Union[str, np.ndarray, List[str]],
|
| 101 |
+
language: str,
|
| 102 |
+
withitn: bool,
|
| 103 |
+
position_encoding: np.ndarray,
|
| 104 |
+
tokenizer=None,
|
| 105 |
+
**kwargs) -> List:
|
| 106 |
+
"""Enhanced model inference with additional features from model.py
|
| 107 |
+
|
| 108 |
+
Args:
|
| 109 |
+
wav_content: Audio data or path
|
| 110 |
+
language: Language code for processing
|
| 111 |
+
withitn: Whether to use ITN (inverse text normalization)
|
| 112 |
+
position_encoding: Position encoding tensor
|
| 113 |
+
tokenizer: Tokenizer for text conversion
|
| 114 |
+
**kwargs: Additional arguments
|
| 115 |
+
"""
|
| 116 |
+
# Start time tracking for metadata
|
| 117 |
+
import time
|
| 118 |
+
meta_data = {}
|
| 119 |
+
time_start = time.perf_counter()
|
| 120 |
+
|
| 121 |
+
# Load waveform data
|
| 122 |
+
waveform_list = self.load_data(wav_content, self.frontend.opts.frame_opts.samp_freq)
|
| 123 |
+
waveform_nums = len(waveform_list)
|
| 124 |
+
time_load = time.perf_counter()
|
| 125 |
+
meta_data["load_data"] = f"{time_load - time_start:0.3f}"
|
| 126 |
+
# Get key for result identification
|
| 127 |
+
key = kwargs.get("key", ["wav_file"])
|
| 128 |
+
if isinstance(wav_content, str):
|
| 129 |
+
wav_name = os.path.splitext(os.path.basename(wav_content))[0]
|
| 130 |
+
if key == ["wav_file"]:
|
| 131 |
+
key = [wav_name]
|
| 132 |
+
|
| 133 |
+
# Load queries from saved numpy files
|
| 134 |
+
language_query = np.load(os.path.join(self.model_dir, f"{language}.npy"))
|
| 135 |
+
textnorm_query = np.load(os.path.join(self.model_dir, "withitn.npy") if withitn
|
| 136 |
+
else os.path.join(self.model_dir, "woitn.npy"))
|
| 137 |
+
event_emo_query = np.load(os.path.join(self.model_dir, "event_emo.npy"))
|
| 138 |
+
|
| 139 |
+
# Concatenate queries to form input_query
|
| 140 |
+
input_query = np.concatenate((language_query, event_emo_query, textnorm_query), axis=1)
|
| 141 |
+
|
| 142 |
+
# Setup dataset directories for saving intermediate files
|
| 143 |
+
dataset = "dataset"
|
| 144 |
+
os.makedirs(dataset, exist_ok=True)
|
| 145 |
+
speech_dir = os.path.join(dataset, "speech", language, "withitn" if withitn else "woitn")
|
| 146 |
+
mask_dir = os.path.join(dataset, "masks", language, "withitn" if withitn else "woitn")
|
| 147 |
+
pe_dir = os.path.join(dataset, "position_encoding", language, "withitn" if withitn else "woitn")
|
| 148 |
+
os.makedirs(speech_dir, exist_ok=True)
|
| 149 |
+
os.makedirs(mask_dir, exist_ok=True)
|
| 150 |
+
os.makedirs(pe_dir, exist_ok=True)
|
| 151 |
+
|
| 152 |
+
# Process features
|
| 153 |
+
results = []
|
| 154 |
+
output_timestamp = kwargs.get("output_timestamp", False)
|
| 155 |
+
ban_emo_unk = kwargs.get("ban_emo_unk", False)
|
| 156 |
+
ibest_writer = None
|
| 157 |
+
# 添加时间偏移变量,用于跟踪连续的时间戳
|
| 158 |
+
time_offset = 0.0
|
| 159 |
+
# 添加合并时间戳变量,用于存储同一音频文件的所有时间戳,仅保留必要功能
|
| 160 |
+
merged_timestamps = []
|
| 161 |
+
merged_words = []
|
| 162 |
+
|
| 163 |
+
|
| 164 |
+
# Handle output_dir without using DatadirWriter (which is not available)
|
| 165 |
+
output_dir = kwargs.get("output_dir")
|
| 166 |
+
|
| 167 |
+
slice_len = self.seq_len - 4
|
| 168 |
+
time_pre = time.perf_counter()
|
| 169 |
+
meta_data["preprocess"] = f"{time_pre - time_load:0.3f}"
|
| 170 |
+
for beg_idx in range(0, waveform_nums, self.batch_size):
|
| 171 |
+
end_idx = min(waveform_nums, beg_idx + self.batch_size)
|
| 172 |
+
feats, feats_len = self.extract_feat(waveform_list[beg_idx:end_idx])
|
| 173 |
+
|
| 174 |
+
time_feat = time.perf_counter()
|
| 175 |
+
meta_data["extract_feat"] = f"{time_feat - time_pre:0.3f}"
|
| 176 |
+
|
| 177 |
+
for i in range(int(np.ceil(feats.shape[1] / slice_len))):
|
| 178 |
+
sub_feats = np.concatenate([input_query, feats[:, i*slice_len : (i+1)*slice_len, :]], axis=1)
|
| 179 |
+
feats_len[0] = sub_feats.shape[1]
|
| 180 |
+
|
| 181 |
+
# 计算当前片段的实际长度(帧数)
|
| 182 |
+
actual_slice_length = min(slice_len, feats.shape[1] - i*slice_len)
|
| 183 |
+
|
| 184 |
+
if feats_len[0] < self.seq_len:
|
| 185 |
+
sub_feats = np.concatenate([sub_feats, np.zeros((1, self.seq_len - feats_len[0], 560), dtype=np.float32)], axis=1)
|
| 186 |
+
|
| 187 |
+
masks = sequence_mask_np([self.seq_len], maxlen=self.seq_len, dtype=np.float32)[:, None, :]
|
| 188 |
+
|
| 189 |
+
# Run inference
|
| 190 |
+
|
| 191 |
+
ctc_logits, encoder_out_lens = self.infer(sub_feats, masks, position_encoding)
|
| 192 |
+
|
| 193 |
+
|
| 194 |
+
# Ban emotion unknown token if requested
|
| 195 |
+
if ban_emo_unk:
|
| 196 |
+
ctc_logits[:, :, self.emo_dict["unk"]] = -float("inf")
|
| 197 |
+
|
| 198 |
+
# Process results for each batch
|
| 199 |
+
b, n, d = ctc_logits.shape
|
| 200 |
+
if isinstance(key, (list, tuple)) and len(key) < b:
|
| 201 |
+
key = key * b
|
| 202 |
+
|
| 203 |
+
for j in range(b):
|
| 204 |
+
enc_len = int(encoder_out_lens[j]) if not hasattr(encoder_out_lens[j], "item") else int(encoder_out_lens[j].item())
|
| 205 |
+
x = ctc_logits[j, : enc_len, :]
|
| 206 |
+
yseq = np.argmax(x, axis=-1)
|
| 207 |
+
yseq = unique_consecutive_np(yseq)
|
| 208 |
+
|
| 209 |
+
mask = yseq != self.blank_id
|
| 210 |
+
token_int = yseq[mask].tolist()
|
| 211 |
+
|
| 212 |
+
# Convert tokens to text
|
| 213 |
+
text = tokenizer.decode(token_int) if tokenizer is not None else str(token_int)
|
| 214 |
+
# 文本处理
|
| 215 |
+
# 简化文本处理,不执行合并操作
|
| 216 |
+
|
| 217 |
+
# Write to output directory if provided
|
| 218 |
+
if ibest_writer is not None:
|
| 219 |
+
ibest_writer["text"][key[j]] = text
|
| 220 |
+
|
| 221 |
+
if output_timestamp:
|
| 222 |
+
# Torch-free build: skip timestamp generation
|
| 223 |
+
output_timestamp = False
|
| 224 |
+
result_i = {"key": key[j] if j < len(key) else f"result_{j}", "text": text}
|
| 225 |
+
|
| 226 |
+
# 直接添加结果,重复处理将在export.py中进行
|
| 227 |
+
results.append(result_i)
|
| 228 |
+
|
| 229 |
+
time_offset = round(time_offset + (actual_slice_length * 60 ) / 1000, 2)
|
| 230 |
+
# 结果处理完成
|
| 231 |
+
time_end = time.perf_counter()
|
| 232 |
+
meta_data["total_time"] = f"{time_end - time_start:0.3f}"
|
| 233 |
+
meta_data["inference_time"] = f"{time_end - time_feat:0.3f}"
|
| 234 |
+
|
| 235 |
+
# 简化元数据处理,只保留时间戳
|
| 236 |
+
if len(results) > 0 and output_timestamp and merged_timestamps:
|
| 237 |
+
# 只保留时间戳数据,以便在export.py中处理
|
| 238 |
+
meta_data["merged_timestamps"] = merged_timestamps
|
| 239 |
+
meta_data["merged_words"] = merged_words
|
| 240 |
+
|
| 241 |
+
return results, meta_data
|
| 242 |
+
|
| 243 |
+
def load_data(self, wav_content: Union[str, np.ndarray, List[str]], fs: int = None) -> List:
|
| 244 |
+
def load_wav(path: str) -> np.ndarray:
|
| 245 |
+
try:
|
| 246 |
+
# Use librosa if available
|
| 247 |
+
if 'librosa' in globals():
|
| 248 |
+
waveform, _ = librosa.load(path, sr=fs)
|
| 249 |
+
else:
|
| 250 |
+
# Use fallback implementation
|
| 251 |
+
waveform, native_sr = load_wav_fallback(path)
|
| 252 |
+
if fs is not None and native_sr != fs:
|
| 253 |
+
# Implement resampling if needed
|
| 254 |
+
print(f"Warning: Resampling from {native_sr} to {fs} is not implemented in fallback mode")
|
| 255 |
+
return waveform
|
| 256 |
+
except Exception as e:
|
| 257 |
+
print(f"Error loading audio file {path}: {e}")
|
| 258 |
+
# Return empty audio in case of error
|
| 259 |
+
return np.zeros(1600, dtype=np.float32)
|
| 260 |
+
|
| 261 |
+
if isinstance(wav_content, np.ndarray):
|
| 262 |
+
return [wav_content]
|
| 263 |
+
|
| 264 |
+
if isinstance(wav_content, str):
|
| 265 |
+
return [load_wav(wav_content)]
|
| 266 |
+
|
| 267 |
+
if isinstance(wav_content, list):
|
| 268 |
+
return [load_wav(path) for path in wav_content]
|
| 269 |
+
|
| 270 |
+
raise TypeError(f"The type of {wav_content} is not in [str, np.ndarray, list]")
|
| 271 |
+
|
| 272 |
+
def extract_feat(self, waveform_list: List[np.ndarray]) -> Tuple[np.ndarray, np.ndarray]:
|
| 273 |
+
feats, feats_len = [], []
|
| 274 |
+
for waveform in waveform_list:
|
| 275 |
+
speech, _ = self.frontend.fbank(waveform)
|
| 276 |
+
|
| 277 |
+
feat, feat_len = self.frontend.lfr_cmvn(speech)
|
| 278 |
+
|
| 279 |
+
feats.append(feat)
|
| 280 |
+
feats_len.append(feat_len)
|
| 281 |
+
|
| 282 |
+
feats = self.pad_feats(feats, np.max(feats_len))
|
| 283 |
+
feats_len = np.array(feats_len).astype(np.int32)
|
| 284 |
+
return feats, feats_len
|
| 285 |
+
|
| 286 |
+
@staticmethod
|
| 287 |
+
def pad_feats(feats: List[np.ndarray], max_feat_len: int) -> np.ndarray:
|
| 288 |
+
def pad_feat(feat: np.ndarray, cur_len: int) -> np.ndarray:
|
| 289 |
+
pad_width = ((0, max_feat_len - cur_len), (0, 0))
|
| 290 |
+
return np.pad(feat, pad_width, "constant", constant_values=0)
|
| 291 |
+
|
| 292 |
+
feat_res = [pad_feat(feat, feat.shape[0]) for feat in feats]
|
| 293 |
+
feats = np.array(feat_res).astype(np.float32)
|
| 294 |
+
return feats
|
| 295 |
+
|
| 296 |
+
def infer(self,
|
| 297 |
+
feats: np.ndarray,
|
| 298 |
+
masks: np.ndarray,
|
| 299 |
+
position_encoding: np.ndarray,
|
| 300 |
+
) -> Tuple[np.ndarray, np.ndarray]:
|
| 301 |
+
#outputs = self.ort_infer([feats, masks, position_encoding])
|
| 302 |
+
outputs =self.session.run(None, {
|
| 303 |
+
'speech': feats,
|
| 304 |
+
'masks': masks,
|
| 305 |
+
'position_encoding': position_encoding
|
| 306 |
+
})
|
| 307 |
+
return outputs
|
ax_meeting/utils/ax_vad_bin.py
ADDED
|
@@ -0,0 +1,158 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# -*- encoding: utf-8 -*-
|
| 2 |
+
# Copyright FunASR (https://github.com/alibaba-damo-academy/FunASR). All Rights Reserved.
|
| 3 |
+
# MIT License (https://opensource.org/licenses/MIT)
|
| 4 |
+
|
| 5 |
+
import os.path
|
| 6 |
+
from typing import List, Tuple
|
| 7 |
+
|
| 8 |
+
import numpy as np
|
| 9 |
+
|
| 10 |
+
from ax_meeting.utils.utils.utils import read_yaml
|
| 11 |
+
from ax_meeting.utils.utils.frontend import WavFrontend
|
| 12 |
+
from ax_meeting.utils.utils.e2e_vad import E2EVadModel
|
| 13 |
+
from ax_meeting.axengine_loader import ensure_axengine
|
| 14 |
+
ensure_axengine()
|
| 15 |
+
import axengine as axe
|
| 16 |
+
|
| 17 |
+
class AX_Fsmn_vad:
|
| 18 |
+
def __init__(self, model_dir, batch_size=1, max_end_sil=None):
|
| 19 |
+
"""Initialize VAD model for inference"""
|
| 20 |
+
|
| 21 |
+
# Export model if needed
|
| 22 |
+
model_file = os.path.join(model_dir, "vad.axmodel")
|
| 23 |
+
|
| 24 |
+
# Load config and frontend
|
| 25 |
+
config_file = os.path.join(model_dir, "vad/config.yaml")
|
| 26 |
+
cmvn_file = os.path.join(model_dir, "vad/am.mvn")
|
| 27 |
+
self.config = read_yaml(config_file)
|
| 28 |
+
self.frontend = WavFrontend(cmvn_file=cmvn_file, **self.config["frontend_conf"])
|
| 29 |
+
self.session = axe.InferenceSession(model_file)
|
| 30 |
+
self.batch_size = batch_size
|
| 31 |
+
self.vad_scorer = E2EVadModel(self.config["model_conf"])
|
| 32 |
+
self.max_end_sil = max_end_sil if max_end_sil is not None else self.config["model_conf"]["max_end_silence_time"]
|
| 33 |
+
|
| 34 |
+
def extract_feat(self, waveform_list):
|
| 35 |
+
"""Extract features from waveform"""
|
| 36 |
+
feats, feats_len = [], []
|
| 37 |
+
for waveform in waveform_list:
|
| 38 |
+
speech, _ = self.frontend.fbank(waveform)
|
| 39 |
+
feat, feat_len = self.frontend.lfr_cmvn(speech)
|
| 40 |
+
feats.append(feat)
|
| 41 |
+
feats_len.append(feat_len)
|
| 42 |
+
|
| 43 |
+
max_len = max(feats_len)
|
| 44 |
+
padded_feats = [np.pad(f, ((0, max_len - f.shape[0]), (0, 0)), 'constant') for f in feats]
|
| 45 |
+
feats = np.array(padded_feats).astype(np.float32)
|
| 46 |
+
feats_len = np.array(feats_len).astype(np.int32)
|
| 47 |
+
return feats, feats_len
|
| 48 |
+
|
| 49 |
+
def infer(self, feats: List) -> Tuple[np.ndarray, np.ndarray]:
|
| 50 |
+
"""Run inference with ONNX Runtime"""
|
| 51 |
+
# Get all input names from the model
|
| 52 |
+
input_names = [input.name for input in self.session.get_inputs()]
|
| 53 |
+
output_names = [x.name for x in self.session.get_outputs()]
|
| 54 |
+
|
| 55 |
+
# Create input dictionary for all inputs
|
| 56 |
+
input_dict = {}
|
| 57 |
+
for i, (name, tensor) in enumerate(zip(input_names, feats)):
|
| 58 |
+
input_dict[name] = tensor
|
| 59 |
+
|
| 60 |
+
# Run inference with all inputs
|
| 61 |
+
outputs = self.session.run(output_names, input_dict)
|
| 62 |
+
scores, out_caches = outputs[0], outputs[1:]
|
| 63 |
+
return scores, out_caches
|
| 64 |
+
|
| 65 |
+
def __call__(self, wav_file, **kwargs):
|
| 66 |
+
"""Process audio file with sliding window approach"""
|
| 67 |
+
# Load audio and prepare data
|
| 68 |
+
# waveform = self.load_wav(wav_file)
|
| 69 |
+
# waveform, _ = librosa.load(wav_file, sr=16000)
|
| 70 |
+
waveform_list = [wav_file]
|
| 71 |
+
waveform_nums = len(waveform_list)
|
| 72 |
+
is_final = kwargs.get("kwargs", False)
|
| 73 |
+
segments = [[]] * self.batch_size
|
| 74 |
+
|
| 75 |
+
for beg_idx in range(0, waveform_nums, self.batch_size):
|
| 76 |
+
vad_scorer = E2EVadModel(self.config["model_conf"])
|
| 77 |
+
end_idx = min(waveform_nums, beg_idx + self.batch_size)
|
| 78 |
+
waveform = waveform_list[beg_idx:end_idx]
|
| 79 |
+
feats, feats_len = self.extract_feat(waveform)
|
| 80 |
+
waveform = np.array(waveform)
|
| 81 |
+
param_dict = kwargs.get("param_dict", dict())
|
| 82 |
+
in_cache = param_dict.get("in_cache", list())
|
| 83 |
+
in_cache = self.prepare_cache(in_cache)
|
| 84 |
+
|
| 85 |
+
t_offset = 0
|
| 86 |
+
feats_len_max = int(feats_len.max()) if hasattr(feats_len, "max") else int(feats_len)
|
| 87 |
+
step = int(min(feats_len_max, 6000))
|
| 88 |
+
for t_offset in range(0, feats_len_max, min(step, feats_len_max - t_offset)):
|
| 89 |
+
if t_offset + step >= feats_len_max - 1:
|
| 90 |
+
step = feats_len_max - t_offset
|
| 91 |
+
is_final = True
|
| 92 |
+
else:
|
| 93 |
+
is_final = False
|
| 94 |
+
|
| 95 |
+
# Extract feature segment
|
| 96 |
+
feats_package = feats[:, t_offset:int(t_offset + step), :]
|
| 97 |
+
|
| 98 |
+
# Pad if it's the final segment
|
| 99 |
+
if is_final:
|
| 100 |
+
pad_length = 6000 - int(step)
|
| 101 |
+
feats_package = np.pad(
|
| 102 |
+
feats_package,
|
| 103 |
+
((0, 0), (0, pad_length), (0, 0)),
|
| 104 |
+
mode='constant',
|
| 105 |
+
constant_values=0
|
| 106 |
+
)
|
| 107 |
+
|
| 108 |
+
# Extract corresponding waveform segment
|
| 109 |
+
waveform_package = waveform[
|
| 110 |
+
:,
|
| 111 |
+
t_offset * 160:min(waveform.shape[-1], (int(t_offset + step) - 1) * 160 + 400),
|
| 112 |
+
]
|
| 113 |
+
|
| 114 |
+
# Pad waveform if it's the final segment
|
| 115 |
+
if is_final:
|
| 116 |
+
expected_wave_length = 6000 * 160 + 240
|
| 117 |
+
current_wave_length = waveform_package.shape[-1]
|
| 118 |
+
pad_wave_length = expected_wave_length - current_wave_length
|
| 119 |
+
if pad_wave_length > 0:
|
| 120 |
+
waveform_package = np.pad(
|
| 121 |
+
waveform_package,
|
| 122 |
+
((0, 0), (0, pad_wave_length)),
|
| 123 |
+
mode='constant',
|
| 124 |
+
constant_values=0
|
| 125 |
+
)
|
| 126 |
+
|
| 127 |
+
# Run inference
|
| 128 |
+
inputs = [feats_package]
|
| 129 |
+
inputs.extend(in_cache)
|
| 130 |
+
scores, out_caches = self.infer(inputs)
|
| 131 |
+
in_cache = out_caches
|
| 132 |
+
|
| 133 |
+
# Get VAD segments for this chunk
|
| 134 |
+
segments_part = vad_scorer(
|
| 135 |
+
scores,
|
| 136 |
+
waveform_package,
|
| 137 |
+
is_final=is_final,
|
| 138 |
+
max_end_sil=self.max_end_sil,
|
| 139 |
+
online=False,
|
| 140 |
+
)
|
| 141 |
+
|
| 142 |
+
# Accumulate segments
|
| 143 |
+
if segments_part:
|
| 144 |
+
for batch_num in range(0, self.batch_size):
|
| 145 |
+
segments[batch_num] += segments_part[batch_num]
|
| 146 |
+
|
| 147 |
+
return segments
|
| 148 |
+
|
| 149 |
+
def prepare_cache(self, in_cache: list = []):
|
| 150 |
+
if len(in_cache) > 0:
|
| 151 |
+
return in_cache
|
| 152 |
+
fsmn_layers = 4
|
| 153 |
+
proj_dim = 128
|
| 154 |
+
lorder = 20
|
| 155 |
+
for i in range(fsmn_layers):
|
| 156 |
+
cache = np.zeros((1, proj_dim, lorder - 1, 1)).astype(np.float32)
|
| 157 |
+
in_cache.append(cache)
|
| 158 |
+
return in_cache
|
ax_meeting/utils/cluster_utils.py
ADDED
|
@@ -0,0 +1,241 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import numpy as np
|
| 2 |
+
import scipy
|
| 3 |
+
from sklearn.cluster._kmeans import k_means
|
| 4 |
+
from sklearn.metrics.pairwise import cosine_similarity
|
| 5 |
+
|
| 6 |
+
import fastcluster
|
| 7 |
+
from scipy.cluster.hierarchy import fcluster
|
| 8 |
+
from scipy.spatial.distance import squareform
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
class SpectralCluster:
|
| 13 |
+
"""A spectral clustering method using unnormalized Laplacian of affinity matrix.
|
| 14 |
+
This implementation is adapted from https://github.com/speechbrain/speechbrain.
|
| 15 |
+
"""
|
| 16 |
+
|
| 17 |
+
def __init__(self, min_num_spks=1, max_num_spks=10, pval=0.02, min_pnum=6, oracle_num=None):
|
| 18 |
+
self.min_num_spks = min_num_spks
|
| 19 |
+
self.max_num_spks = max_num_spks
|
| 20 |
+
self.min_pnum = min_pnum
|
| 21 |
+
self.pval = pval
|
| 22 |
+
self.k = oracle_num
|
| 23 |
+
|
| 24 |
+
def __call__(self, X, **kwargs):
|
| 25 |
+
pval = kwargs.get('pval', None)
|
| 26 |
+
oracle_num = kwargs.get('speaker_num', None)
|
| 27 |
+
|
| 28 |
+
# Similarity matrix computation
|
| 29 |
+
sim_mat = self.get_sim_mat(X)
|
| 30 |
+
|
| 31 |
+
# Refining similarity matrix with pval
|
| 32 |
+
prunned_sim_mat = self.p_pruning(sim_mat, pval)
|
| 33 |
+
|
| 34 |
+
# Symmetrization
|
| 35 |
+
sym_prund_sim_mat = 0.5 * (prunned_sim_mat + prunned_sim_mat.T)
|
| 36 |
+
|
| 37 |
+
# Laplacian calculation
|
| 38 |
+
laplacian = self.get_laplacian(sym_prund_sim_mat)
|
| 39 |
+
|
| 40 |
+
# Get Spectral Embeddings
|
| 41 |
+
emb, num_of_spk = self.get_spec_embs(laplacian, oracle_num)
|
| 42 |
+
|
| 43 |
+
# Perform clustering
|
| 44 |
+
labels = self.cluster_embs(emb, num_of_spk)
|
| 45 |
+
|
| 46 |
+
return labels
|
| 47 |
+
|
| 48 |
+
def get_sim_mat(self, X):
|
| 49 |
+
# Cosine similarities
|
| 50 |
+
M = cosine_similarity(X, X)
|
| 51 |
+
return M
|
| 52 |
+
|
| 53 |
+
def p_pruning(self, A, pval=None):
|
| 54 |
+
if pval is None:
|
| 55 |
+
pval = self.pval
|
| 56 |
+
n_elems = int((1 - pval) * A.shape[0])
|
| 57 |
+
n_elems = min(n_elems, A.shape[0]-self.min_pnum)
|
| 58 |
+
|
| 59 |
+
# For each row in a affinity matrix
|
| 60 |
+
for i in range(A.shape[0]):
|
| 61 |
+
low_indexes = np.argsort(A[i, :])
|
| 62 |
+
low_indexes = low_indexes[0:n_elems]
|
| 63 |
+
|
| 64 |
+
# Replace smaller similarity values by 0s
|
| 65 |
+
A[i, low_indexes] = 0
|
| 66 |
+
return A
|
| 67 |
+
|
| 68 |
+
def get_laplacian(self, M):
|
| 69 |
+
M[np.diag_indices(M.shape[0])] = 0
|
| 70 |
+
D = np.sum(np.abs(M), axis=1)
|
| 71 |
+
D = np.diag(D)
|
| 72 |
+
L = D - M
|
| 73 |
+
return L
|
| 74 |
+
|
| 75 |
+
def get_spec_embs(self, L, k_oracle=None):
|
| 76 |
+
if k_oracle is None:
|
| 77 |
+
k_oracle = self.k
|
| 78 |
+
|
| 79 |
+
lambdas, eig_vecs = scipy.sparse.linalg.eigsh(L, k=min(self.max_num_spks+1, L.shape[0]), which='SM')
|
| 80 |
+
|
| 81 |
+
if k_oracle is not None:
|
| 82 |
+
num_of_spk = k_oracle
|
| 83 |
+
else:
|
| 84 |
+
lambda_gap_list = self.getEigenGaps(
|
| 85 |
+
lambdas[self.min_num_spks - 1:self.max_num_spks + 1])
|
| 86 |
+
num_of_spk = np.argmax(lambda_gap_list) + self.min_num_spks
|
| 87 |
+
|
| 88 |
+
emb = eig_vecs[:, :num_of_spk]
|
| 89 |
+
return emb, num_of_spk
|
| 90 |
+
|
| 91 |
+
def cluster_embs(self, emb, k):
|
| 92 |
+
# k-means
|
| 93 |
+
_, labels, _ = k_means(emb, k)
|
| 94 |
+
return labels
|
| 95 |
+
|
| 96 |
+
def getEigenGaps(self, eig_vals):
|
| 97 |
+
eig_vals_gap_list = []
|
| 98 |
+
for i in range(len(eig_vals) - 1):
|
| 99 |
+
gap = float(eig_vals[i + 1]) - float(eig_vals[i])
|
| 100 |
+
eig_vals_gap_list.append(gap)
|
| 101 |
+
return eig_vals_gap_list
|
| 102 |
+
|
| 103 |
+
|
| 104 |
+
class UmapHdbscan:
|
| 105 |
+
"""
|
| 106 |
+
Reference:
|
| 107 |
+
- Siqi Zheng, Hongbin Suo. Reformulating Speaker Diarization as Community Detection With
|
| 108 |
+
Emphasis On Topological Structure. ICASSP2022
|
| 109 |
+
"""
|
| 110 |
+
|
| 111 |
+
def __init__(self, n_neighbors=20, n_components=60, min_samples=20, min_cluster_size=10, metric='euclidean'):
|
| 112 |
+
self.n_neighbors = n_neighbors
|
| 113 |
+
self.n_components = n_components
|
| 114 |
+
self.min_samples = min_samples
|
| 115 |
+
self.min_cluster_size = min_cluster_size
|
| 116 |
+
self.metric = metric
|
| 117 |
+
import umap # type: ignore
|
| 118 |
+
import hdbscan # type: ignore
|
| 119 |
+
self._umap = umap
|
| 120 |
+
self._hdbscan = hdbscan
|
| 121 |
+
|
| 122 |
+
def __call__(self, X, **kwargs):
|
| 123 |
+
umap_X = self._umap.UMAP(
|
| 124 |
+
n_neighbors=self.n_neighbors,
|
| 125 |
+
min_dist=0.0,
|
| 126 |
+
n_components=min(self.n_components, X.shape[0]-2),
|
| 127 |
+
metric=self.metric,
|
| 128 |
+
).fit_transform(X)
|
| 129 |
+
labels = self._hdbscan.HDBSCAN(
|
| 130 |
+
min_samples=self.min_samples,
|
| 131 |
+
min_cluster_size=self.min_cluster_size
|
| 132 |
+
).fit_predict(umap_X)
|
| 133 |
+
return labels
|
| 134 |
+
|
| 135 |
+
class AHCluster:
|
| 136 |
+
"""
|
| 137 |
+
Agglomerative Hierarchical Clustering, a bottom-up approach which iteratively merges
|
| 138 |
+
the closest clusters until a termination condition is reached.
|
| 139 |
+
This implementation is adapted from https://github.com/BUTSpeechFIT/VBx.
|
| 140 |
+
"""
|
| 141 |
+
|
| 142 |
+
def __init__(self, fix_cos_thr=0.4):
|
| 143 |
+
self.fix_cos_thr = fix_cos_thr
|
| 144 |
+
|
| 145 |
+
def __call__(self, X, **kwargs):
|
| 146 |
+
scr_mx = cosine_similarity(X)
|
| 147 |
+
scr_mx = squareform(-scr_mx, checks=False)
|
| 148 |
+
lin_mat = fastcluster.linkage(scr_mx, method='average', preserve_input='False')
|
| 149 |
+
adjust = abs(lin_mat[:, 2].min())
|
| 150 |
+
lin_mat[:, 2] += adjust
|
| 151 |
+
labels = fcluster(lin_mat, -self.fix_cos_thr + adjust, criterion='distance') - 1
|
| 152 |
+
return labels
|
| 153 |
+
|
| 154 |
+
|
| 155 |
+
class CommonClustering:
|
| 156 |
+
"""Perfom clustering for input embeddings and output the labels.
|
| 157 |
+
"""
|
| 158 |
+
|
| 159 |
+
def __init__(self, cluster_type, cluster_line=40, mer_cos=None, min_cluster_size=4, **kwargs):
|
| 160 |
+
self.cluster_type = cluster_type
|
| 161 |
+
self.cluster_line = cluster_line
|
| 162 |
+
self.min_cluster_size = min_cluster_size
|
| 163 |
+
self.mer_cos = mer_cos
|
| 164 |
+
|
| 165 |
+
# Initialize main cluster
|
| 166 |
+
if self.cluster_type == 'spectral':
|
| 167 |
+
self.cluster = SpectralCluster(**kwargs)
|
| 168 |
+
elif self.cluster_type == 'umap_hdbscan':
|
| 169 |
+
kwargs['min_cluster_size'] = min_cluster_size
|
| 170 |
+
self.cluster = UmapHdbscan(**kwargs)
|
| 171 |
+
elif self.cluster_type == 'AHC':
|
| 172 |
+
self.cluster = AHCluster(**kwargs)
|
| 173 |
+
else:
|
| 174 |
+
raise ValueError(
|
| 175 |
+
'%s is not currently supported.' % self.cluster_type
|
| 176 |
+
)
|
| 177 |
+
|
| 178 |
+
# Initialize short cluster
|
| 179 |
+
if self.cluster_type != 'AHC':
|
| 180 |
+
self.cluster_for_short = AHCluster()
|
| 181 |
+
else:
|
| 182 |
+
self.cluster_for_short = self.cluster
|
| 183 |
+
|
| 184 |
+
|
| 185 |
+
def __call__(self, X, **kwargs):
|
| 186 |
+
# clustering and return the labels
|
| 187 |
+
assert len(X.shape) == 2, 'Shape of input should be [N, C]'
|
| 188 |
+
if X.shape[0] <= 1:
|
| 189 |
+
return np.zeros(X.shape[0], dtype=int)
|
| 190 |
+
|
| 191 |
+
if X.shape[0] < self.cluster_line:
|
| 192 |
+
labels = self.cluster_for_short(X)
|
| 193 |
+
else:
|
| 194 |
+
labels = self.cluster(X, **kwargs)
|
| 195 |
+
|
| 196 |
+
# remove extremely minor cluster
|
| 197 |
+
labels = self.filter_minor_cluster(labels, X, self.min_cluster_size)
|
| 198 |
+
|
| 199 |
+
# merge similar speaker
|
| 200 |
+
if self.mer_cos is not None:
|
| 201 |
+
labels = self.merge_by_cos(labels, X, self.mer_cos)
|
| 202 |
+
return labels
|
| 203 |
+
|
| 204 |
+
def filter_minor_cluster(self, labels, x, min_cluster_size):
|
| 205 |
+
cset = np.unique(labels)
|
| 206 |
+
csize = np.array([(labels == i).sum() for i in cset])
|
| 207 |
+
minor_idx = np.where(csize <= self.min_cluster_size)[0]
|
| 208 |
+
if len(minor_idx) == 0:
|
| 209 |
+
return labels
|
| 210 |
+
|
| 211 |
+
minor_cset = cset[minor_idx]
|
| 212 |
+
major_idx = np.where(csize > self.min_cluster_size)[0]
|
| 213 |
+
if len(major_idx) == 0:
|
| 214 |
+
return np.zeros_like(labels)
|
| 215 |
+
major_cset = cset[major_idx]
|
| 216 |
+
major_center = np.stack([x[labels == i].mean(0) \
|
| 217 |
+
for i in major_cset])
|
| 218 |
+
for i in range(len(labels)):
|
| 219 |
+
if labels[i] in minor_cset:
|
| 220 |
+
cos_sim = cosine_similarity(x[i][np.newaxis], major_center)
|
| 221 |
+
labels[i] = major_cset[cos_sim.argmax()]
|
| 222 |
+
|
| 223 |
+
return labels
|
| 224 |
+
|
| 225 |
+
def merge_by_cos(self, labels, x, cos_thr):
|
| 226 |
+
# merge the similar speakers by cosine similarity
|
| 227 |
+
assert cos_thr > 0 and cos_thr <= 1
|
| 228 |
+
while True:
|
| 229 |
+
cset = np.unique(labels)
|
| 230 |
+
if len(cset) == 1:
|
| 231 |
+
break
|
| 232 |
+
centers = np.stack([x[labels == i].mean(0) \
|
| 233 |
+
for i in cset])
|
| 234 |
+
affinity = cosine_similarity(centers, centers)
|
| 235 |
+
affinity = np.triu(affinity, 1)
|
| 236 |
+
idx = np.unravel_index(np.argmax(affinity), affinity.shape)
|
| 237 |
+
if affinity[idx] < cos_thr:
|
| 238 |
+
break
|
| 239 |
+
c1, c2 = cset[np.array(idx)]
|
| 240 |
+
labels[labels==c2]=c1
|
| 241 |
+
return labels
|
ax_meeting/utils/ctc_alignment.py
ADDED
|
@@ -0,0 +1,76 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
|
| 3 |
+
def ctc_forced_align(
|
| 4 |
+
log_probs: torch.Tensor,
|
| 5 |
+
targets: torch.Tensor,
|
| 6 |
+
input_lengths: torch.Tensor,
|
| 7 |
+
target_lengths: torch.Tensor,
|
| 8 |
+
blank: int = 0,
|
| 9 |
+
ignore_id: int = -1,
|
| 10 |
+
) -> torch.Tensor:
|
| 11 |
+
"""Align a CTC label sequence to an emission.
|
| 12 |
+
|
| 13 |
+
Args:
|
| 14 |
+
log_probs (Tensor): log probability of CTC emission output.
|
| 15 |
+
Tensor of shape `(B, T, C)`. where `B` is the batch size, `T` is the input length,
|
| 16 |
+
`C` is the number of characters in alphabet including blank.
|
| 17 |
+
targets (Tensor): Target sequence. Tensor of shape `(B, L)`,
|
| 18 |
+
where `L` is the target length.
|
| 19 |
+
input_lengths (Tensor):
|
| 20 |
+
Lengths of the inputs (max value must each be <= `T`). 1-D Tensor of shape `(B,)`.
|
| 21 |
+
target_lengths (Tensor):
|
| 22 |
+
Lengths of the targets. 1-D Tensor of shape `(B,)`.
|
| 23 |
+
blank_id (int, optional): The index of blank symbol in CTC emission. (Default: 0)
|
| 24 |
+
ignore_id (int, optional): The index of ignore symbol in CTC emission. (Default: -1)
|
| 25 |
+
"""
|
| 26 |
+
targets[targets == ignore_id] = blank
|
| 27 |
+
|
| 28 |
+
batch_size, input_time_size, _ = log_probs.size()
|
| 29 |
+
bsz_indices = torch.arange(batch_size, device=input_lengths.device)
|
| 30 |
+
|
| 31 |
+
_t_a_r_g_e_t_s_ = torch.cat(
|
| 32 |
+
(
|
| 33 |
+
torch.stack((torch.full_like(targets, blank), targets), dim=-1).flatten(start_dim=1),
|
| 34 |
+
torch.full_like(targets[:, :1], blank),
|
| 35 |
+
),
|
| 36 |
+
dim=-1,
|
| 37 |
+
)
|
| 38 |
+
diff_labels = torch.cat(
|
| 39 |
+
(
|
| 40 |
+
torch.as_tensor([[False, False]], device=targets.device).expand(batch_size, -1),
|
| 41 |
+
_t_a_r_g_e_t_s_[:, 2:] != _t_a_r_g_e_t_s_[:, :-2],
|
| 42 |
+
),
|
| 43 |
+
dim=1,
|
| 44 |
+
)
|
| 45 |
+
|
| 46 |
+
neg_inf = torch.tensor(float("-inf"), device=log_probs.device, dtype=log_probs.dtype)
|
| 47 |
+
padding_num = 2
|
| 48 |
+
padded_t = padding_num + _t_a_r_g_e_t_s_.size(-1)
|
| 49 |
+
best_score = torch.full((batch_size, padded_t), neg_inf, device=log_probs.device, dtype=log_probs.dtype)
|
| 50 |
+
best_score[:, padding_num + 0] = log_probs[:, 0, blank]
|
| 51 |
+
best_score[:, padding_num + 1] = log_probs[bsz_indices, 0, _t_a_r_g_e_t_s_[:, 1]]
|
| 52 |
+
|
| 53 |
+
backpointers = torch.zeros((batch_size, input_time_size, padded_t), device=log_probs.device, dtype=targets.dtype)
|
| 54 |
+
|
| 55 |
+
for t in range(1, input_time_size):
|
| 56 |
+
prev = torch.stack(
|
| 57 |
+
(best_score[:, 2:], best_score[:, 1:-1], torch.where(diff_labels, best_score[:, :-2], neg_inf))
|
| 58 |
+
)
|
| 59 |
+
prev_max_value, prev_max_idx = prev.max(dim=0)
|
| 60 |
+
best_score[:, padding_num:] = log_probs[:, t].gather(-1, _t_a_r_g_e_t_s_) + prev_max_value
|
| 61 |
+
backpointers[:, t, padding_num:] = prev_max_idx
|
| 62 |
+
|
| 63 |
+
l1l2 = best_score.gather(
|
| 64 |
+
-1, torch.stack((padding_num + target_lengths * 2 - 1, padding_num + target_lengths * 2), dim=-1)
|
| 65 |
+
)
|
| 66 |
+
|
| 67 |
+
path = torch.zeros((batch_size, input_time_size), device=best_score.device, dtype=torch.long)
|
| 68 |
+
path[bsz_indices, input_lengths - 1] = padding_num + target_lengths * 2 - 1 + l1l2.argmax(dim=-1)
|
| 69 |
+
|
| 70 |
+
for t in range(input_time_size - 1, 0, -1):
|
| 71 |
+
target_indices = path[:, t]
|
| 72 |
+
prev_max_idx = backpointers[bsz_indices, t, target_indices]
|
| 73 |
+
path[:, t - 1] += target_indices - prev_max_idx
|
| 74 |
+
|
| 75 |
+
alignments = _t_a_r_g_e_t_s_.gather(dim=-1, index=(path - padding_num).clamp(min=0))
|
| 76 |
+
return alignments
|
ax_meeting/utils/frontend.py
ADDED
|
@@ -0,0 +1,433 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# -*- encoding: utf-8 -*-
|
| 2 |
+
from pathlib import Path
|
| 3 |
+
from typing import Any, Dict, Iterable, List, NamedTuple, Set, Tuple, Union
|
| 4 |
+
import copy
|
| 5 |
+
|
| 6 |
+
import numpy as np
|
| 7 |
+
import kaldi_native_fbank as knf
|
| 8 |
+
|
| 9 |
+
root_dir = Path(__file__).resolve().parent
|
| 10 |
+
|
| 11 |
+
logger_initialized = {}
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
class WavFrontend:
|
| 15 |
+
"""Conventional frontend structure for ASR."""
|
| 16 |
+
|
| 17 |
+
def __init__(
|
| 18 |
+
self,
|
| 19 |
+
cmvn_file: str = None,
|
| 20 |
+
fs: int = 16000,
|
| 21 |
+
window: str = "hamming",
|
| 22 |
+
n_mels: int = 80,
|
| 23 |
+
frame_length: int = 25,
|
| 24 |
+
frame_shift: int = 10,
|
| 25 |
+
lfr_m: int = 1,
|
| 26 |
+
lfr_n: int = 1,
|
| 27 |
+
dither: float = 1.0,
|
| 28 |
+
**kwargs,
|
| 29 |
+
) -> None:
|
| 30 |
+
|
| 31 |
+
opts = knf.FbankOptions()
|
| 32 |
+
opts.frame_opts.samp_freq = fs
|
| 33 |
+
opts.frame_opts.dither = dither
|
| 34 |
+
opts.frame_opts.window_type = window
|
| 35 |
+
opts.frame_opts.frame_shift_ms = float(frame_shift)
|
| 36 |
+
opts.frame_opts.frame_length_ms = float(frame_length)
|
| 37 |
+
opts.mel_opts.num_bins = n_mels
|
| 38 |
+
opts.energy_floor = 0
|
| 39 |
+
opts.frame_opts.snip_edges = True
|
| 40 |
+
opts.mel_opts.debug_mel = False
|
| 41 |
+
self.opts = opts
|
| 42 |
+
|
| 43 |
+
self.lfr_m = lfr_m
|
| 44 |
+
self.lfr_n = lfr_n
|
| 45 |
+
self.cmvn_file = cmvn_file
|
| 46 |
+
|
| 47 |
+
if self.cmvn_file:
|
| 48 |
+
self.cmvn = self.load_cmvn()
|
| 49 |
+
self.fbank_fn = None
|
| 50 |
+
self.fbank_beg_idx = 0
|
| 51 |
+
self.reset_status()
|
| 52 |
+
|
| 53 |
+
def fbank(self, waveform: np.ndarray) -> Tuple[np.ndarray, np.ndarray]:
|
| 54 |
+
waveform = waveform * (1 << 15)
|
| 55 |
+
self.fbank_fn = knf.OnlineFbank(self.opts)
|
| 56 |
+
self.fbank_fn.accept_waveform(self.opts.frame_opts.samp_freq, waveform.tolist())
|
| 57 |
+
frames = self.fbank_fn.num_frames_ready
|
| 58 |
+
mat = np.empty([frames, self.opts.mel_opts.num_bins])
|
| 59 |
+
for i in range(frames):
|
| 60 |
+
mat[i, :] = self.fbank_fn.get_frame(i)
|
| 61 |
+
feat = mat.astype(np.float32)
|
| 62 |
+
feat_len = np.array(mat.shape[0]).astype(np.int32)
|
| 63 |
+
return feat, feat_len
|
| 64 |
+
|
| 65 |
+
def fbank_online(self, waveform: np.ndarray) -> Tuple[np.ndarray, np.ndarray]:
|
| 66 |
+
waveform = waveform * (1 << 15)
|
| 67 |
+
# self.fbank_fn = knf.OnlineFbank(self.opts)
|
| 68 |
+
self.fbank_fn.accept_waveform(self.opts.frame_opts.samp_freq, waveform.tolist())
|
| 69 |
+
frames = self.fbank_fn.num_frames_ready
|
| 70 |
+
mat = np.empty([frames, self.opts.mel_opts.num_bins])
|
| 71 |
+
for i in range(self.fbank_beg_idx, frames):
|
| 72 |
+
mat[i, :] = self.fbank_fn.get_frame(i)
|
| 73 |
+
# self.fbank_beg_idx += (frames-self.fbank_beg_idx)
|
| 74 |
+
feat = mat.astype(np.float32)
|
| 75 |
+
feat_len = np.array(mat.shape[0]).astype(np.int32)
|
| 76 |
+
return feat, feat_len
|
| 77 |
+
|
| 78 |
+
def reset_status(self):
|
| 79 |
+
self.fbank_fn = knf.OnlineFbank(self.opts)
|
| 80 |
+
self.fbank_beg_idx = 0
|
| 81 |
+
|
| 82 |
+
def lfr_cmvn(self, feat: np.ndarray) -> Tuple[np.ndarray, np.ndarray]:
|
| 83 |
+
if self.lfr_m != 1 or self.lfr_n != 1:
|
| 84 |
+
feat = self.apply_lfr(feat, self.lfr_m, self.lfr_n)
|
| 85 |
+
|
| 86 |
+
if self.cmvn_file:
|
| 87 |
+
feat = self.apply_cmvn(feat)
|
| 88 |
+
|
| 89 |
+
feat_len = np.array(feat.shape[0]).astype(np.int32)
|
| 90 |
+
return feat, feat_len
|
| 91 |
+
|
| 92 |
+
@staticmethod
|
| 93 |
+
def apply_lfr(inputs: np.ndarray, lfr_m: int, lfr_n: int) -> np.ndarray:
|
| 94 |
+
LFR_inputs = []
|
| 95 |
+
|
| 96 |
+
T = inputs.shape[0]
|
| 97 |
+
T_lfr = int(np.ceil(T / lfr_n))
|
| 98 |
+
left_padding = np.tile(inputs[0], ((lfr_m - 1) // 2, 1))
|
| 99 |
+
inputs = np.vstack((left_padding, inputs))
|
| 100 |
+
T = T + (lfr_m - 1) // 2
|
| 101 |
+
for i in range(T_lfr):
|
| 102 |
+
if lfr_m <= T - i * lfr_n:
|
| 103 |
+
LFR_inputs.append((inputs[i * lfr_n : i * lfr_n + lfr_m]).reshape(1, -1))
|
| 104 |
+
else:
|
| 105 |
+
# process last LFR frame
|
| 106 |
+
num_padding = lfr_m - (T - i * lfr_n)
|
| 107 |
+
frame = inputs[i * lfr_n :].reshape(-1)
|
| 108 |
+
for _ in range(num_padding):
|
| 109 |
+
frame = np.hstack((frame, inputs[-1]))
|
| 110 |
+
|
| 111 |
+
LFR_inputs.append(frame)
|
| 112 |
+
LFR_outputs = np.vstack(LFR_inputs).astype(np.float32)
|
| 113 |
+
return LFR_outputs
|
| 114 |
+
|
| 115 |
+
def apply_cmvn(self, inputs: np.ndarray) -> np.ndarray:
|
| 116 |
+
"""
|
| 117 |
+
Apply CMVN with mvn data
|
| 118 |
+
"""
|
| 119 |
+
frame, dim = inputs.shape
|
| 120 |
+
means = np.tile(self.cmvn[0:1, :dim], (frame, 1))
|
| 121 |
+
vars = np.tile(self.cmvn[1:2, :dim], (frame, 1))
|
| 122 |
+
inputs = (inputs + means) * vars
|
| 123 |
+
return inputs
|
| 124 |
+
|
| 125 |
+
def load_cmvn(
|
| 126 |
+
self,
|
| 127 |
+
) -> np.ndarray:
|
| 128 |
+
with open(self.cmvn_file, "r", encoding="utf-8") as f:
|
| 129 |
+
lines = f.readlines()
|
| 130 |
+
|
| 131 |
+
means_list = []
|
| 132 |
+
vars_list = []
|
| 133 |
+
for i in range(len(lines)):
|
| 134 |
+
line_item = lines[i].split()
|
| 135 |
+
if line_item[0] == "<AddShift>":
|
| 136 |
+
line_item = lines[i + 1].split()
|
| 137 |
+
if line_item[0] == "<LearnRateCoef>":
|
| 138 |
+
add_shift_line = line_item[3 : (len(line_item) - 1)]
|
| 139 |
+
means_list = list(add_shift_line)
|
| 140 |
+
continue
|
| 141 |
+
elif line_item[0] == "<Rescale>":
|
| 142 |
+
line_item = lines[i + 1].split()
|
| 143 |
+
if line_item[0] == "<LearnRateCoef>":
|
| 144 |
+
rescale_line = line_item[3 : (len(line_item) - 1)]
|
| 145 |
+
vars_list = list(rescale_line)
|
| 146 |
+
continue
|
| 147 |
+
|
| 148 |
+
means = np.array(means_list).astype(np.float64)
|
| 149 |
+
vars = np.array(vars_list).astype(np.float64)
|
| 150 |
+
cmvn = np.array([means, vars])
|
| 151 |
+
return cmvn
|
| 152 |
+
|
| 153 |
+
|
| 154 |
+
class WavFrontendOnline(WavFrontend):
|
| 155 |
+
def __init__(self, **kwargs):
|
| 156 |
+
super().__init__(**kwargs)
|
| 157 |
+
# self.fbank_fn = knf.OnlineFbank(self.opts)
|
| 158 |
+
# add variables
|
| 159 |
+
self.frame_sample_length = int(
|
| 160 |
+
self.opts.frame_opts.frame_length_ms * self.opts.frame_opts.samp_freq / 1000
|
| 161 |
+
)
|
| 162 |
+
self.frame_shift_sample_length = int(
|
| 163 |
+
self.opts.frame_opts.frame_shift_ms * self.opts.frame_opts.samp_freq / 1000
|
| 164 |
+
)
|
| 165 |
+
self.waveform = None
|
| 166 |
+
self.reserve_waveforms = None
|
| 167 |
+
self.input_cache = None
|
| 168 |
+
self.lfr_splice_cache = []
|
| 169 |
+
|
| 170 |
+
@staticmethod
|
| 171 |
+
# inputs has catted the cache
|
| 172 |
+
def apply_lfr(
|
| 173 |
+
inputs: np.ndarray, lfr_m: int, lfr_n: int, is_final: bool = False
|
| 174 |
+
) -> Tuple[np.ndarray, np.ndarray, int]:
|
| 175 |
+
"""
|
| 176 |
+
Apply lfr with data
|
| 177 |
+
"""
|
| 178 |
+
|
| 179 |
+
LFR_inputs = []
|
| 180 |
+
T = inputs.shape[0] # include the right context
|
| 181 |
+
T_lfr = int(
|
| 182 |
+
np.ceil((T - (lfr_m - 1) // 2) / lfr_n)
|
| 183 |
+
) # minus the right context: (lfr_m - 1) // 2
|
| 184 |
+
splice_idx = T_lfr
|
| 185 |
+
for i in range(T_lfr):
|
| 186 |
+
if lfr_m <= T - i * lfr_n:
|
| 187 |
+
LFR_inputs.append((inputs[i * lfr_n : i * lfr_n + lfr_m]).reshape(1, -1))
|
| 188 |
+
else: # process last LFR frame
|
| 189 |
+
if is_final:
|
| 190 |
+
num_padding = lfr_m - (T - i * lfr_n)
|
| 191 |
+
frame = (inputs[i * lfr_n :]).reshape(-1)
|
| 192 |
+
for _ in range(num_padding):
|
| 193 |
+
frame = np.hstack((frame, inputs[-1]))
|
| 194 |
+
LFR_inputs.append(frame)
|
| 195 |
+
else:
|
| 196 |
+
# update splice_idx and break the circle
|
| 197 |
+
splice_idx = i
|
| 198 |
+
break
|
| 199 |
+
splice_idx = min(T - 1, splice_idx * lfr_n)
|
| 200 |
+
lfr_splice_cache = inputs[splice_idx:, :]
|
| 201 |
+
LFR_outputs = np.vstack(LFR_inputs)
|
| 202 |
+
return LFR_outputs.astype(np.float32), lfr_splice_cache, splice_idx
|
| 203 |
+
|
| 204 |
+
@staticmethod
|
| 205 |
+
def compute_frame_num(
|
| 206 |
+
sample_length: int, frame_sample_length: int, frame_shift_sample_length: int
|
| 207 |
+
) -> int:
|
| 208 |
+
frame_num = int((sample_length - frame_sample_length) / frame_shift_sample_length + 1)
|
| 209 |
+
return frame_num if frame_num >= 1 and sample_length >= frame_sample_length else 0
|
| 210 |
+
|
| 211 |
+
def fbank(
|
| 212 |
+
self, input: np.ndarray, input_lengths: np.ndarray
|
| 213 |
+
) -> Tuple[np.ndarray, np.ndarray, np.ndarray]:
|
| 214 |
+
self.fbank_fn = knf.OnlineFbank(self.opts)
|
| 215 |
+
batch_size = input.shape[0]
|
| 216 |
+
if self.input_cache is None:
|
| 217 |
+
self.input_cache = np.empty((batch_size, 0), dtype=np.float32)
|
| 218 |
+
input = np.concatenate((self.input_cache, input), axis=1)
|
| 219 |
+
frame_num = self.compute_frame_num(
|
| 220 |
+
input.shape[-1], self.frame_sample_length, self.frame_shift_sample_length
|
| 221 |
+
)
|
| 222 |
+
# update self.in_cache
|
| 223 |
+
self.input_cache = input[
|
| 224 |
+
:, -(input.shape[-1] - frame_num * self.frame_shift_sample_length) :
|
| 225 |
+
]
|
| 226 |
+
waveforms = np.empty(0, dtype=np.float32)
|
| 227 |
+
feats_pad = np.empty(0, dtype=np.float32)
|
| 228 |
+
feats_lens = np.empty(0, dtype=np.int32)
|
| 229 |
+
if frame_num:
|
| 230 |
+
waveforms = []
|
| 231 |
+
feats = []
|
| 232 |
+
feats_lens = []
|
| 233 |
+
for i in range(batch_size):
|
| 234 |
+
waveform = input[i]
|
| 235 |
+
waveforms.append(
|
| 236 |
+
waveform[
|
| 237 |
+
: (
|
| 238 |
+
(frame_num - 1) * self.frame_shift_sample_length
|
| 239 |
+
+ self.frame_sample_length
|
| 240 |
+
)
|
| 241 |
+
]
|
| 242 |
+
)
|
| 243 |
+
waveform = waveform * (1 << 15)
|
| 244 |
+
|
| 245 |
+
self.fbank_fn.accept_waveform(self.opts.frame_opts.samp_freq, waveform.tolist())
|
| 246 |
+
frames = self.fbank_fn.num_frames_ready
|
| 247 |
+
mat = np.empty([frames, self.opts.mel_opts.num_bins])
|
| 248 |
+
for i in range(frames):
|
| 249 |
+
mat[i, :] = self.fbank_fn.get_frame(i)
|
| 250 |
+
feat = mat.astype(np.float32)
|
| 251 |
+
feat_len = np.array(mat.shape[0]).astype(np.int32)
|
| 252 |
+
feats.append(feat)
|
| 253 |
+
feats_lens.append(feat_len)
|
| 254 |
+
|
| 255 |
+
waveforms = np.stack(waveforms)
|
| 256 |
+
feats_lens = np.array(feats_lens)
|
| 257 |
+
feats_pad = np.array(feats)
|
| 258 |
+
self.fbanks = feats_pad
|
| 259 |
+
self.fbanks_lens = copy.deepcopy(feats_lens)
|
| 260 |
+
return waveforms, feats_pad, feats_lens
|
| 261 |
+
|
| 262 |
+
def get_fbank(self) -> Tuple[np.ndarray, np.ndarray]:
|
| 263 |
+
return self.fbanks, self.fbanks_lens
|
| 264 |
+
|
| 265 |
+
def lfr_cmvn(
|
| 266 |
+
self, input: np.ndarray, input_lengths: np.ndarray, is_final: bool = False
|
| 267 |
+
) -> Tuple[np.ndarray, np.ndarray, List[int]]:
|
| 268 |
+
batch_size = input.shape[0]
|
| 269 |
+
feats = []
|
| 270 |
+
feats_lens = []
|
| 271 |
+
lfr_splice_frame_idxs = []
|
| 272 |
+
for i in range(batch_size):
|
| 273 |
+
mat = input[i, : input_lengths[i], :]
|
| 274 |
+
lfr_splice_frame_idx = -1
|
| 275 |
+
if self.lfr_m != 1 or self.lfr_n != 1:
|
| 276 |
+
# update self.lfr_splice_cache in self.apply_lfr
|
| 277 |
+
mat, self.lfr_splice_cache[i], lfr_splice_frame_idx = self.apply_lfr(
|
| 278 |
+
mat, self.lfr_m, self.lfr_n, is_final
|
| 279 |
+
)
|
| 280 |
+
if self.cmvn_file is not None:
|
| 281 |
+
mat = self.apply_cmvn(mat)
|
| 282 |
+
feat_length = mat.shape[0]
|
| 283 |
+
feats.append(mat)
|
| 284 |
+
feats_lens.append(feat_length)
|
| 285 |
+
lfr_splice_frame_idxs.append(lfr_splice_frame_idx)
|
| 286 |
+
|
| 287 |
+
feats_lens = np.array(feats_lens)
|
| 288 |
+
feats_pad = np.array(feats)
|
| 289 |
+
return feats_pad, feats_lens, lfr_splice_frame_idxs
|
| 290 |
+
|
| 291 |
+
def extract_fbank(
|
| 292 |
+
self, input: np.ndarray, input_lengths: np.ndarray, is_final: bool = False
|
| 293 |
+
) -> Tuple[np.ndarray, np.ndarray]:
|
| 294 |
+
batch_size = input.shape[0]
|
| 295 |
+
assert (
|
| 296 |
+
batch_size == 1
|
| 297 |
+
), "we support to extract feature online only when the batch size is equal to 1 now"
|
| 298 |
+
waveforms, feats, feats_lengths = self.fbank(input, input_lengths) # input shape: B T D
|
| 299 |
+
if feats.shape[0]:
|
| 300 |
+
self.waveforms = (
|
| 301 |
+
waveforms
|
| 302 |
+
if self.reserve_waveforms is None
|
| 303 |
+
else np.concatenate((self.reserve_waveforms, waveforms), axis=1)
|
| 304 |
+
)
|
| 305 |
+
if not self.lfr_splice_cache:
|
| 306 |
+
for i in range(batch_size):
|
| 307 |
+
self.lfr_splice_cache.append(
|
| 308 |
+
np.expand_dims(feats[i][0, :], axis=0).repeat((self.lfr_m - 1) // 2, axis=0)
|
| 309 |
+
)
|
| 310 |
+
|
| 311 |
+
if feats_lengths[0] + self.lfr_splice_cache[0].shape[0] >= self.lfr_m:
|
| 312 |
+
lfr_splice_cache_np = np.stack(self.lfr_splice_cache) # B T D
|
| 313 |
+
feats = np.concatenate((lfr_splice_cache_np, feats), axis=1)
|
| 314 |
+
feats_lengths += lfr_splice_cache_np[0].shape[0]
|
| 315 |
+
frame_from_waveforms = int(
|
| 316 |
+
(self.waveforms.shape[1] - self.frame_sample_length)
|
| 317 |
+
/ self.frame_shift_sample_length
|
| 318 |
+
+ 1
|
| 319 |
+
)
|
| 320 |
+
minus_frame = (self.lfr_m - 1) // 2 if self.reserve_waveforms is None else 0
|
| 321 |
+
feats, feats_lengths, lfr_splice_frame_idxs = self.lfr_cmvn(
|
| 322 |
+
feats, feats_lengths, is_final
|
| 323 |
+
)
|
| 324 |
+
if self.lfr_m == 1:
|
| 325 |
+
self.reserve_waveforms = None
|
| 326 |
+
else:
|
| 327 |
+
reserve_frame_idx = lfr_splice_frame_idxs[0] - minus_frame
|
| 328 |
+
# print('reserve_frame_idx: ' + str(reserve_frame_idx))
|
| 329 |
+
# print('frame_frame: ' + str(frame_from_waveforms))
|
| 330 |
+
self.reserve_waveforms = self.waveforms[
|
| 331 |
+
:,
|
| 332 |
+
reserve_frame_idx
|
| 333 |
+
* self.frame_shift_sample_length : frame_from_waveforms
|
| 334 |
+
* self.frame_shift_sample_length,
|
| 335 |
+
]
|
| 336 |
+
sample_length = (
|
| 337 |
+
frame_from_waveforms - 1
|
| 338 |
+
) * self.frame_shift_sample_length + self.frame_sample_length
|
| 339 |
+
self.waveforms = self.waveforms[:, :sample_length]
|
| 340 |
+
else:
|
| 341 |
+
# update self.reserve_waveforms and self.lfr_splice_cache
|
| 342 |
+
self.reserve_waveforms = self.waveforms[
|
| 343 |
+
:, : -(self.frame_sample_length - self.frame_shift_sample_length)
|
| 344 |
+
]
|
| 345 |
+
for i in range(batch_size):
|
| 346 |
+
self.lfr_splice_cache[i] = np.concatenate(
|
| 347 |
+
(self.lfr_splice_cache[i], feats[i]), axis=0
|
| 348 |
+
)
|
| 349 |
+
return np.empty(0, dtype=np.float32), feats_lengths
|
| 350 |
+
else:
|
| 351 |
+
if is_final:
|
| 352 |
+
self.waveforms = (
|
| 353 |
+
waveforms if self.reserve_waveforms is None else self.reserve_waveforms
|
| 354 |
+
)
|
| 355 |
+
feats = np.stack(self.lfr_splice_cache)
|
| 356 |
+
feats_lengths = np.zeros(batch_size, dtype=np.int32) + feats.shape[1]
|
| 357 |
+
feats, feats_lengths, _ = self.lfr_cmvn(feats, feats_lengths, is_final)
|
| 358 |
+
if is_final:
|
| 359 |
+
self.cache_reset()
|
| 360 |
+
return feats, feats_lengths
|
| 361 |
+
|
| 362 |
+
def get_waveforms(self):
|
| 363 |
+
return self.waveforms
|
| 364 |
+
|
| 365 |
+
def cache_reset(self):
|
| 366 |
+
self.fbank_fn = knf.OnlineFbank(self.opts)
|
| 367 |
+
self.reserve_waveforms = None
|
| 368 |
+
self.input_cache = None
|
| 369 |
+
self.lfr_splice_cache = []
|
| 370 |
+
|
| 371 |
+
|
| 372 |
+
def load_bytes(input):
|
| 373 |
+
middle_data = np.frombuffer(input, dtype=np.int16)
|
| 374 |
+
middle_data = np.asarray(middle_data)
|
| 375 |
+
if middle_data.dtype.kind not in "iu":
|
| 376 |
+
raise TypeError("'middle_data' must be an array of integers")
|
| 377 |
+
dtype = np.dtype("float32")
|
| 378 |
+
if dtype.kind != "f":
|
| 379 |
+
raise TypeError("'dtype' must be a floating point type")
|
| 380 |
+
|
| 381 |
+
i = np.iinfo(middle_data.dtype)
|
| 382 |
+
abs_max = 2 ** (i.bits - 1)
|
| 383 |
+
offset = i.min + abs_max
|
| 384 |
+
array = np.frombuffer((middle_data.astype(dtype) - offset) / abs_max, dtype=np.float32)
|
| 385 |
+
return array
|
| 386 |
+
|
| 387 |
+
|
| 388 |
+
class SinusoidalPositionEncoderOnline:
|
| 389 |
+
"""Streaming Positional encoding."""
|
| 390 |
+
|
| 391 |
+
def encode(self, positions: np.ndarray = None, depth: int = None, dtype: np.dtype = np.float32):
|
| 392 |
+
batch_size = positions.shape[0]
|
| 393 |
+
positions = positions.astype(dtype)
|
| 394 |
+
log_timescale_increment = np.log(np.array([10000], dtype=dtype)) / (depth / 2 - 1)
|
| 395 |
+
inv_timescales = np.exp(np.arange(depth / 2).astype(dtype) * (-log_timescale_increment))
|
| 396 |
+
inv_timescales = np.reshape(inv_timescales, [batch_size, -1])
|
| 397 |
+
scaled_time = np.reshape(positions, [1, -1, 1]) * np.reshape(inv_timescales, [1, 1, -1])
|
| 398 |
+
encoding = np.concatenate((np.sin(scaled_time), np.cos(scaled_time)), axis=2)
|
| 399 |
+
return encoding.astype(dtype)
|
| 400 |
+
|
| 401 |
+
def forward(self, x, start_idx=0):
|
| 402 |
+
batch_size, timesteps, input_dim = x.shape
|
| 403 |
+
positions = np.arange(1, timesteps + 1 + start_idx)[None, :]
|
| 404 |
+
position_encoding = self.encode(positions, input_dim, x.dtype)
|
| 405 |
+
|
| 406 |
+
return x + position_encoding[:, start_idx : start_idx + timesteps]
|
| 407 |
+
|
| 408 |
+
|
| 409 |
+
def test():
|
| 410 |
+
path = "/nfs/zhifu.gzf/export/damo/speech_paraformer-large_asr_nat-zh-cn-16k-common-vocab8404-pytorch/example/asr_example.wav"
|
| 411 |
+
import librosa
|
| 412 |
+
|
| 413 |
+
cmvn_file = "/nfs/zhifu.gzf/export/damo/speech_paraformer-large_asr_nat-zh-cn-16k-common-vocab8404-pytorch/am.mvn"
|
| 414 |
+
config_file = "/nfs/zhifu.gzf/export/damo/speech_paraformer-large_asr_nat-zh-cn-16k-common-vocab8404-pytorch/config.yaml"
|
| 415 |
+
from funasr.runtime.python.onnxruntime.rapid_paraformer.utils.utils import read_yaml
|
| 416 |
+
|
| 417 |
+
config = read_yaml(config_file)
|
| 418 |
+
waveform, _ = librosa.load(path, sr=None)
|
| 419 |
+
frontend = WavFrontend(
|
| 420 |
+
cmvn_file=cmvn_file,
|
| 421 |
+
**config["frontend_conf"],
|
| 422 |
+
)
|
| 423 |
+
speech, _ = frontend.fbank_online(waveform) # 1d, (sample,), numpy
|
| 424 |
+
feat, feat_len = frontend.lfr_cmvn(
|
| 425 |
+
speech
|
| 426 |
+
) # 2d, (frame, 450), np.float32 -> torch, torch.from_numpy(), dtype, (1, frame, 450)
|
| 427 |
+
|
| 428 |
+
frontend.reset_status() # clear cache
|
| 429 |
+
return feat, feat_len
|
| 430 |
+
|
| 431 |
+
|
| 432 |
+
if __name__ == "__main__":
|
| 433 |
+
test()
|
ax_meeting/utils/infer_func.py
ADDED
|
@@ -0,0 +1,273 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import numpy as np
|
| 3 |
+
from typing import List, Tuple
|
| 4 |
+
from tqdm import tqdm
|
| 5 |
+
from axengine import InferenceSession
|
| 6 |
+
import os
|
| 7 |
+
import re
|
| 8 |
+
from ml_dtypes import bfloat16
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
# Discover model files automatically from model_dir.
|
| 12 |
+
# We expect files like: <prefix>_p128_l<idx>_together.axmodel and <prefix>_post.axmodel
|
| 13 |
+
# we try to detect model prefix and layer files automatically
|
| 14 |
+
def _find_axmodel_files(base_dir: str, expected_layers: int = None, expected_prefill: int = 128):
|
| 15 |
+
files = os.listdir(base_dir)
|
| 16 |
+
# match prefix, prefill size (dynamic), and layer index
|
| 17 |
+
layer_pattern = re.compile(r"^(?P<prefix>.*)_p(?P<prefill>\d+)_l(?P<idx>\d+)_together\.axmodel$")
|
| 18 |
+
post_pattern = re.compile(r"^(?P<prefix>.*)_post\.axmodel$")
|
| 19 |
+
|
| 20 |
+
# collect prefix -> [(idx, fname)]
|
| 21 |
+
prefix_map = {}
|
| 22 |
+
for fname in files:
|
| 23 |
+
m = layer_pattern.match(fname)
|
| 24 |
+
if m:
|
| 25 |
+
prefix = m.group("prefix")
|
| 26 |
+
idx = int(m.group("idx"))
|
| 27 |
+
prefix_map.setdefault(prefix, []).append((idx, fname))
|
| 28 |
+
|
| 29 |
+
if not prefix_map:
|
| 30 |
+
# fallback to hardcoded pattern if nothing detected
|
| 31 |
+
prefix = "gemma3_text"
|
| 32 |
+
layer_files = [(
|
| 33 |
+
i, f"{prefix}_p{expected_prefill}_l{i}_together.axmodel"
|
| 34 |
+
) for i in range(expected_layers or 0)]
|
| 35 |
+
else:
|
| 36 |
+
# choose the prefix with the most layers (most likely the correct one)
|
| 37 |
+
prefix = max(prefix_map.items(), key=lambda kv: len(kv[1]))[0]
|
| 38 |
+
# debug info
|
| 39 |
+
print(f"Detected prefixes: {list(prefix_map.keys())}, chosen: {prefix}, layers: {len(prefix_map[prefix])}")
|
| 40 |
+
layer_files = sorted(prefix_map[prefix], key=lambda it: it[0])
|
| 41 |
+
|
| 42 |
+
# find post process file
|
| 43 |
+
post_file = None
|
| 44 |
+
for fname in files:
|
| 45 |
+
m = post_pattern.match(fname)
|
| 46 |
+
if m and m.group("prefix") == prefix:
|
| 47 |
+
post_file = fname
|
| 48 |
+
break
|
| 49 |
+
if post_file is None:
|
| 50 |
+
candidate = os.path.join(base_dir, f"{prefix}_post.axmodel")
|
| 51 |
+
if os.path.exists(candidate):
|
| 52 |
+
post_file = f"{prefix}_post.axmodel"
|
| 53 |
+
else:
|
| 54 |
+
for fname in files:
|
| 55 |
+
if fname.endswith("_post.axmodel"):
|
| 56 |
+
post_file = fname
|
| 57 |
+
break
|
| 58 |
+
|
| 59 |
+
return layer_files, post_file, prefix
|
| 60 |
+
|
| 61 |
+
class InferManager:
|
| 62 |
+
def __init__(self, config, model_dir, max_seq_len=4095, max_prefill_len=1023):
|
| 63 |
+
|
| 64 |
+
self.config = config
|
| 65 |
+
self.max_seq_len = max_seq_len
|
| 66 |
+
self.max_prefill_len = max_prefill_len
|
| 67 |
+
|
| 68 |
+
self.sub_dim = config.hidden_size // config.num_attention_heads if not config.head_dim else config.head_dim
|
| 69 |
+
self.kv_dim = self.sub_dim * config.num_key_value_heads
|
| 70 |
+
|
| 71 |
+
self.k_caches = [
|
| 72 |
+
np.zeros((1, self.max_seq_len, self.kv_dim), dtype=bfloat16)
|
| 73 |
+
for _ in range(config.num_hidden_layers)
|
| 74 |
+
]
|
| 75 |
+
self.v_caches = [
|
| 76 |
+
np.zeros((1, self.max_seq_len, self.kv_dim), dtype=bfloat16)
|
| 77 |
+
for _ in range(config.num_hidden_layers)
|
| 78 |
+
]
|
| 79 |
+
|
| 80 |
+
layer_files, post_file, prefix = _find_axmodel_files(model_dir, config.num_hidden_layers)
|
| 81 |
+
|
| 82 |
+
self.decoder_sessions = []
|
| 83 |
+
for _, fname in tqdm(layer_files, desc="Init InferenceSession"):
|
| 84 |
+
session = InferenceSession(os.path.join(model_dir, fname))
|
| 85 |
+
self.decoder_sessions.append(session)
|
| 86 |
+
|
| 87 |
+
# post_file was returned by _find_axmodel_files; ensure it was found
|
| 88 |
+
if post_file is None:
|
| 89 |
+
raise FileNotFoundError("Cannot find post process .axmodel file in model_dir")
|
| 90 |
+
self.post_process_session = InferenceSession(os.path.join(model_dir, post_file))
|
| 91 |
+
print("Model loaded successfully!")
|
| 92 |
+
|
| 93 |
+
@staticmethod
|
| 94 |
+
def _top_p(probs: np.ndarray, p: float) -> np.ndarray:
|
| 95 |
+
sorted_indices = np.argsort(probs)
|
| 96 |
+
filtered = probs.copy()
|
| 97 |
+
cumulative = 0
|
| 98 |
+
for idx in sorted_indices[::-1]:
|
| 99 |
+
if cumulative >= p:
|
| 100 |
+
filtered[idx] = 0
|
| 101 |
+
cumulative += filtered[idx]
|
| 102 |
+
return filtered / cumulative
|
| 103 |
+
|
| 104 |
+
@staticmethod
|
| 105 |
+
def _softmax(logits: np.ndarray) -> np.ndarray:
|
| 106 |
+
logits = logits - logits.max()
|
| 107 |
+
exp_logits = np.exp(logits)
|
| 108 |
+
return (exp_logits / np.sum(exp_logits)).astype(np.float64)
|
| 109 |
+
|
| 110 |
+
def post_process(self, logits, top_k=1, top_p=0.9, temperature=0.6):
|
| 111 |
+
logits = logits.astype(np.float32).flatten()
|
| 112 |
+
candidate_indices = np.argpartition(logits, -top_k)[-top_k:]
|
| 113 |
+
candidate_logits = logits[candidate_indices] / temperature
|
| 114 |
+
candidate_probs = self._softmax(candidate_logits)
|
| 115 |
+
candidate_probs = self._top_p(candidate_probs, top_p)
|
| 116 |
+
candidate_probs = candidate_probs.astype(np.float64) / candidate_probs.sum()
|
| 117 |
+
chosen_idx = np.random.multinomial(1, candidate_probs).argmax()
|
| 118 |
+
next_token = candidate_indices[chosen_idx]
|
| 119 |
+
return next_token, candidate_indices, candidate_probs
|
| 120 |
+
|
| 121 |
+
def gen_slice_indices(self, token_len, prefill=128, expand=128):
|
| 122 |
+
remaining = max(0, token_len - prefill)
|
| 123 |
+
extra_blocks = (remaining + expand - 1) // expand
|
| 124 |
+
return list(range(extra_blocks + 1))
|
| 125 |
+
|
| 126 |
+
def prefill(
|
| 127 |
+
self,
|
| 128 |
+
tokenizer,
|
| 129 |
+
token_ids,
|
| 130 |
+
embed_data,
|
| 131 |
+
slice_len=128,
|
| 132 |
+
):
|
| 133 |
+
"""
|
| 134 |
+
Prefill step for chunked inference.
|
| 135 |
+
"""
|
| 136 |
+
seq_len = len(token_ids)
|
| 137 |
+
slice_indices = [i for i in range(min(seq_len, self.max_prefill_len) // slice_len + 1)]
|
| 138 |
+
print(f"slice_indices: {slice_indices}")
|
| 139 |
+
# total_prefill_len = (
|
| 140 |
+
# slice_len * slice_indices[-1]
|
| 141 |
+
# if slice_indices[-1] != 0
|
| 142 |
+
# else slice_len
|
| 143 |
+
# )
|
| 144 |
+
total_prefill_len = slice_len * (slice_indices[-1] + 1)
|
| 145 |
+
# slice_indices = self.gen_slice_indices(seq_len)
|
| 146 |
+
|
| 147 |
+
if total_prefill_len > 0:
|
| 148 |
+
for slice_idx in slice_indices:
|
| 149 |
+
indices = np.arange(
|
| 150 |
+
slice_idx * slice_len,
|
| 151 |
+
(slice_idx + 1) * slice_len,
|
| 152 |
+
dtype=np.uint32
|
| 153 |
+
).reshape((1, slice_len))
|
| 154 |
+
|
| 155 |
+
mask = (
|
| 156 |
+
np.zeros((1, slice_len, slice_len * (slice_idx + 1)))
|
| 157 |
+
- 65536
|
| 158 |
+
)
|
| 159 |
+
data = np.zeros((1, slice_len, self.config.hidden_size)).astype(bfloat16)
|
| 160 |
+
for i, t in enumerate(
|
| 161 |
+
range(
|
| 162 |
+
slice_idx * slice_len,
|
| 163 |
+
(slice_idx + 1) * slice_len,
|
| 164 |
+
)
|
| 165 |
+
):
|
| 166 |
+
if t < len(token_ids):
|
| 167 |
+
mask[:, i, : slice_idx * slice_len + i + 1] = 0
|
| 168 |
+
data[:, i : i + 1, :] = (
|
| 169 |
+
embed_data[t]
|
| 170 |
+
.reshape((1, 1, self.config.hidden_size))
|
| 171 |
+
.astype(bfloat16)
|
| 172 |
+
)
|
| 173 |
+
|
| 174 |
+
remain_len = (
|
| 175 |
+
seq_len - slice_idx * slice_len
|
| 176 |
+
if slice_idx == slice_indices[-1]
|
| 177 |
+
else slice_len
|
| 178 |
+
)
|
| 179 |
+
mask = mask.astype(bfloat16)
|
| 180 |
+
for layer_idx in range(self.config.num_hidden_layers):
|
| 181 |
+
input_feed = {
|
| 182 |
+
"K_cache": (
|
| 183 |
+
self.k_caches[layer_idx][:, 0 : slice_len * slice_idx, :]
|
| 184 |
+
if slice_idx
|
| 185 |
+
else np.zeros((1, 1, self.config.hidden_size), dtype=bfloat16)
|
| 186 |
+
),
|
| 187 |
+
"V_cache": (
|
| 188 |
+
self.v_caches[layer_idx][:, 0 : slice_len * slice_idx, :]
|
| 189 |
+
if slice_idx
|
| 190 |
+
else np.zeros((1, 1, self.config.hidden_size), dtype=bfloat16)
|
| 191 |
+
),
|
| 192 |
+
"indices": indices,
|
| 193 |
+
"input": data,
|
| 194 |
+
"mask": mask,
|
| 195 |
+
}
|
| 196 |
+
outputs = self.decoder_sessions[layer_idx].run(None, input_feed, shape_group=slice_idx + 1)
|
| 197 |
+
self.k_caches[layer_idx][
|
| 198 |
+
:,
|
| 199 |
+
slice_idx * slice_len : slice_idx * slice_len + remain_len,
|
| 200 |
+
:,
|
| 201 |
+
] = outputs[0][:, :remain_len, :]
|
| 202 |
+
self.v_caches[layer_idx][
|
| 203 |
+
:,
|
| 204 |
+
slice_idx * slice_len : slice_idx * slice_len + remain_len,
|
| 205 |
+
:,
|
| 206 |
+
] = outputs[1][:, :remain_len, :]
|
| 207 |
+
data = outputs[2]
|
| 208 |
+
|
| 209 |
+
print("Slice prefill done:", slice_idx)
|
| 210 |
+
|
| 211 |
+
# return data[:, :remain_len, :]
|
| 212 |
+
post_out = self.post_process_session.run(
|
| 213 |
+
None,
|
| 214 |
+
{
|
| 215 |
+
"input": data[
|
| 216 |
+
:, seq_len - (len(slice_indices) - 1) * slice_len - 1, None, :
|
| 217 |
+
]
|
| 218 |
+
}
|
| 219 |
+
)[0]
|
| 220 |
+
next_token, possible_tokens, possible_probs = self.post_process(post_out)
|
| 221 |
+
possible_decoded = [tokenizer.decode([t]) for t in possible_tokens]
|
| 222 |
+
possible_probs_str = [str((t, p)) for t, p in zip(possible_decoded, possible_probs)]
|
| 223 |
+
token_ids.append(next_token)
|
| 224 |
+
return token_ids
|
| 225 |
+
|
| 226 |
+
def decode(
|
| 227 |
+
self,
|
| 228 |
+
tokenizer,
|
| 229 |
+
token_ids,
|
| 230 |
+
embed_matrix,
|
| 231 |
+
prefill_len=128,
|
| 232 |
+
slice_len=128,
|
| 233 |
+
eos_token_id=None, # 某些模型有多个 eos_token_id
|
| 234 |
+
):
|
| 235 |
+
print("answer >>", tokenizer.decode(token_ids[-1], skip_special_tokens=True), end='', flush=True)
|
| 236 |
+
mask = np.zeros((1, 1, self.max_seq_len + 1), dtype=np.float32).astype(bfloat16)
|
| 237 |
+
mask[:, :, :self.max_seq_len] -= 65536
|
| 238 |
+
seq_len = len(token_ids) - 1
|
| 239 |
+
if prefill_len > 0:
|
| 240 |
+
mask[:, :, :seq_len] = 0
|
| 241 |
+
for step_idx in range(self.max_seq_len):
|
| 242 |
+
if prefill_len > 0 and step_idx < seq_len:
|
| 243 |
+
continue
|
| 244 |
+
cur_token = token_ids[step_idx]
|
| 245 |
+
indices = np.array([step_idx], np.uint32).reshape((1, 1))
|
| 246 |
+
data = embed_matrix[cur_token, :].reshape((1, 1, self.config.hidden_size)).astype(bfloat16)
|
| 247 |
+
for layer_idx in range(self.config.num_hidden_layers):
|
| 248 |
+
input_feed = {
|
| 249 |
+
"K_cache": self.k_caches[layer_idx],
|
| 250 |
+
"V_cache": self.v_caches[layer_idx],
|
| 251 |
+
"indices": indices,
|
| 252 |
+
"input": data,
|
| 253 |
+
"mask": mask,
|
| 254 |
+
}
|
| 255 |
+
outputs = self.decoder_sessions[layer_idx].run(None, input_feed, shape_group=0)
|
| 256 |
+
self.k_caches[layer_idx][:, step_idx, :] = outputs[0][:, :, :]
|
| 257 |
+
self.v_caches[layer_idx][:, step_idx, :] = outputs[1][:, :, :]
|
| 258 |
+
data = outputs[2]
|
| 259 |
+
mask[..., step_idx] = 0
|
| 260 |
+
if step_idx < seq_len - 1:
|
| 261 |
+
continue
|
| 262 |
+
else:
|
| 263 |
+
post_out = self.post_process_session.run(None, {"input": data})[0]
|
| 264 |
+
next_token, possible_tokens, possible_probs = self.post_process(post_out)
|
| 265 |
+
if eos_token_id is not None and next_token in eos_token_id:
|
| 266 |
+
break
|
| 267 |
+
elif next_token == tokenizer.eos_token_id:
|
| 268 |
+
break
|
| 269 |
+
else:
|
| 270 |
+
pass
|
| 271 |
+
token_ids.append(next_token)
|
| 272 |
+
print(tokenizer.decode(next_token, skip_special_tokens=True), end='', flush=True)
|
| 273 |
+
|
ax_meeting/utils/infer_utils.py
ADDED
|
@@ -0,0 +1,312 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# -*- encoding: utf-8 -*-
|
| 2 |
+
|
| 3 |
+
import functools
|
| 4 |
+
import logging
|
| 5 |
+
from pathlib import Path
|
| 6 |
+
from typing import Any, Dict, Iterable, List, NamedTuple, Set, Tuple, Union
|
| 7 |
+
|
| 8 |
+
import re
|
| 9 |
+
import numpy as np
|
| 10 |
+
import yaml
|
| 11 |
+
|
| 12 |
+
import jieba
|
| 13 |
+
import warnings
|
| 14 |
+
|
| 15 |
+
root_dir = Path(__file__).resolve().parent
|
| 16 |
+
|
| 17 |
+
logger_initialized = {}
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
def pad_list(xs, pad_value, max_len=None):
|
| 21 |
+
n_batch = len(xs)
|
| 22 |
+
if max_len is None:
|
| 23 |
+
max_len = max(x.size(0) for x in xs)
|
| 24 |
+
# pad = xs[0].new(n_batch, max_len, *xs[0].size()[1:]).fill_(pad_value)
|
| 25 |
+
# numpy format
|
| 26 |
+
pad = (np.zeros((n_batch, max_len)) + pad_value).astype(np.int32)
|
| 27 |
+
for i in range(n_batch):
|
| 28 |
+
pad[i, : xs[i].shape[0]] = xs[i]
|
| 29 |
+
|
| 30 |
+
return pad
|
| 31 |
+
|
| 32 |
+
|
| 33 |
+
"""
|
| 34 |
+
def make_pad_mask(lengths, xs=None, length_dim=-1, maxlen=None):
|
| 35 |
+
if length_dim == 0:
|
| 36 |
+
raise ValueError("length_dim cannot be 0: {}".format(length_dim))
|
| 37 |
+
|
| 38 |
+
if not isinstance(lengths, list):
|
| 39 |
+
lengths = lengths.tolist()
|
| 40 |
+
bs = int(len(lengths))
|
| 41 |
+
if maxlen is None:
|
| 42 |
+
if xs is None:
|
| 43 |
+
maxlen = int(max(lengths))
|
| 44 |
+
else:
|
| 45 |
+
maxlen = xs.size(length_dim)
|
| 46 |
+
else:
|
| 47 |
+
assert xs is None
|
| 48 |
+
assert maxlen >= int(max(lengths))
|
| 49 |
+
|
| 50 |
+
seq_range = torch.arange(0, maxlen, dtype=torch.int64)
|
| 51 |
+
seq_range_expand = seq_range.unsqueeze(0).expand(bs, maxlen)
|
| 52 |
+
seq_length_expand = seq_range_expand.new(lengths).unsqueeze(-1)
|
| 53 |
+
mask = seq_range_expand >= seq_length_expand
|
| 54 |
+
|
| 55 |
+
if xs is not None:
|
| 56 |
+
assert xs.size(0) == bs, (xs.size(0), bs)
|
| 57 |
+
|
| 58 |
+
if length_dim < 0:
|
| 59 |
+
length_dim = xs.dim() + length_dim
|
| 60 |
+
# ind = (:, None, ..., None, :, , None, ..., None)
|
| 61 |
+
ind = tuple(
|
| 62 |
+
slice(None) if i in (0, length_dim) else None for i in range(xs.dim())
|
| 63 |
+
)
|
| 64 |
+
mask = mask[ind].expand_as(xs).to(xs.device)
|
| 65 |
+
return mask
|
| 66 |
+
"""
|
| 67 |
+
|
| 68 |
+
|
| 69 |
+
class TokenIDConverter:
|
| 70 |
+
def __init__(
|
| 71 |
+
self,
|
| 72 |
+
token_list: Union[List, str],
|
| 73 |
+
):
|
| 74 |
+
|
| 75 |
+
self.token_list = token_list
|
| 76 |
+
self.unk_symbol = token_list[-1]
|
| 77 |
+
self.token2id = {v: i for i, v in enumerate(self.token_list)}
|
| 78 |
+
self.unk_id = self.token2id[self.unk_symbol]
|
| 79 |
+
|
| 80 |
+
def get_num_vocabulary_size(self) -> int:
|
| 81 |
+
return len(self.token_list)
|
| 82 |
+
|
| 83 |
+
def ids2tokens(self, integers: Union[np.ndarray, Iterable[int]]) -> List[str]:
|
| 84 |
+
if isinstance(integers, np.ndarray) and integers.ndim != 1:
|
| 85 |
+
raise TokenIDConverterError(f"Must be 1 dim ndarray, but got {integers.ndim}")
|
| 86 |
+
return [self.token_list[i] for i in integers]
|
| 87 |
+
|
| 88 |
+
def tokens2ids(self, tokens: Iterable[str]) -> List[int]:
|
| 89 |
+
|
| 90 |
+
return [self.token2id.get(i, self.unk_id) for i in tokens]
|
| 91 |
+
|
| 92 |
+
|
| 93 |
+
class CharTokenizer:
|
| 94 |
+
def __init__(
|
| 95 |
+
self,
|
| 96 |
+
symbol_value: Union[Path, str, Iterable[str]] = None,
|
| 97 |
+
space_symbol: str = "<space>",
|
| 98 |
+
remove_non_linguistic_symbols: bool = False,
|
| 99 |
+
):
|
| 100 |
+
|
| 101 |
+
self.space_symbol = space_symbol
|
| 102 |
+
self.non_linguistic_symbols = self.load_symbols(symbol_value)
|
| 103 |
+
self.remove_non_linguistic_symbols = remove_non_linguistic_symbols
|
| 104 |
+
|
| 105 |
+
@staticmethod
|
| 106 |
+
def load_symbols(value: Union[Path, str, Iterable[str]] = None) -> Set:
|
| 107 |
+
if value is None:
|
| 108 |
+
return set()
|
| 109 |
+
|
| 110 |
+
if isinstance(value, Iterable[str]):
|
| 111 |
+
return set(value)
|
| 112 |
+
|
| 113 |
+
file_path = Path(value)
|
| 114 |
+
if not file_path.exists():
|
| 115 |
+
logging.warning("%s doesn't exist.", file_path)
|
| 116 |
+
return set()
|
| 117 |
+
|
| 118 |
+
with file_path.open("r", encoding="utf-8") as f:
|
| 119 |
+
return set(line.rstrip() for line in f)
|
| 120 |
+
|
| 121 |
+
def text2tokens(self, line: Union[str, list]) -> List[str]:
|
| 122 |
+
tokens = []
|
| 123 |
+
while len(line) != 0:
|
| 124 |
+
for w in self.non_linguistic_symbols:
|
| 125 |
+
if line.startswith(w):
|
| 126 |
+
if not self.remove_non_linguistic_symbols:
|
| 127 |
+
tokens.append(line[: len(w)])
|
| 128 |
+
line = line[len(w) :]
|
| 129 |
+
break
|
| 130 |
+
else:
|
| 131 |
+
t = line[0]
|
| 132 |
+
if t == " ":
|
| 133 |
+
t = "<space>"
|
| 134 |
+
tokens.append(t)
|
| 135 |
+
line = line[1:]
|
| 136 |
+
return tokens
|
| 137 |
+
|
| 138 |
+
def tokens2text(self, tokens: Iterable[str]) -> str:
|
| 139 |
+
tokens = [t if t != self.space_symbol else " " for t in tokens]
|
| 140 |
+
return "".join(tokens)
|
| 141 |
+
|
| 142 |
+
def __repr__(self):
|
| 143 |
+
return (
|
| 144 |
+
f"{self.__class__.__name__}("
|
| 145 |
+
f'space_symbol="{self.space_symbol}"'
|
| 146 |
+
f'non_linguistic_symbols="{self.non_linguistic_symbols}"'
|
| 147 |
+
f")"
|
| 148 |
+
)
|
| 149 |
+
|
| 150 |
+
|
| 151 |
+
class Hypothesis(NamedTuple):
|
| 152 |
+
"""Hypothesis data type."""
|
| 153 |
+
|
| 154 |
+
yseq: np.ndarray
|
| 155 |
+
score: Union[float, np.ndarray] = 0
|
| 156 |
+
scores: Dict[str, Union[float, np.ndarray]] = dict()
|
| 157 |
+
states: Dict[str, Any] = dict()
|
| 158 |
+
|
| 159 |
+
def asdict(self) -> dict:
|
| 160 |
+
"""Convert data to JSON-friendly dict."""
|
| 161 |
+
return self._replace(
|
| 162 |
+
yseq=self.yseq.tolist(),
|
| 163 |
+
score=float(self.score),
|
| 164 |
+
scores={k: float(v) for k, v in self.scores.items()},
|
| 165 |
+
)._asdict()
|
| 166 |
+
|
| 167 |
+
|
| 168 |
+
class TokenIDConverterError(Exception):
|
| 169 |
+
pass
|
| 170 |
+
|
| 171 |
+
|
| 172 |
+
class ONNXRuntimeError(Exception):
|
| 173 |
+
pass
|
| 174 |
+
|
| 175 |
+
|
| 176 |
+
def split_to_mini_sentence(words: list, word_limit: int = 20):
|
| 177 |
+
assert word_limit > 1
|
| 178 |
+
if len(words) <= word_limit:
|
| 179 |
+
return [words]
|
| 180 |
+
sentences = []
|
| 181 |
+
length = len(words)
|
| 182 |
+
sentence_len = length // word_limit
|
| 183 |
+
for i in range(sentence_len):
|
| 184 |
+
sentences.append(words[i * word_limit : (i + 1) * word_limit])
|
| 185 |
+
if length % word_limit > 0:
|
| 186 |
+
sentences.append(words[sentence_len * word_limit :])
|
| 187 |
+
return sentences
|
| 188 |
+
|
| 189 |
+
|
| 190 |
+
def code_mix_split_words(text: str):
|
| 191 |
+
words = []
|
| 192 |
+
segs = text.split()
|
| 193 |
+
for seg in segs:
|
| 194 |
+
# There is no space in seg.
|
| 195 |
+
current_word = ""
|
| 196 |
+
for c in seg:
|
| 197 |
+
if len(c.encode()) == 1:
|
| 198 |
+
# This is an ASCII char.
|
| 199 |
+
current_word += c
|
| 200 |
+
else:
|
| 201 |
+
# This is a Chinese char.
|
| 202 |
+
if len(current_word) > 0:
|
| 203 |
+
words.append(current_word)
|
| 204 |
+
current_word = ""
|
| 205 |
+
words.append(c)
|
| 206 |
+
if len(current_word) > 0:
|
| 207 |
+
words.append(current_word)
|
| 208 |
+
return words
|
| 209 |
+
|
| 210 |
+
|
| 211 |
+
def isEnglish(text: str):
|
| 212 |
+
if re.search("^[a-zA-Z']+$", text):
|
| 213 |
+
return True
|
| 214 |
+
else:
|
| 215 |
+
return False
|
| 216 |
+
|
| 217 |
+
|
| 218 |
+
def join_chinese_and_english(input_list):
|
| 219 |
+
line = ""
|
| 220 |
+
for token in input_list:
|
| 221 |
+
if isEnglish(token):
|
| 222 |
+
line = line + " " + token
|
| 223 |
+
else:
|
| 224 |
+
line = line + token
|
| 225 |
+
|
| 226 |
+
line = line.strip()
|
| 227 |
+
return line
|
| 228 |
+
|
| 229 |
+
|
| 230 |
+
def code_mix_split_words_jieba(seg_dict_file: str):
|
| 231 |
+
jieba.load_userdict(seg_dict_file)
|
| 232 |
+
|
| 233 |
+
def _fn(text: str):
|
| 234 |
+
input_list = text.split()
|
| 235 |
+
token_list_all = []
|
| 236 |
+
langauge_list = []
|
| 237 |
+
token_list_tmp = []
|
| 238 |
+
language_flag = None
|
| 239 |
+
for token in input_list:
|
| 240 |
+
if isEnglish(token) and language_flag == "Chinese":
|
| 241 |
+
token_list_all.append(token_list_tmp)
|
| 242 |
+
langauge_list.append("Chinese")
|
| 243 |
+
token_list_tmp = []
|
| 244 |
+
elif not isEnglish(token) and language_flag == "English":
|
| 245 |
+
token_list_all.append(token_list_tmp)
|
| 246 |
+
langauge_list.append("English")
|
| 247 |
+
token_list_tmp = []
|
| 248 |
+
|
| 249 |
+
token_list_tmp.append(token)
|
| 250 |
+
|
| 251 |
+
if isEnglish(token):
|
| 252 |
+
language_flag = "English"
|
| 253 |
+
else:
|
| 254 |
+
language_flag = "Chinese"
|
| 255 |
+
|
| 256 |
+
if token_list_tmp:
|
| 257 |
+
token_list_all.append(token_list_tmp)
|
| 258 |
+
langauge_list.append(language_flag)
|
| 259 |
+
|
| 260 |
+
result_list = []
|
| 261 |
+
for token_list_tmp, language_flag in zip(token_list_all, langauge_list):
|
| 262 |
+
if language_flag == "English":
|
| 263 |
+
result_list.extend(token_list_tmp)
|
| 264 |
+
else:
|
| 265 |
+
seg_list = jieba.cut(join_chinese_and_english(token_list_tmp), HMM=False)
|
| 266 |
+
result_list.extend(seg_list)
|
| 267 |
+
|
| 268 |
+
return result_list
|
| 269 |
+
|
| 270 |
+
return _fn
|
| 271 |
+
|
| 272 |
+
|
| 273 |
+
def read_yaml(yaml_path: Union[str, Path]) -> Dict:
|
| 274 |
+
if not Path(yaml_path).exists():
|
| 275 |
+
raise FileExistsError(f"The {yaml_path} does not exist.")
|
| 276 |
+
|
| 277 |
+
with open(str(yaml_path), "rb") as f:
|
| 278 |
+
data = yaml.load(f, Loader=yaml.Loader)
|
| 279 |
+
return data
|
| 280 |
+
|
| 281 |
+
|
| 282 |
+
@functools.lru_cache()
|
| 283 |
+
def get_logger(name="funasr_onnx"):
|
| 284 |
+
"""Initialize and get a logger by name.
|
| 285 |
+
If the logger has not been initialized, this method will initialize the
|
| 286 |
+
logger by adding one or two handlers, otherwise the initialized logger will
|
| 287 |
+
be directly returned. During initialization, a StreamHandler will always be
|
| 288 |
+
added.
|
| 289 |
+
Args:
|
| 290 |
+
name (str): Logger name.
|
| 291 |
+
Returns:
|
| 292 |
+
logging.Logger: The expected logger.
|
| 293 |
+
"""
|
| 294 |
+
logger = logging.getLogger(name)
|
| 295 |
+
if name in logger_initialized:
|
| 296 |
+
return logger
|
| 297 |
+
|
| 298 |
+
for logger_name in logger_initialized:
|
| 299 |
+
if name.startswith(logger_name):
|
| 300 |
+
return logger
|
| 301 |
+
|
| 302 |
+
formatter = logging.Formatter(
|
| 303 |
+
"[%(asctime)s] %(name)s %(levelname)s: %(message)s", datefmt="%Y/%m/%d %H:%M:%S"
|
| 304 |
+
)
|
| 305 |
+
|
| 306 |
+
sh = logging.StreamHandler()
|
| 307 |
+
sh.setFormatter(formatter)
|
| 308 |
+
logger.addHandler(sh)
|
| 309 |
+
logger_initialized[name] = True
|
| 310 |
+
logger.propagate = False
|
| 311 |
+
logging.basicConfig(level=logging.ERROR)
|
| 312 |
+
return logger
|
ax_meeting/utils/sentencepiece_tokenizer.py
ADDED
|
@@ -0,0 +1,46 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# -*- coding: utf-8 -*-
|
| 2 |
+
from pathlib import Path
|
| 3 |
+
from typing import Iterable, List, Union
|
| 4 |
+
|
| 5 |
+
import sentencepiece as spm
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
class SentencepiecesTokenizer:
|
| 9 |
+
def __init__(self, bpemodel: Union[Path, str], **kwargs):
|
| 10 |
+
self.bpemodel = str(bpemodel)
|
| 11 |
+
self.sp = None
|
| 12 |
+
self._build()
|
| 13 |
+
|
| 14 |
+
def __repr__(self):
|
| 15 |
+
return f'{self.__class__.__name__}(model="{self.bpemodel}")'
|
| 16 |
+
|
| 17 |
+
def _build(self):
|
| 18 |
+
if self.sp is None:
|
| 19 |
+
self.sp = spm.SentencePieceProcessor()
|
| 20 |
+
self.sp.load(self.bpemodel)
|
| 21 |
+
|
| 22 |
+
def text2tokens(self, line: str) -> List[str]:
|
| 23 |
+
self._build()
|
| 24 |
+
return self.sp.EncodeAsPieces(line)
|
| 25 |
+
|
| 26 |
+
def tokens2text(self, tokens: Iterable[str]) -> str:
|
| 27 |
+
self._build()
|
| 28 |
+
return self.sp.DecodePieces(list(tokens))
|
| 29 |
+
|
| 30 |
+
def encode(self, line: str) -> List[int]:
|
| 31 |
+
self._build()
|
| 32 |
+
return self.sp.EncodeAsIds(line)
|
| 33 |
+
|
| 34 |
+
def decode(self, line: List[int]):
|
| 35 |
+
self._build()
|
| 36 |
+
return self.sp.DecodeIds(line)
|
| 37 |
+
|
| 38 |
+
def get_vocab_size(self):
|
| 39 |
+
self._build()
|
| 40 |
+
return self.sp.GetPieceSize()
|
| 41 |
+
|
| 42 |
+
def ids2tokens(self, *args, **kwargs):
|
| 43 |
+
return self.decode(*args, **kwargs)
|
| 44 |
+
|
| 45 |
+
def tokens2ids(self, *args, **kwargs):
|
| 46 |
+
return self.encode(*args, **kwargs)
|
ax_meeting/utils/speaker_fbank.py
ADDED
|
@@ -0,0 +1,22 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# -*- coding: utf-8 -*-
|
| 2 |
+
import numpy as np
|
| 3 |
+
import kaldi_native_fbank as knf
|
| 4 |
+
|
| 5 |
+
|
| 6 |
+
def compute_fbank(wav: np.ndarray, sample_rate: int, n_mels: int = 80, mean_nor: bool = True) -> np.ndarray:
|
| 7 |
+
if wav.ndim != 1:
|
| 8 |
+
wav = wav.reshape(-1)
|
| 9 |
+
opts = knf.FbankOptions()
|
| 10 |
+
opts.frame_opts.samp_freq = sample_rate
|
| 11 |
+
opts.mel_opts.num_bins = n_mels
|
| 12 |
+
opts.frame_opts.dither = 0.0
|
| 13 |
+
fbank = knf.OnlineFbank(opts)
|
| 14 |
+
fbank.accept_waveform(sample_rate, wav.astype(np.float32))
|
| 15 |
+
fbank.input_finished()
|
| 16 |
+
num_frames = fbank.num_frames_ready
|
| 17 |
+
if num_frames == 0:
|
| 18 |
+
return np.zeros((0, n_mels), dtype=np.float32)
|
| 19 |
+
feats = np.stack([fbank.get_frame(i) for i in range(num_frames)]).astype(np.float32)
|
| 20 |
+
if mean_nor and feats.size > 0:
|
| 21 |
+
feats = feats - feats.mean(axis=0, keepdims=True)
|
| 22 |
+
return feats
|