wangli commited on
Commit
1494b1b
·
verified ·
1 Parent(s): 9b44474

Upload folder using huggingface_hub

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. .gitattributes +1 -0
  2. .gitignore +2 -1
  3. MANIFEST.in +5 -0
  4. README.md +55 -7
  5. ax_meeting/__init__.py +5 -0
  6. ax_meeting/ax_model/.gitattributes +2 -0
  7. ax_meeting/ax_model/auto.npy +3 -0
  8. ax_meeting/ax_model/campplus.axmodel +3 -0
  9. ax_meeting/ax_model/chn_jpn_yue_eng_ko_spectok.bpe.model +3 -0
  10. ax_meeting/ax_model/en.npy +3 -0
  11. ax_meeting/ax_model/event_emo.npy +3 -0
  12. ax_meeting/ax_model/ja.npy +3 -0
  13. ax_meeting/ax_model/ko.npy +3 -0
  14. ax_meeting/ax_model/sensevoice.axmodel +3 -0
  15. ax_meeting/ax_model/sensevoice/am.mvn +8 -0
  16. ax_meeting/ax_model/sensevoice/config.yaml +97 -0
  17. ax_meeting/ax_model/vad.axmodel +3 -0
  18. ax_meeting/ax_model/vad/am.mvn +8 -0
  19. ax_meeting/ax_model/vad/config.yaml +56 -0
  20. ax_meeting/ax_model/withitn.npy +3 -0
  21. ax_meeting/ax_model/yue.npy +3 -0
  22. ax_meeting/ax_model/zh.npy +3 -0
  23. ax_meeting/axengine_loader.py +34 -0
  24. ax_meeting/certs/cert.pem +19 -0
  25. ax_meeting/certs/key.pem +28 -0
  26. ax_meeting/config.py +19 -0
  27. ax_meeting/diar_asr_cli.py +93 -0
  28. ax_meeting/diar_utils.py +13 -0
  29. ax_meeting/engines.py +207 -0
  30. ax_meeting/model_bundle.py +87 -0
  31. ax_meeting/pipeline.py +163 -0
  32. ax_meeting/positional.py +18 -0
  33. ax_meeting/server.py +135 -0
  34. ax_meeting/static/app.js +187 -0
  35. ax_meeting/static/index.html +50 -0
  36. ax_meeting/static/style.css +147 -0
  37. ax_meeting/summarize_cli.py +33 -0
  38. ax_meeting/summarizer.py +75 -0
  39. ax_meeting/text_cleaner.py +41 -0
  40. ax_meeting/utils/__init__.py +0 -0
  41. ax_meeting/utils/ax_cam_bin.py +231 -0
  42. ax_meeting/utils/ax_model_bin.py +307 -0
  43. ax_meeting/utils/ax_vad_bin.py +158 -0
  44. ax_meeting/utils/cluster_utils.py +241 -0
  45. ax_meeting/utils/ctc_alignment.py +76 -0
  46. ax_meeting/utils/frontend.py +433 -0
  47. ax_meeting/utils/infer_func.py +273 -0
  48. ax_meeting/utils/infer_utils.py +312 -0
  49. ax_meeting/utils/sentencepiece_tokenizer.py +46 -0
  50. 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 app.server
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 app.server
 
 
 
 
 
 
100
  ```
101
 
 
 
 
 
 
 
 
 
102
  ![meeting_demo.png](assert/meeting_demo.png)
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 app/cli_batch.py --wav_file wav/vad_example.wav --output_dir output_dir
111
  ```
112
 
113
- 带会议总结(可选参数覆盖 LLM 配置):
114
 
115
  ```bash
116
- python app/cli_batch.py \\
117
  --wav_file wav/vad_example.wav \\
118
- --output_dir output_dir \\
119
- --summary \\
 
 
 
 
 
 
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
  ![meeting_demo.png](assert/meeting_demo.png)
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