pangkaiyu commited on
Commit
3d4f331
·
verified ·
1 Parent(s): 70a552a

Add files using upload-large-folder tool

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 +34 -0
  2. almeval/models/glm4voice/resources/architecture.jpeg +3 -0
  3. almeval/models/glm4voice/resources/web_demo.png +3 -0
  4. almeval/models/kimi_audio/assets/kimia_framework.png +3 -0
  5. almeval/models/kimi_audio/assets/kimia_logo.png +3 -0
  6. almeval/models/kimi_audio/assets/kimia_radar_chart.png +3 -0
  7. almeval/models/kimi_audio/assets/kimia_report.pdf +3 -0
  8. almeval/models/kimi_audio/kimia_infer/models/tokenizer/glm4/resources/architecture.jpeg +3 -0
  9. almeval/models/kimi_audio/kimia_infer/models/tokenizer/glm4/resources/web_demo.png +3 -0
  10. almeval/models/kimi_audio/kimia_infer/models/tokenizer/whisper_Lv3/mel_filters.npz +3 -0
  11. almeval/models/kimi_audio/test_audios/asr_example.wav +3 -0
  12. almeval/models/kimi_audio/test_audios/qa_example.wav +3 -0
  13. almeval/models/stepaudio/Dockerfile-vllm +35 -0
  14. almeval/models/stepaudio/README_JP.md +674 -0
  15. almeval/models/stepaudio/assets/Step-Audio.pdf +3 -0
  16. almeval/models/stepaudio/assets/architecture.png +3 -0
  17. almeval/models/stepaudio/assets/logo.png +0 -0
  18. almeval/models/stepaudio/assets/pipeline.png +3 -0
  19. almeval/models/stepaudio/assets/rlhf.png +3 -0
  20. almeval/models/stepaudio/assets/stepeval_radar_chart.png +3 -0
  21. almeval/models/stepaudio/assets/yuewen.jpeg +0 -0
  22. almeval/models/stepaudio/examples/clone_wav_lixueqin.wav +3 -0
  23. almeval/models/stepaudio/examples/clone_wav_yuqian.wav +3 -0
  24. almeval/models/stepaudio/examples/emotional_control1.wav +3 -0
  25. almeval/models/stepaudio/examples/emotional_control2.wav +3 -0
  26. almeval/models/stepaudio/examples/multilingual1.wav +3 -0
  27. almeval/models/stepaudio/examples/multilingual2.wav +3 -0
  28. almeval/models/stepaudio/examples/multilingual_singing.wav +3 -0
  29. almeval/models/stepaudio/examples/prompt_wav_lixueqin.wav +3 -0
  30. almeval/models/stepaudio/examples/prompt_wav_yuqian.wav +3 -0
  31. almeval/models/stepaudio/examples/prompt_wav_zhaobenshan.wav +3 -0
  32. almeval/models/stepaudio/examples/rap.wav +3 -0
  33. almeval/models/stepaudio/examples/singing.wav +3 -0
  34. almeval/models/stepaudio/examples/speed_control1.wav +3 -0
  35. almeval/models/stepaudio/examples/speed_control2.wav +3 -0
  36. almeval/models/stepaudio/examples/tone_control.wav +3 -0
  37. almeval/models/stepaudio/funasr_detach/__init__.py +38 -0
  38. almeval/models/stepaudio/funasr_detach/auto/__init__.py +0 -0
  39. almeval/models/stepaudio/funasr_detach/auto/auto_frontend.py +90 -0
  40. almeval/models/stepaudio/funasr_detach/auto/auto_model.py +573 -0
  41. almeval/models/stepaudio/funasr_detach/auto/auto_tokenizer.py +7 -0
  42. almeval/models/stepaudio/funasr_detach/bin/__init__.py +0 -0
  43. almeval/models/stepaudio/funasr_detach/bin/compute_audio_cmvn.py +152 -0
  44. almeval/models/stepaudio/funasr_detach/bin/inference.py +33 -0
  45. almeval/models/stepaudio/funasr_detach/bin/tokenize_text.py +281 -0
  46. almeval/models/stepaudio/funasr_detach/bin/train.py +227 -0
  47. almeval/models/stepaudio/funasr_detach/datasets/__init__.py +0 -0
  48. almeval/models/stepaudio/funasr_detach/datasets/audio_datasets/__init__.py +0 -0
  49. almeval/models/stepaudio/funasr_detach/datasets/audio_datasets/datasets.py +112 -0
  50. almeval/models/stepaudio/funasr_detach/datasets/audio_datasets/index_ds.py +150 -0
.gitattributes CHANGED
@@ -33,3 +33,37 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
36
+ almeval/models/glm4voice/resources/web_demo.png filter=lfs diff=lfs merge=lfs -text
37
+ almeval/models/glm4voice/resources/architecture.jpeg filter=lfs diff=lfs merge=lfs -text
38
+ almeval/models/kimi_audio/test_audios/qa_example.wav filter=lfs diff=lfs merge=lfs -text
39
+ almeval/models/kimi_audio/test_audios/asr_example.wav filter=lfs diff=lfs merge=lfs -text
40
+ almeval/models/kimi_audio/assets/kimia_radar_chart.png filter=lfs diff=lfs merge=lfs -text
41
+ almeval/models/kimi_audio/assets/kimia_framework.png filter=lfs diff=lfs merge=lfs -text
42
+ almeval/models/kimi_audio/assets/kimia_logo.png filter=lfs diff=lfs merge=lfs -text
43
+ almeval/models/kimi_audio/assets/kimia_report.pdf filter=lfs diff=lfs merge=lfs -text
44
+ almeval/models/kimi_audio/kimia_infer/models/tokenizer/glm4/resources/architecture.jpeg filter=lfs diff=lfs merge=lfs -text
45
+ almeval/models/kimi_audio/kimia_infer/models/tokenizer/glm4/resources/web_demo.png filter=lfs diff=lfs merge=lfs -text
46
+ almeval/models/stepaudio/assets/architecture.png filter=lfs diff=lfs merge=lfs -text
47
+ almeval/models/stepaudio/assets/rlhf.png filter=lfs diff=lfs merge=lfs -text
48
+ almeval/models/stepaudio/assets/pipeline.png filter=lfs diff=lfs merge=lfs -text
49
+ almeval/models/stepaudio/assets/Step-Audio.pdf filter=lfs diff=lfs merge=lfs -text
50
+ almeval/models/stepaudio/assets/stepeval_radar_chart.png filter=lfs diff=lfs merge=lfs -text
51
+ almeval/models/stepaudio/speakers/TingtingRAP_prompt.wav filter=lfs diff=lfs merge=lfs -text
52
+ almeval/models/stepaudio/speakers/Tingting_prompt.wav filter=lfs diff=lfs merge=lfs -text
53
+ almeval/models/stepaudio/speakers/Tingting哼唱_prompt.wav filter=lfs diff=lfs merge=lfs -text
54
+ almeval/models/stepaudio/examples/clone_wav_yuqian.wav filter=lfs diff=lfs merge=lfs -text
55
+ almeval/models/stepaudio/examples/multilingual2.wav filter=lfs diff=lfs merge=lfs -text
56
+ almeval/models/stepaudio/examples/multilingual_singing.wav filter=lfs diff=lfs merge=lfs -text
57
+ almeval/models/stepaudio/examples/prompt_wav_lixueqin.wav filter=lfs diff=lfs merge=lfs -text
58
+ almeval/models/stepaudio/examples/prompt_wav_zhaobenshan.wav filter=lfs diff=lfs merge=lfs -text
59
+ almeval/models/stepaudio/examples/singing.wav filter=lfs diff=lfs merge=lfs -text
60
+ almeval/models/stepaudio/examples/clone_wav_lixueqin.wav filter=lfs diff=lfs merge=lfs -text
61
+ almeval/models/stepaudio/examples/speed_control1.wav filter=lfs diff=lfs merge=lfs -text
62
+ almeval/models/stepaudio/examples/emotional_control1.wav filter=lfs diff=lfs merge=lfs -text
63
+ almeval/models/stepaudio/examples/emotional_control2.wav filter=lfs diff=lfs merge=lfs -text
64
+ almeval/models/stepaudio/examples/multilingual1.wav filter=lfs diff=lfs merge=lfs -text
65
+ almeval/models/stepaudio/examples/prompt_wav_yuqian.wav filter=lfs diff=lfs merge=lfs -text
66
+ almeval/models/stepaudio/examples/rap.wav filter=lfs diff=lfs merge=lfs -text
67
+ almeval/models/stepaudio/examples/speed_control2.wav filter=lfs diff=lfs merge=lfs -text
68
+ almeval/models/stepaudio/examples/tone_control.wav filter=lfs diff=lfs merge=lfs -text
69
+ tests/data/asr_example.wav filter=lfs diff=lfs merge=lfs -text
almeval/models/glm4voice/resources/architecture.jpeg ADDED

Git LFS Details

  • SHA256: a1f8a08e4fdedb82a776ed949d1ff5b48f04f9996681b5646ac3b939dda88eb9
  • Pointer size: 131 Bytes
  • Size of remote file: 143 kB
almeval/models/glm4voice/resources/web_demo.png ADDED

Git LFS Details

  • SHA256: 565cfa8fe6ee3998e6425c252bf646a68aae5b8d9957e0b51733b1001138b346
  • Pointer size: 131 Bytes
  • Size of remote file: 143 kB
almeval/models/kimi_audio/assets/kimia_framework.png ADDED

Git LFS Details

  • SHA256: 88e34c7745dd95c77f7020cfcd42810deaddc8e57dd257585126b613b8e2ddcd
  • Pointer size: 131 Bytes
  • Size of remote file: 599 kB
almeval/models/kimi_audio/assets/kimia_logo.png ADDED

Git LFS Details

  • SHA256: fcc60532e50061e8d3d65a519178306e864b9cc5de7ccee41d4d892e9ce3fb23
  • Pointer size: 131 Bytes
  • Size of remote file: 284 kB
almeval/models/kimi_audio/assets/kimia_radar_chart.png ADDED

Git LFS Details

  • SHA256: 6d2b32e9ccd6b877e289a5d45e1c912bbb0ee4b7ccf439700106aa45c4e4f46b
  • Pointer size: 132 Bytes
  • Size of remote file: 2.39 MB
almeval/models/kimi_audio/assets/kimia_report.pdf ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:8928bbfc59d21811a2ecdbaabe277763973aa3ff761682efb8c614bc6fc77666
3
+ size 2433108
almeval/models/kimi_audio/kimia_infer/models/tokenizer/glm4/resources/architecture.jpeg ADDED

Git LFS Details

  • SHA256: a1f8a08e4fdedb82a776ed949d1ff5b48f04f9996681b5646ac3b939dda88eb9
  • Pointer size: 131 Bytes
  • Size of remote file: 143 kB
almeval/models/kimi_audio/kimia_infer/models/tokenizer/glm4/resources/web_demo.png ADDED

Git LFS Details

  • SHA256: 565cfa8fe6ee3998e6425c252bf646a68aae5b8d9957e0b51733b1001138b346
  • Pointer size: 131 Bytes
  • Size of remote file: 143 kB
almeval/models/kimi_audio/kimia_infer/models/tokenizer/whisper_Lv3/mel_filters.npz ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:7450ae70723a5ef9d341e3cee628c7cb0177f36ce42c44b7ed2bf3325f0f6d4c
3
+ size 4271
almeval/models/kimi_audio/test_audios/asr_example.wav ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:5214e16584b6498ddbd22f321738c0f62b19117ab1880b21e405fba983986b6a
3
+ size 332204
almeval/models/kimi_audio/test_audios/qa_example.wav ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:9518e276b1d7a1e64fa03138524e515719ab98adcbb82b9b201c8d6db145f26c
3
+ size 118146
almeval/models/stepaudio/Dockerfile-vllm ADDED
@@ -0,0 +1,35 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ FROM nvidia/cuda:12.4.1-cudnn-runtime-ubuntu22.04
2
+
3
+ ENV TZ=Asia/Shanghai
4
+ RUN ln -snf /usr/share/zoneinfo/$TZ /etc/localtime \
5
+ && echo $TZ > /etc/timezone
6
+
7
+ RUN apt-get update \
8
+ && apt-get install -y build-essential \
9
+ && apt-get install -y wget \
10
+ && apt-get install -y software-properties-common curl zip unzip git-lfs awscli libssl-dev openssh-server vim \
11
+ && apt-get install -y net-tools iputils-ping iproute2 \
12
+ && apt-get clean \
13
+ && rm -rf /var/lib/apt/lists/*
14
+
15
+ RUN apt-get install --reinstall ca-certificates && update-ca-certificates \
16
+ && apt-get clean \
17
+ && rm -rf /var/lib/apt/lists/*
18
+
19
+ RUN add-apt-repository -y 'ppa:deadsnakes/ppa' && apt update
20
+ RUN apt install python3.10 python3.10-dev python3.10-distutils python3.10-venv -y \
21
+ && apt-get clean \
22
+ && rm -rf /var/lib/apt/lists/*
23
+
24
+ RUN wget -qO- https://bootstrap.pypa.io/get-pip.py | python3.10
25
+ RUN ln -s /usr/bin/python3.10 /usr/bin/python
26
+ RUN pip uninstall -y Pillow && pip install pillow
27
+
28
+ COPY requirements-vllm.txt /tmp/requirements.txt
29
+ RUN pip3 install -r /tmp/requirements.txt
30
+ # update vllm
31
+ RUN VLLM_PYTHON_DIR=$(pip3 show vllm | grep Location | awk '{print $2}')/vllm \
32
+ && git clone -b add-step1-model https://github.com/stepfun-ai/vllm.git /tmp/vllm \
33
+ && cd /tmp/vllm/vllm \
34
+ && find . -name '*.py' -exec cp -v --parents -t $VLLM_PYTHON_DIR {} + \
35
+ && rm -rf /tmp/vllm
almeval/models/stepaudio/README_JP.md ADDED
@@ -0,0 +1,674 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ <p align="left">
2
+ <a href="README_CN.md">中文</a> &nbsp | &nbsp<a href="README.md">English</a>&nbsp | &nbsp 日本語&nbsp
3
+ </p>
4
+ <br><br>
5
+
6
+ # Step-Audio
7
+ <p align="center">
8
+ <img src="assets/logo.png" height=100>
9
+ </p>
10
+ <div align="center">
11
+ <a href="https://arxiv.org/abs/2502.11946"><img src="https://img.shields.io/static/v1?label=Tech Report&message=Arxiv&color=red"></a> &ensp;
12
+ <a href="https://x.com/StepFun_ai"><img src="https://img.shields.io/static/v1?label=X.com&message=Web&color=blue"></a> &ensp;
13
+ </div>
14
+
15
+ <div align="center">
16
+ <a href="https://huggingface.co/stepfun-ai/Step-Audio-Chat"><img src="https://img.shields.io/static/v1?label=Step-Audio-Chat&message=HuggingFace&color=yellow"></a> &ensp;
17
+ <a href="https://huggingface.co/stepfun-ai/Step-Audio-TTS-3B"><img src="https://img.shields.io/static/v1?label=Step-Audio-TTS-3B&message=HuggingFace&color=yellow"></a> &ensp;
18
+ </div>
19
+ <div align="center">
20
+ <a href="https://huggingface.co/stepfun-ai/Step-Audio-Tokenizer"><img src="https://img.shields.io/static/v1?label=Step-Audio-Tokenier&message=HuggingFace&color=yellow"></a> &ensp;
21
+ <a href="https://huggingface.co/datasets/stepfun-ai/StepEval-Audio-360"><img src="https://img.shields.io/static/v1?label=StepEval-Audio-360&message=HuggingFace&color=yellow"></a> &ensp;
22
+ </div>
23
+
24
+ ## 🔥🔥🔥 ニュース!!
25
+ * 2025年2月17日: 👋 推論コードとモデルの重みをリリースしました。[Step-Audio-Chat](https://huggingface.co/stepfun-ai/Step-Audio-Chat), [Step-Audio-TTS-3B](https://huggingface.co/stepfun-ai/Step-Audio-TTS-3B) および [Step-Audio-Tokenizer](https://huggingface.co/stepfun-ai/Step-Audio-Tokenizer)。
26
+ * 2025年2月17日: 👋 マルチターンオーディオベンチマーク [StepEval-Audio-360](https://huggingface.co/datasets/stepfun-ai/StepEval-Audio-360) をリリースしました。
27
+ * 2025年2月17日: 👋 技術レポート [Step-Audio-Report](https://arxiv.org/abs/2502.11946) をリリースしました。
28
+
29
+ ## 目次
30
+
31
+ 1. [紹介](#1-紹介)
32
+ 2. [モデル概要](#2-モデル概要)
33
+ 3. [モデルのダウンロード](#3-モデルのダウンロード)
34
+ 4. [モデルの使用方法](#4-モデルの使用方法)
35
+ 5. [ベンチマーク](#5-ベンチマーク)
36
+ 6. [オンラインエンジン](#6-オンラインエンジン)
37
+ 7. [引用](#7-引用)
38
+
39
+ ## 1. 紹介
40
+
41
+ Step-Audioは、音声理解と生成を統合した業界初の製品レベルのオープンソースリアルタイム音声対話システムであり、多言語対話(例:日本語、英語、中国語)、音声感情(例:喜び、悲しみ)、方言(例:関西弁、広東語)、音声速度および韻律スタイルの調整をサポートします。Step-Audioは、以下の4つの主要な技術革新を示しています:
42
+
43
+ - **1300億パラメータのマルチモーダルモデル**:単一の統合モデルで、音声認識、意味理解、対話、音声クローン、音声生成を実行します。1300億パラメータのStep-Audio-Chatバリアントをオープンソース化しました。
44
+
45
+ - **生成データエンジン**:従来のTTSが手動データ収集に依存することを排除し、1300億パラメータのマルチモーダルモデルを使用して高品質の音声を生成します。このデータを活用して、リソース効率の高いStep-Audio-TTS-3Bモデルをトレーニングし、制御可能な音声合成のための指示フォロー機能を強化しました。
46
+
47
+ - **細かい音声制御**:指示ベースの制御設計を通じて、複数の感情(怒り、喜び、悲しみ)、方言(関西弁、広東語など)、および音声スタイル(ラップ、アカペラハミング)をサポートし、多様な音声生成ニーズに対応します。
48
+
49
+ - **強化されたインテリジェンス**:ToolCallメカニズムの統合とロールプレイングの強化を通じて、エージェントの複雑なタスクにおけるパフォーマンスを向上させます。
50
+
51
+ ## 2. モデル概要
52
+ Step-Audioでは、音声ストリームをトークン化するために、並列のセマンティック(16.7Hz、1024エントリのコードブック)および音響(25Hz、4096エントリのコードブック)トークナイザーを組み合わせたデュアルコードブックフレームワークを使用し、2:3の時間的インターリーブを行います。1300億パラメータのLLM基盤(Step-1)は、音声コンテキスト化継続的事前トレーニングおよびタスク固有の後トレーニングを通じて強化され、強力なクロスモーダル音声理解を実現します。フローマッチングとニューラルボコーダを組み合わせたハイブリッド音声デコーダを使用し、リアルタイムの波形生成を最適化します。推論パイプラインは、投機的応答生成(40%のコミット率)およびテキストベースのコンテキスト管理(14:1の圧縮率)を特���とするストリーミング対応アーキテクチャを備えています。
53
+ ![Architecture](assets/architecture.png)
54
+
55
+ ### 2.1 トークナイザー
56
+
57
+ セマンティックトークナイザーと音響トークナイザーを効果的に統合するために、トークンレベルのインターリーブアプローチを実装しています。セマンティックトークナイザーは1024のコードブックサイズを使用し、音響トークナイザーはより大きな4096のコードブックサイズを使用して、より細かい音響の詳細をキャプチャします。異なるトークンレートを考慮して、2つのセマンティックトークンごとに3つの音響トークンをペアリングする2:3の時間的アライメント比を確立します。
58
+
59
+ ### 2.2 言語モデル
60
+
61
+ Step-Audioの音声情報を効果的に処理し、正確な音声-テキストアライメントを実現するために、1300億パラメータの事前トレーニングされたテキストベースの大規模言語モデル(LLM)であるStep-1に基づいて、音声継続的事前トレーニングを実施しました。
62
+
63
+ ### 2.3 音声デコーダ
64
+ Step-Audioの音声デコーダは、セマンティックおよび音響情報を含む離散音声トークンを、自然な音声を表す連続的な時間領域の波形に変換する重要な機能を果たします。デコーダアーキテクチャには、フローマッチングモデルとメルから波形へのボコーダが組み込まれています。生成された音声の明瞭度と自然さを最適化するために、音声デコーダはデュアルコードインターリーブアプローチを使用してトレーニングされ、生成プロセス全体でセマンティックおよび音響機能のシームレスな統合を確保します。
65
+
66
+ ### 2.4 リアルタイム推論パイプライン
67
+ リアルタイムの対話を可能にするために、最適化された推論パイプラインを設計しました。その中心には、状態遷移を管理し、投機的応答生成を調整し、重要なサブシステム間のシームレスな調整を確保するコントローラーモジュールがあります。これらのサブシステムには、ユーザーの音声を検出する音声活動検出(VAD)、リアルタイムで音声を処理するストリーミングオーディオトークナイザー、応答を処理および生成するStep-Audio言語モデルおよび音声デコーダ、および会話の連続性を維持するコンテキストマネージャが含まれます。
68
+ ![Inference Pipeline](assets/pipeline.png)
69
+
70
+ ### 2.5 後トレーニングの詳細
71
+ 後トレーニングフェーズでは、自動音声認識(ASR)およびテキストから音声への変換(TTS)のタスク固有の監督付き微調整(SFT)を実施しました。音声入力テキスト出力(AQTA)タスクについては、多様な高品質データセットを使用してSFTを実施し、人間のフィードバックからの強化学習(RLHF)を組み合わせて応答品質を向上させ、感情表現、音声速度、方言、および韻律の細かい制御を可能にしました。
72
+ ![RLHF](assets/rlhf.png)
73
+
74
+
75
+ ## 3. モデルのダウンロード
76
+ ### 3.1 Huggingface
77
+ | モデル | リンク |
78
+ |-------|-------|
79
+ | Step-Audio-Tokenizer | [🤗huggingface](https://huggingface.co/stepfun-ai/Step-Audio-Tokenizer) |
80
+ | Step-Audio-Chat | [🤗huggingface](https://huggingface.co/stepfun-ai/Step-Audio-Chat) |
81
+ | Step-Audio-TTS-3B | [🤗huggingface](https://huggingface.co/stepfun-ai/Step-Audio-TTS-3B) |
82
+
83
+ ### 3.2 Modelscope
84
+ | モデル | リンク |
85
+ |-------|-------|
86
+ | Step-Audio-Tokenizer | [modelscope](https://modelscope.cn/models/stepfun-ai/Step-Audio-Tokenizer) |
87
+ | Step-Audio-Chat | [modelscope](https://modelscope.cn/models/stepfun-ai/Step-Audio-Chat) |
88
+ | Step-Audio-TTS-3B | [modelscope](https://modelscope.cn/models/stepfun-ai/Step-Audio-TTS-3B) |
89
+
90
+ ## 4. モデルの使用方法
91
+ ### 📜 4.1 要件
92
+ 次の表は、Step-Audioモデル(バッチサイズ=1)を実行するための要件を示しています:
93
+
94
+ | モデル | 設定<br/>(サンプル周波数) | GPU最小メモリ |
95
+ |------------|--------------------------------|----------------|
96
+ | Step-Audio-Tokenizer | 41.6Hz | 1.5GB |
97
+ | Step-Audio-Chat | 41.6Hz | 265GB |
98
+ | Step-Audio-TTS-3B | 41.6Hz | 8GB |
99
+
100
+ * CUDAサポートのあるNVIDIA GPUが必要です。
101
+ * モデルは、4つのA800 80G GPUでテストされています。
102
+ * **推奨**:より良い生成品質のために、80GBメモリを持つ4つのA800/H800 GPUを使用することをお勧めします。
103
+ * テストされたオペレーティングシステム:Linux
104
+
105
+ ### 🔧 4.2 依存関係とインストール
106
+ - Python >= 3.10.0([Anaconda](https://www.anaconda.com/download/#linux)または[Miniconda](https://docs.conda.io/en/latest/miniconda.html)の使用を推奨)
107
+ - [PyTorch >= 2.3-cu121](https://pytorch.org/)
108
+ - [CUDA Toolkit](https://developer.nvidia.com/cuda-downloads)
109
+
110
+ ```bash
111
+ git clone https://github.com/stepfun-ai/Step-Audio.git
112
+ conda create -n stepaudio python=3.10
113
+ conda activate stepaudio
114
+
115
+ cd Step-Audio
116
+ pip install -r requirements.txt
117
+
118
+ git lfs install
119
+ git clone https://huggingface.co/stepfun-ai/Step-Audio-Tokenizer
120
+ git clone https://huggingface.co/stepfun-ai/Step-Audio-Chat
121
+ git clone https://huggingface.co/stepfun-ai/Step-Audio-TTS-3B
122
+
123
+ ```
124
+
125
+ モデルをダウンロードした後、where_you_download_dirは次の構造を持つ必要があります:
126
+ ```
127
+ where_you_download_dir
128
+ ├── Step-Audio-Tokenizer
129
+ ├── Step-Audio-Chat
130
+ ├── Step-Audio-TTS-3B
131
+ ```
132
+
133
+ #### Docker 実行環境
134
+
135
+ dockerを使用してStep-Audioの実行に必要な環境を作成します
136
+
137
+ ```bash
138
+ # Dockerイメージのビルド
139
+ docker build . -t step-audio
140
+
141
+ # Dockerコンテナの実行
142
+ docker run --rm -ti --gpus all \
143
+ -v /your/code/path:/app -v /your/model/path:/model \
144
+ -p 7860:7860 \
145
+ step-audio \
146
+ -- bash
147
+
148
+ # vLLM Dockerイメージのビルド
149
+ docker build -f Dockerfile-vllm -t step-audio-vllm .
150
+
151
+ # vLLM Dockerコンテナの実行
152
+ docker run --rm -ti --gpus all \
153
+ -v /your/code/path:/app -v /your/model/path:/model \
154
+ -p 7860:7860 \
155
+ -p 8000:8000 \
156
+ step-audio-vllm \
157
+ -- bash
158
+ ```
159
+
160
+
161
+ ### 🚀 4.3 推論スクリプト
162
+ #### オフライン推論
163
+ エンドツーエンドの音声/テキスト入力と音声/テキスト出力で推論を行います。
164
+ ```bash
165
+ python offline_inference.py --model-path where_you_download_dir
166
+ ```
167
+ #### TTS推論
168
+ デフォルトのスピーカーを使用してTTSを推論するか、新しいスピーカーでクローンを作成します
169
+ ```bash
170
+ python tts_inference.py --model-path where_you_download_dir --output-path where_you_save_audio_dir --synthesis-type use_tts_or_clone
171
+ ```
172
+ クローンモードには、次の形式のスピーカー情報辞書が必要です:
173
+ ```bash
174
+ {
175
+ "speaker": "speaker id",
176
+ "prompt_text": "content of prompt wav",
177
+ "wav_path": "prompt wav path"
178
+ }
179
+ ```
180
+
181
+ #### Webデモの起動
182
+ オンライン推論のためにローカルサーバーを起動します。
183
+ 4つのGPUが利用可能で、すべてのモデルをダウンロード済みであると仮定します。
184
+
185
+ ```bash
186
+ # Step-Audio-Chat デモ
187
+ python app.py --model-path where_you_download_dir
188
+
189
+ # Step-Audio-TTS-3B デモ
190
+ python tts_app.py --model-path where_you_download_dir
191
+
192
+ ```
193
+
194
+ #### vLLMを用いた対話モデル推論(推奨)
195
+ Step-Audio-Chatは130Bパラメータの大規模言語モデルであり、テンソル並列処理をサポートするvLLMを使用した推論を推奨します。
196
+ * vLLMはTokenizerおよびTTSをロードしないため、音声入力による推論には対応していません
197
+
198
+ 現在の公式vLLMはStep 1モデルアーキテクチャに対応していないため、当社の[開発ブランチ](https://github.com/stepfun-ai/vllm/tree/add-step1-model)を使用したローカルインストールを推奨します。
199
+
200
+ 本モデルのAttentionメカニズムはALIBIの変種実装を採用しているため、公式flash attentionライブラリとの互換性がありません。[Step-Audio-Chat](https://huggingface.co/stepfun-ai/Step-Audio-Chat/tree/main/lib)リポジトリにカスタム版flash attentionライブラリを提供しています。モデル実行前に必ず環境変数へカスタムライブラリのパスを追加してください。
201
+
202
+ ```bash
203
+ export OPTIMUS_LIB_PATH=where_you_download_dir/Step-Audio-Chat/lib
204
+
205
+ vllm serve where_you_download_dir/Step-Audio-Chat --dtype auto -tp $tp --served-model-name step-audio-chat --trust-remote-code
206
+
207
+ # vLLMチャットの呼び出し例
208
+ python call_vllm_chat.py
209
+ ```
210
+
211
+ ## 5. ベンチマーク
212
+
213
+ ### 5.1 ASR結果の比較
214
+
215
+ <table>
216
+ <thead>
217
+ <tr>
218
+ <th style="text-align:center"></th>
219
+ <th colspan="4" style="text-align:center">隠れた特徴モデリング</th>
220
+ <th colspan="5" style="text-align:center">離散音声トークンモデリング</th>
221
+ </tr>
222
+ <tr>
223
+ <th style="text-align:center"></th>
224
+ <th style="text-align:center">Whisper Large-v3</th>
225
+ <th style="text-align:center">Qwen2-Audio</th>
226
+ <th style="text-align:center">MinMo</th>
227
+ <th style="text-align:center">LUCY</th>
228
+ <th style="text-align:center">Moshi</th>
229
+ <th style="text-align:center">GLM-4-voice Base</th>
230
+ <th style="text-align:center">GLM-4-voice Chat</th>
231
+ <th style="text-align:center">Step-Audio Pretrain</th>
232
+ <th style="text-align:center">Step-Audio-Chat</th>
233
+ </tr>
234
+ </thead>
235
+ <tbody>
236
+ <tr>
237
+ <td>Aishell-1</td>
238
+ <td style="text-align:center">5.14</td>
239
+ <td style="text-align:center">1.53</td>
240
+ <td style="text-align:center">-</td>
241
+ <td style="text-align:center">2.4</td>
242
+ <td style="text-align:center">-</td>
243
+ <td style="text-align:center">2.46</td>
244
+ <td style="text-align:center">226.47</td>
245
+ <td style="text-align:center"><strong>0.87</strong></td>
246
+ <td style="text-align:center">1.95</td>
247
+ </tr>
248
+ <tr>
249
+ <td>Aishell-2 ios</td>
250
+ <td style="text-align:center">4.76</td>
251
+ <td style="text-align:center">3.06</td>
252
+ <td style="text-align:center"><strong>2.69</strong></td>
253
+ <td style="text-align:center">-</td>
254
+ <td style="text-align:center">-</td>
255
+ <td style="text-align:center">-</td>
256
+ <td style="text-align:center">211.3</td>
257
+ <td style="text-align:center">2.91</td>
258
+ <td style="text-align:center">3.57</td>
259
+ </tr>
260
+ <tr>
261
+ <td>Wenetspeech test-net</td>
262
+ <td style="text-align:center">9.68</td>
263
+ <td style="text-align:center">7.72</td>
264
+ <td style="text-align:center"><strong>6.64</strong></td>
265
+ <td style="text-align:center">8.78</td>
266
+ <td style="text-align:center">-</td>
267
+ <td style="text-align:center">-</td>
268
+ <td style="text-align:center">146.05</td>
269
+ <td style="text-align:center">7.62</td>
270
+ <td style="text-align:center">8.75</td>
271
+ </tr>
272
+ <tr>
273
+ <td>Wenet test-meeting</td>
274
+ <td style="text-align:center">18.54</td>
275
+ <td style="text-align:center">8.4</td>
276
+ <td style="text-align:center"><strong>7.6</strong></td>
277
+ <td style="text-align:center">10.42</td>
278
+ <td style="text-align:center">-</td>
279
+ <td style="text-align:center">-</td>
280
+ <td style="text-align:center">140.82</td>
281
+ <td style="text-align:center">7.78</td>
282
+ <td style="text-align:center">9.52</td>
283
+ </tr>
284
+ <tr>
285
+ <td>Librispeech test-clean</td>
286
+ <td style="text-align:center">1.9</td>
287
+ <td style="text-align:center"><strong>1.6</strong></td>
288
+ <td style="text-align:center"><strong>1.6</strong></td>
289
+ <td style="text-align:center">3.36</td>
290
+ <td style="text-align:center">5.7</td>
291
+ <td style="text-align:center">2.82</td>
292
+ <td style="text-align:center">75.39</td>
293
+ <td style="text-align:center">2.36</td>
294
+ <td style="text-align:center">3.11</td>
295
+ </tr>
296
+ <tr>
297
+ <td>Librispeech test-other</td>
298
+ <td style="text-align:center">3.65</td>
299
+ <td style="text-align:center"><strong>3.6</strong></td>
300
+ <td style="text-align:center">3.82</td>
301
+ <td style="text-align:center">8.05</td>
302
+ <td style="text-align:center">-</td>
303
+ <td style="text-align:center">7.66</td>
304
+ <td style="text-align:center">80.3</td>
305
+ <td style="text-align:center">6.32</td>
306
+ <td style="text-align:center">8.44</td>
307
+ </tr>
308
+ <tr>
309
+ <td>AVG</td>
310
+ <td style="text-align:center">7.28</td>
311
+ <td style="text-align:center"><strong>4.32</strong></td>
312
+ <td style="text-align:center">-</td>
313
+ <td style="text-align:center">-</td>
314
+ <td style="text-align:center">-</td>
315
+ <td style="text-align:center">-</td>
316
+ <td style="text-align:center">146.74</td>
317
+ <td style="text-align:center">4.64</td>
318
+ <td style="text-align:center">5.89</td>
319
+ </tr>
320
+ </tbody>
321
+ </table>
322
+
323
+ ### 5.2 TTS
324
+ #### 5.2.1 GLM-4-VoiceとMinMoのコンテンツ一貫性(CER/WER)のパフォーマンス比較。
325
+
326
+ <table>
327
+ <thead>
328
+ <tr>
329
+ <th rowspan="2">モデル</th>
330
+ <th style="text-align:center" colspan="1">test-zh</th>
331
+ <th style="text-align:center" colspan="1">test-en</th>
332
+ </tr>
333
+ <tr>
334
+ <th style="text-align:center">CER (%) &darr;</th>
335
+ <th style="text-align:center">WER (%) &darr;</th>
336
+ </tr>
337
+ </thead>
338
+ <tbody>
339
+ <tr>
340
+ <td>GLM-4-Voice</td>
341
+ <td style="text-align:center">2.19</td>
342
+ <td style="text-align:center">2.91</td>
343
+ </tr>
344
+ <tr>
345
+ <td>MinMo</td>
346
+ <td style="text-align:center">2.48</td>
347
+ <td style="text-align:center">2.90</td>
348
+ </tr>
349
+ <tr>
350
+ <td><strong>Step-Audio</strong></td>
351
+ <td style="text-align:center"><strong>1.53</strong></td>
352
+ <td style="text-align:center"><strong>2.71</strong></td>
353
+ </tr>
354
+ </tbody>
355
+ </table>
356
+
357
+ #### 5.2.2 SEEDテストセットでのTTSモデルの結果。
358
+ * StepAudio-TTS-3B-Singleは、デュアルコードブックバックボーンとシングルコードブックボコーダの組み合���せを示します。
359
+
360
+ <table>
361
+ <thead>
362
+ <tr>
363
+ <th rowspan="2">モデル</th>
364
+ <th style="text-align:center" colspan="2">test-zh</th>
365
+ <th style="text-align:center" colspan="2">test-en</th>
366
+ </tr>
367
+ <tr>
368
+ <th style="text-align:center">CER (%) &darr;</th>
369
+ <th style="text-align:center">SS &uarr;</th>
370
+ <th style="text-align:center">WER (%) &darr;</th>
371
+ <th style="text-align:center">SS &uarr;</th>
372
+ </tr>
373
+ </thead>
374
+ <tbody>
375
+ <tr>
376
+ <td>FireRedTTS</td>
377
+ <td style="text-align:center">1.51</td>
378
+ <td style="text-align:center">0.630</td>
379
+ <td style="text-align:center">3.82</td>
380
+ <td style="text-align:center">0.460</td>
381
+ </tr>
382
+ <tr>
383
+ <td>MaskGCT</td>
384
+ <td style="text-align:center">2.27</td>
385
+ <td style="text-align:center">0.774</td>
386
+ <td style="text-align:center">2.62</td>
387
+ <td style="text-align:center">0.774</td>
388
+ </tr>
389
+ <tr>
390
+ <td>CosyVoice</td>
391
+ <td style="text-align:center">3.63</td>
392
+ <td style="text-align:center">0.775</td>
393
+ <td style="text-align:center">4.29</td>
394
+ <td style="text-align:center">0.699</td>
395
+ </tr>
396
+ <tr>
397
+ <td>CosyVoice 2</td>
398
+ <td style="text-align:center">1.45</td>
399
+ <td style="text-align:center">0.806</td>
400
+ <td style="text-align:center">2.57</td>
401
+ <td style="text-align:center">0.736</td>
402
+ </tr>
403
+ <tr>
404
+ <td>CosyVoice 2-S</td>
405
+ <td style="text-align:center">1.45</td>
406
+ <td style="text-align:center">0.812</td>
407
+ <td style="text-align:center">2.38</td>
408
+ <td style="text-align:center">0.743</td>
409
+ </tr>
410
+ <tr>
411
+ <td><strong>Step-Audio-TTS-3B-Single</strong></td>
412
+ <td style="text-align:center">1.37</td>
413
+ <td style="text-align:center">0.802</td>
414
+ <td style="text-align:center">2.52</td>
415
+ <td style="text-align:center">0.704</td>
416
+ </tr>
417
+ <tr>
418
+ <td><strong>Step-Audio-TTS-3B</strong></td>
419
+ <td style="text-align:center"><strong>1.31</strong></td>
420
+ <td style="text-align:center">0.733</td>
421
+ <td style="text-align:center"><strong>2.31</strong></td>
422
+ <td style="text-align:center">0.660</td>
423
+ </tr>
424
+ <tr>
425
+ <td><strong>Step-Audio-TTS</strong></td>
426
+ <td style="text-align:center"><strong>1.17</strong></td>
427
+ <td style="text-align:center">0.73</td>
428
+ <td style="text-align:center"><strong>2.0</strong></td>
429
+ <td style="text-align:center">0.660</td>
430
+ </tr>
431
+ </tbody>
432
+ </table>
433
+
434
+ #### 5.2.3 デュアルコードブック再合成とCosyVoiceのパフォーマンス比較。
435
+
436
+ <table>
437
+ <thead>
438
+ <tr>
439
+ <th style="text-align:center" rowspan="2">トークン</th>
440
+ <th style="text-align:center" colspan="2">test-zh</th>
441
+ <th style="text-align:center" colspan="2">test-en</th>
442
+ </tr>
443
+ <tr>
444
+ <th style="text-align:center">CER (%) &darr;</th>
445
+ <th style="text-align:center">SS &uarr;</th>
446
+ <th style="text-align:center">WER (%) &darr;</th>
447
+ <th style="text-align:center">SS &uarr;</th>
448
+ </tr>
449
+ </thead>
450
+ <tbody>
451
+ <tr>
452
+ <td style="text-align:center">Groundtruth</td>
453
+ <td style="text-align:center">0.972</td>
454
+ <td style="text-align:center">-</td>
455
+ <td style="text-align:center">2.156</td>
456
+ <td style="text-align:center">-</td>
457
+ </tr>
458
+ <tr>
459
+ <td style="text-align:center">CosyVoice</td>
460
+ <td style="text-align:center">2.857</td>
461
+ <td style="text-align:center"><strong>0.849</strong></td>
462
+ <td style="text-align:center">4.519</td>
463
+ <td style="text-align:center"><strong>0.807</strong></td>
464
+ </tr>
465
+ <tr>
466
+ <td style="text-align:center">Step-Audio-TTS-3B</td>
467
+ <td style="text-align:center"><strong>2.192</strong></td>
468
+ <td style="text-align:center">0.784</td>
469
+ <td style="text-align:center"><strong>3.585</strong></td>
470
+ <td style="text-align:center">0.742</td>
471
+ </tr>
472
+ </tbody>
473
+ </table>
474
+
475
+ ### 5.3 AQTAチャット
476
+ [**StepEval-Audio-360**](https://huggingface.co/datasets/stepfun-ai/StepEval-Audio-360) を新しいベンチマークとしてリリースしました。これは、実際のユーザーからの137のマルチターンの日本語プロンプトで構成されており、生成された応答の品質を次の次元で評価するように設計されています:音声指示のフォロー、音声理解、論理的推論、ロールプレイング、創造性、歌唱、言語能力、音���感情制御、ゲーム。
477
+
478
+ #### 5.3.1 StepEval-Audio-360
479
+
480
+ #### LLM評価指標(GPT-4o)
481
+ <table>
482
+ <caption>StepEval-Audio-360での音声チャットの基本機能の比較。</caption>
483
+ <thead>
484
+ <tr>
485
+ <th>モデル</th>
486
+ <th style="text-align:center">事実性(% &uarr;)</th>
487
+ <th style="text-align:center">関連性(% &uarr;)</th>
488
+ <th style="text-align:center">チャットスコア &uarr;</th>
489
+ </tr>
490
+ </thead>
491
+ <tbody>
492
+ <tr>
493
+ <td>GLM4-Voice</td>
494
+ <td style="text-align:center">54.7</td>
495
+ <td style="text-align:center">66.4</td>
496
+ <td style="text-align:center">3.49</td>
497
+ </tr>
498
+ <tr>
499
+ <td>Qwen2-Audio</td>
500
+ <td style="text-align:center">22.6</td>
501
+ <td style="text-align:center">26.3</td>
502
+ <td style="text-align:center">2.27</td>
503
+ </tr>
504
+ <tr>
505
+ <td>Moshi<sup>*</sup></td>
506
+ <td style="text-align:center">1.0</td>
507
+ <td style="text-align:center">0</td>
508
+ <td style="text-align:center">1.49</td>
509
+ </tr>
510
+ <tr>
511
+ <td><strong>Step-Audio-Chat</strong></td>
512
+ <td style="text-align:center"><strong>66.4</strong></td>
513
+ <td style="text-align:center"><strong>75.2</strong></td>
514
+ <td style="text-align:center"><strong>4.11</strong></td>
515
+ </tr>
516
+ </tbody>
517
+ </table>
518
+
519
+ * 注:Moshiは「\*」でマークされており、参考として考慮する必要があります。
520
+
521
+ #### レーダーチャート(人間の評価)
522
+ <img src="./assets/stepeval_radar_chart.png" width="600" alt="QR code">
523
+
524
+ #### 5.3.2 公開テストセット
525
+
526
+ <table>
527
+ <thead>
528
+ <tr>
529
+ <th>モデル</th>
530
+ <th style="text-align:center">Llama Question</th>
531
+ <th style="text-align:center">Web Questions</th>
532
+ <th style="text-align:center">TriviaQA*</th>
533
+ <th style="text-align:center">ComplexBench</th>
534
+ <th style="text-align:center">HSK-6</th>
535
+ </tr>
536
+ </thead>
537
+ <tbody>
538
+ <tr>
539
+ <td>GLM4-Voice</td>
540
+ <td style="text-align:center">64.7</td>
541
+ <td style="text-align:center">32.2</td>
542
+ <td style="text-align:center">39.1</td>
543
+ <td style="text-align:center">66.0</td>
544
+ <td style="text-align:center">74.0</td>
545
+ </tr>
546
+ <tr>
547
+ <td>Moshi</td>
548
+ <td style="text-align:center">62.3</td>
549
+ <td style="text-align:center">26.6</td>
550
+ <td style="text-align:center">22.8</td>
551
+ <td style="text-align:center">-</td>
552
+ <td style="text-align:center">-</td>
553
+ </tr>
554
+ <tr>
555
+ <td>Freeze-Omni</td>
556
+ <td style="text-align:center">72.0</td>
557
+ <td style="text-align:center">44.7</td>
558
+ <td style="text-align:center">53.9</td>
559
+ <td style="text-align:center">-</td>
560
+ <td style="text-align:center">-</td>
561
+ </tr>
562
+ <tr>
563
+ <td>LUCY</td>
564
+ <td style="text-align:center">59.7</td>
565
+ <td style="text-align:center">29.3</td>
566
+ <td style="text-align:center">27.0</td>
567
+ <td style="text-align:center">-</td>
568
+ <td style="text-align:center">-</td>
569
+ </tr>
570
+ <tr>
571
+ <td>MinMo</td>
572
+ <td style="text-align:center">78.9</td>
573
+ <td style="text-align:center">55.0</td>
574
+ <td style="text-align:center">48.3</td>
575
+ <td style="text-align:center">-</td>
576
+ <td style="text-align:center">-</td>
577
+ </tr>
578
+ <tr>
579
+ <td>Qwen2-Audio</td>
580
+ <td style="text-align:center">52.0</td>
581
+ <td style="text-align:center">27.0</td>
582
+ <td style="text-align:center">37.3</td>
583
+ <td style="text-align:center">54.0</td>
584
+ <td style="text-align:center">-</td>
585
+ </tr>
586
+ <tr>
587
+ <td><strong>Step-Audio-Chat</strong></td>
588
+ <td style="text-align:center"><strong><i>81.0</i></strong></td>
589
+ <td style="text-align:center"><strong>75.1</strong></td>
590
+ <td style="text-align:center"><strong>58.0</strong></td>
591
+ <td style="text-align:center"><strong>74.0</strong></td>
592
+ <td style="text-align:center"><strong>86.0</strong></td>
593
+ </tr>
594
+ </tbody>
595
+ </table>
596
+
597
+ * 注:TriviaQAデータセットで「\*」でマークされた結果は参考として考慮されます。
598
+
599
+ #### 5.3.3 音声指示のフォロー
600
+ <table>
601
+ <thead>
602
+ <tr>
603
+ <th rowspan="2">カテゴリ</th>
604
+ <th colspan="2" style="text-align:center">指示のフォロー</th>
605
+ <th colspan="2" style="text-align:center">音声品質</th>
606
+ </tr>
607
+ <tr>
608
+ <th style="text-align:center">GLM-4-Voice</th>
609
+ <th style="text-align:center">Step-Audio</th>
610
+ <th style="text-align:center">GLM-4-Voice</th>
611
+ <th style="text-align:center">Step-Audio</th>
612
+ </tr>
613
+ </thead>
614
+ <tbody>
615
+ <tr>
616
+ <td>言語</td>
617
+ <td style="text-align:center">1.9</td>
618
+ <td style="text-align:center">3.8</td>
619
+ <td style="text-align:center">2.9</td>
620
+ <td style="text-align:center">3.3</td>
621
+ </tr>
622
+ <tr>
623
+ <td>ロールプレイング</td>
624
+ <td style="text-align:center">3.8</td>
625
+ <td style="text-align:center">4.2</td>
626
+ <td style="text-align:center">3.2</td>
627
+ <td style="text-align:center">3.6</td>
628
+ </tr>
629
+ <tr>
630
+ <td>歌唱 / ラップ</td>
631
+ <td style="text-align:center">2.1</td>
632
+ <td style="text-align:center">2.4</td>
633
+ <td style="text-align:center">2.4</td>
634
+ <td style="text-align:center">4</td>
635
+ </tr>
636
+ <tr>
637
+ <td>音声制御</td>
638
+ <td style="text-align:center">3.6</td>
639
+ <td style="text-align:center">4.4</td>
640
+ <td style="text-align:center">3.3</td>
641
+ <td style="text-align:center">4.1</td>
642
+ </tr>
643
+ </tbody>
644
+ </table>
645
+
646
+ ## 6. オンラインエンジン
647
+ Step-Audioのオンラインバージョンは、[跃问](https://yuewen.cn)のアプリバージョンからアクセスでき、いくつかの印象的な例も見つけることができます。
648
+
649
+ <img src="./assets/yuewen.jpeg" width="200" alt="QR code">
650
+
651
+ ## 7. 例
652
+ ### 音声クローン
653
+ | 役割 | プロンプト音声 | クローン音声 |
654
+ |:-------:|:-------:|:-------:|
655
+ |于谦| [google drive](https://drive.google.com/file/d/1N9EJypafFwmeL0R152GoL_CVGbYn1_9A/preview)<br>[audio file](https://github.com/stepfun-ai/Step-Audio/tree/main/examples/prompt_wav_yuqian.wav)|[google drive](https://drive.google.com/file/d/1Zs_1QrCUuoSqtUSdn2ENIor-k5baQdDV/preview)<br>[audio file](https://github.com/stepfun-ai/Step-Audio/tree/main/examples/clone_wav_yuqian.wav)|
656
+ |李雪琴| [google drive](https://drive.google.com/file/d/15SkZ29hksELYi1NDOxYOPu-kRTLSyke_/preview)<br>[audio file](https://github.com/stepfun-ai/Step-Audio/tree/main/examples/prompt_wav_lixueqin.wav)|[google drive](https://drive.google.com/file/d/11Le4qMqL2DmWpf7RFRpKUXERIR9TtKC0/preview)<br>[audio file](https://github.com/stepfun-ai/Step-Audio/tree/main/examples/clone_wav_lixueqin.wav)|
657
+
658
+ ### 速度制御
659
+ | プロンプト | 応答 |
660
+ |:-------:|:-------:|
661
+ |Human: 早口言葉を言ってください<br>Assistant: すもももももももものうち<br>Human: もっと早く言えますか?|[google drive](https://drive.google.com/file/d/1mAH-NRrOVZo4tv6gdAZkyJg8kRuTNNGC/preview)<br>[audio file](https://github.com/stepfun-ai/Step-Audio/tree/main/examples/speed_control1.wav)|
662
+ |Human: 早口言葉を言ってください<br>Assistant: すもももももももものうち<br>Human: もっと早く言えますか?<br>Assistant: すもももももももものうち<br>Human: もっとゆっくり言ってください。|[google drive](https://drive.google.com/file/d/1FhRnKo8uGrtO-cWg4qkrg8iDoNRbtqSX/preview)<br>[audio file](https://github.com/stepfun-ai/Step-Audio/tree/main/examples/speed_control2.wav)|
663
+
664
+ ### 高EQ(感情制御 & トーン制御)
665
+ | プロンプト | 応答 |
666
+ |:-------:|:-------:|
667
+ |Human: もっとかわいく話してみてください。|[google drive](https://drive.google.com/file/d/19IROE6_6h2UQVNniCmDTnrhxKRMOFHq3/preview)<br>[audio file](https://github.com/stepfun-ai/Step-Audio/tree/main/examples/tone_control.wav)|
668
+ |Human: どうしよう、人生がうまくいかない。|[google drive](https://drive.google.com/file/d/1JlLbOlzmdrokVdxtwy1S8eeWqsZR2Vmc/preview)<br>[audio file](https://github.com/stepfun-ai/Step-Audio/tree/main/examples/emotional_control1.wav)|
669
+ |Human: すごいですね。|[google drive](https://drive.google.com/file/d/19ga1RpguDP5r0Xfl1r5GY1J-kzbmHvJb/preview)<br>[audio file](https://github.com/stepfun-ai/Step-Audio/tree/main/examples/emotional_control2.wav)|
670
+
671
+ ### 多言語(例:日本語、英語、中国語)
672
+ | プロンプト | 応答 |
673
+ |:-------:|:-------:|
674
+ |Human: "It's raining cats and dogs" ってどういう意味ですか?<br>Assistant: "It's raining cats and dogs" というのは、非常に激しい雨が降っていることを意味します。実際に猫や犬が空から降ってくるわけではありません!これは激しい雨を表現するための面白い言い方です。|[google drive](https://drive.google.com/file/d/1LEIvdR5ANMzWX8GOTqUPTNrynNS1xx--/preview)<br>[audio file](https://github.com
almeval/models/stepaudio/assets/Step-Audio.pdf ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:de3a061acb18ba15f965113aa179ec76a9f1e513451c995f7e0ed18ae15c7ef2
3
+ size 6952313
almeval/models/stepaudio/assets/architecture.png ADDED

Git LFS Details

  • SHA256: 478749d7110ad9e4f31225d76e3e896fe95391a1711bbe379766731f1857c377
  • Pointer size: 131 Bytes
  • Size of remote file: 802 kB
almeval/models/stepaudio/assets/logo.png ADDED
almeval/models/stepaudio/assets/pipeline.png ADDED

Git LFS Details

  • SHA256: de38317a91553ea2a89ffb812697a8163fc05269a6c0e009e76b92e95c15e060
  • Pointer size: 131 Bytes
  • Size of remote file: 794 kB
almeval/models/stepaudio/assets/rlhf.png ADDED

Git LFS Details

  • SHA256: c0de6c940771ced3beaaf204c4fee6356bc44834e7f0990b66e9ad6b557248bc
  • Pointer size: 131 Bytes
  • Size of remote file: 880 kB
almeval/models/stepaudio/assets/stepeval_radar_chart.png ADDED

Git LFS Details

  • SHA256: cb17b166e608380c674a8a84bb64b598e547143eb650a51f54d9ecf10276f70e
  • Pointer size: 131 Bytes
  • Size of remote file: 365 kB
almeval/models/stepaudio/assets/yuewen.jpeg ADDED
almeval/models/stepaudio/examples/clone_wav_lixueqin.wav ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:4e9de6d9c0e98c466fec26a806f64a60d56faa99f641d389f9de8c259201bad5
3
+ size 285774
almeval/models/stepaudio/examples/clone_wav_yuqian.wav ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:5d89b264b1698dabe293e8d48c94499b8d19bea42859c97ff041769db62fbdcb
3
+ size 619086
almeval/models/stepaudio/examples/emotional_control1.wav ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:75ef4b31d860813e0a178c90df03e07e763c5ff3fc35d654cdf5f5f6c1451a12
3
+ size 980012
almeval/models/stepaudio/examples/emotional_control2.wav ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:abd39f0ebfb6607af65a77e25d808de7435cbe8fb3e74d53074224b353787ef1
3
+ size 327724
almeval/models/stepaudio/examples/multilingual1.wav ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:e8f8ee61cf128f9037a544a80fd72e421d627882228988543874dc97d4314906
3
+ size 100396
almeval/models/stepaudio/examples/multilingual2.wav ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:5c039642bff22bb884169dca731d851ceae43c58599438f85a95728b897c98c6
3
+ size 665678
almeval/models/stepaudio/examples/multilingual_singing.wav ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:58a5d3ee4067654abca8a6ca1afd6b9b83dfeadab2f40b112f203a4872bbe850
3
+ size 370694
almeval/models/stepaudio/examples/prompt_wav_lixueqin.wav ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:f051b9aaa2d20f5c9894673cedec15fa7b2ae1baaea16cb7653740cd68ef76e9
3
+ size 275242
almeval/models/stepaudio/examples/prompt_wav_yuqian.wav ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:6c50f24fec74f8e18b2fdcab401465d3a4992bd7e7028e2a35a9203db1f08df9
3
+ size 714284
almeval/models/stepaudio/examples/prompt_wav_zhaobenshan.wav ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:416beacd10c063c412497e9a2c4624ff0972a31ff4ef38c73c95768032241480
3
+ size 672044
almeval/models/stepaudio/examples/rap.wav ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:2241d9b6bf7eb272822251bbeff5c851e0d3bc82c67dd4d7e7e84766cb7e0aac
3
+ size 507948
almeval/models/stepaudio/examples/singing.wav ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:6787572f0eb81ccaee36306a603dbaa802103237113e029475083d311c16e69d
3
+ size 549966
almeval/models/stepaudio/examples/speed_control1.wav ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:e53b586f3fa992aa90fdecb05cc566438c2194d2994e4287333739c88724384b
3
+ size 189996
almeval/models/stepaudio/examples/speed_control2.wav ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:87147a26106aa5f63a61e57bcc674791396f87173b3cfbf1b0fb47a8dffe3269
3
+ size 264236
almeval/models/stepaudio/examples/tone_control.wav ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:24e4c974a0d614665a2d69829fbc93ff0cd6f1f2618b15a2850d81d9092fb6f4
3
+ size 550444
almeval/models/stepaudio/funasr_detach/__init__.py ADDED
@@ -0,0 +1,38 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Initialize funasr package."""
2
+
3
+ import os
4
+ import pkgutil
5
+ import importlib
6
+
7
+ dirname = os.path.dirname(__file__)
8
+ version_file = os.path.join(dirname, "version.txt")
9
+ with open(version_file, "r") as f:
10
+ __version__ = f.read().strip()
11
+
12
+
13
+ import importlib
14
+ import pkgutil
15
+
16
+
17
+ def import_submodules(package, recursive=True):
18
+ if isinstance(package, str):
19
+ package = importlib.import_module(package)
20
+ results = {}
21
+ for loader, name, is_pkg in pkgutil.walk_packages(
22
+ package.__path__, package.__name__ + "."
23
+ ):
24
+ try:
25
+ results[name] = importlib.import_module(name)
26
+ except Exception as e:
27
+ # 如果想要看到导入错误的具体信息,可以取消注释下面的行
28
+ # print(f"Failed to import {name}: {e}")
29
+ pass
30
+ if recursive and is_pkg:
31
+ results.update(import_submodules(name))
32
+ return results
33
+
34
+
35
+ import_submodules(__name__)
36
+
37
+ from funasr_detach.auto.auto_model import AutoModel
38
+ from funasr_detach.auto.auto_frontend import AutoFrontend
almeval/models/stepaudio/funasr_detach/auto/__init__.py ADDED
File without changes
almeval/models/stepaudio/funasr_detach/auto/auto_frontend.py ADDED
@@ -0,0 +1,90 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import time
2
+ import logging
3
+ from tqdm import tqdm
4
+
5
+ from funasr_detach.register import tables
6
+ from funasr_detach.download.download_from_hub import download_model
7
+ from funasr_detach.utils.load_utils import load_audio_text_image_video, extract_fbank
8
+ from funasr_detach.auto.auto_model import prepare_data_iterator
9
+ from funasr_detach.auto.auto_model import prepare_data_iterator
10
+
11
+
12
+ class AutoFrontend:
13
+ def __init__(self, **kwargs):
14
+ assert "model" in kwargs
15
+ if "model_conf" not in kwargs:
16
+ logging.info(
17
+ "download models from model hub: {}".format(
18
+ kwargs.get("model_hub", "ms")
19
+ )
20
+ )
21
+ kwargs = download_model(**kwargs)
22
+
23
+ # build frontend
24
+ frontend = kwargs.get("frontend", None)
25
+ if frontend is not None:
26
+ frontend_class = tables.frontend_classes.get(frontend)
27
+ frontend = frontend_class(**kwargs["frontend_conf"])
28
+
29
+ self.frontend = frontend
30
+ if "frontend" in kwargs:
31
+ del kwargs["frontend"]
32
+ self.kwargs = kwargs
33
+
34
+ def __call__(self, input, input_len=None, kwargs=None, **cfg):
35
+
36
+ kwargs = self.kwargs if kwargs is None else kwargs
37
+ kwargs.update(cfg)
38
+
39
+ key_list, data_list = prepare_data_iterator(input, input_len=input_len)
40
+ batch_size = kwargs.get("batch_size", 1)
41
+ device = kwargs.get("device", "cpu")
42
+ if device == "cpu":
43
+ batch_size = 1
44
+
45
+ meta_data = {}
46
+
47
+ result_list = []
48
+ num_samples = len(data_list)
49
+ pbar = tqdm(colour="blue", total=num_samples + 1, dynamic_ncols=True)
50
+
51
+ time0 = time.perf_counter()
52
+ for beg_idx in range(0, num_samples, batch_size):
53
+ end_idx = min(num_samples, beg_idx + batch_size)
54
+ data_batch = data_list[beg_idx:end_idx]
55
+ key_batch = key_list[beg_idx:end_idx]
56
+
57
+ # extract fbank feats
58
+ time1 = time.perf_counter()
59
+ audio_sample_list = load_audio_text_image_video(
60
+ data_batch, fs=self.frontend.fs, audio_fs=kwargs.get("fs", 16000)
61
+ )
62
+ time2 = time.perf_counter()
63
+ meta_data["load_data"] = f"{time2 - time1:0.3f}"
64
+ speech, speech_lengths = extract_fbank(
65
+ audio_sample_list,
66
+ data_type=kwargs.get("data_type", "sound"),
67
+ frontend=self.frontend,
68
+ **kwargs,
69
+ )
70
+ time3 = time.perf_counter()
71
+ meta_data["extract_feat"] = f"{time3 - time2:0.3f}"
72
+ meta_data["batch_data_time"] = (
73
+ speech_lengths.sum().item()
74
+ * self.frontend.frame_shift
75
+ * self.frontend.lfr_n
76
+ / 1000
77
+ )
78
+
79
+ speech.to(device=device), speech_lengths.to(device=device)
80
+ batch = {"input": speech, "input_len": speech_lengths, "key": key_batch}
81
+ result_list.append(batch)
82
+
83
+ pbar.update(1)
84
+ description = f"{meta_data}, "
85
+ pbar.set_description(description)
86
+
87
+ time_end = time.perf_counter()
88
+ pbar.set_description(f"time escaped total: {time_end - time0:0.3f}")
89
+
90
+ return result_list
almeval/models/stepaudio/funasr_detach/auto/auto_model.py ADDED
@@ -0,0 +1,573 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import json
2
+ import time
3
+ import copy
4
+ import torch
5
+ import random
6
+ import string
7
+ import logging
8
+ import os.path
9
+ import numpy as np
10
+ from tqdm import tqdm
11
+
12
+ from funasr_detach.register import tables
13
+ from funasr_detach.utils.load_utils import load_bytes
14
+ from funasr_detach.download.file import download_from_url
15
+ from funasr_detach.download.download_from_hub import download_model
16
+ from funasr_detach.utils.vad_utils import slice_padding_audio_samples
17
+ from funasr_detach.train_utils.set_all_random_seed import set_all_random_seed
18
+ from funasr_detach.train_utils.load_pretrained_model import load_pretrained_model
19
+ from funasr_detach.utils.load_utils import load_audio_text_image_video
20
+ from funasr_detach.utils.timestamp_tools import timestamp_sentence
21
+ from funasr_detach.models.campplus.utils import sv_chunk, postprocess, distribute_spk
22
+
23
+ try:
24
+ from funasr_detach.models.campplus.cluster_backend import ClusterBackend
25
+ except:
26
+ print("If you want to use the speaker diarization, please `pip install hdbscan`")
27
+
28
+
29
+ def prepare_data_iterator(data_in, input_len=None, data_type=None, key=None):
30
+ """
31
+
32
+ :param input:
33
+ :param input_len:
34
+ :param data_type:
35
+ :param frontend:
36
+ :return:
37
+ """
38
+ data_list = []
39
+ key_list = []
40
+ filelist = [".scp", ".txt", ".json", ".jsonl"]
41
+
42
+ chars = string.ascii_letters + string.digits
43
+ if isinstance(data_in, str) and data_in.startswith("http"): # url
44
+ data_in = download_from_url(data_in)
45
+ if isinstance(data_in, str) and os.path.exists(
46
+ data_in
47
+ ): # wav_path; filelist: wav.scp, file.jsonl;text.txt;
48
+ _, file_extension = os.path.splitext(data_in)
49
+ file_extension = file_extension.lower()
50
+ if file_extension in filelist: # filelist: wav.scp, file.jsonl;text.txt;
51
+ with open(data_in, encoding="utf-8") as fin:
52
+ for line in fin:
53
+ key = "rand_key_" + "".join(random.choice(chars) for _ in range(13))
54
+ if data_in.endswith(
55
+ ".jsonl"
56
+ ): # file.jsonl: json.dumps({"source": data})
57
+ lines = json.loads(line.strip())
58
+ data = lines["source"]
59
+ key = data["key"] if "key" in data else key
60
+ else: # filelist, wav.scp, text.txt: id \t data or data
61
+ lines = line.strip().split(maxsplit=1)
62
+ data = lines[1] if len(lines) > 1 else lines[0]
63
+ key = lines[0] if len(lines) > 1 else key
64
+
65
+ data_list.append(data)
66
+ key_list.append(key)
67
+ else:
68
+ key = "rand_key_" + "".join(random.choice(chars) for _ in range(13))
69
+ data_list = [data_in]
70
+ key_list = [key]
71
+ elif isinstance(data_in, (list, tuple)):
72
+ if data_type is not None and isinstance(
73
+ data_type, (list, tuple)
74
+ ): # mutiple inputs
75
+ data_list_tmp = []
76
+ for data_in_i, data_type_i in zip(data_in, data_type):
77
+ key_list, data_list_i = prepare_data_iterator(
78
+ data_in=data_in_i, data_type=data_type_i
79
+ )
80
+ data_list_tmp.append(data_list_i)
81
+ data_list = []
82
+ for item in zip(*data_list_tmp):
83
+ data_list.append(item)
84
+ else:
85
+ # [audio sample point, fbank, text]
86
+ data_list = data_in
87
+ key_list = [
88
+ "rand_key_" + "".join(random.choice(chars) for _ in range(13))
89
+ for _ in range(len(data_in))
90
+ ]
91
+ else: # raw text; audio sample point, fbank; bytes
92
+ if isinstance(data_in, bytes): # audio bytes
93
+ data_in = load_bytes(data_in)
94
+ if key is None:
95
+ key = "rand_key_" + "".join(random.choice(chars) for _ in range(13))
96
+ data_list = [data_in]
97
+ key_list = [key]
98
+
99
+ return key_list, data_list
100
+
101
+
102
+ class AutoModel:
103
+
104
+ def __init__(self, **kwargs):
105
+ if not kwargs.get("disable_log", False):
106
+ tables.print()
107
+
108
+ model, kwargs = self.build_model(**kwargs)
109
+
110
+ # if vad_model is not None, build vad model else None
111
+ vad_model = kwargs.get("vad_model", None)
112
+ vad_kwargs = kwargs.get("vad_model_revision", None)
113
+ if vad_model is not None:
114
+ logging.info("Building VAD model.")
115
+ vad_kwargs = {
116
+ "model": vad_model,
117
+ "model_revision": vad_kwargs,
118
+ "device": kwargs["device"],
119
+ }
120
+ vad_model, vad_kwargs = self.build_model(**vad_kwargs)
121
+
122
+ # if punc_model is not None, build punc model else None
123
+ punc_model = kwargs.get("punc_model", None)
124
+ punc_kwargs = kwargs.get("punc_model_revision", None)
125
+ if punc_model is not None:
126
+ logging.info("Building punc model.")
127
+ punc_kwargs = {
128
+ "model": punc_model,
129
+ "model_revision": punc_kwargs,
130
+ "device": kwargs["device"],
131
+ }
132
+ punc_model, punc_kwargs = self.build_model(**punc_kwargs)
133
+
134
+ # if spk_model is not None, build spk model else None
135
+ spk_model = kwargs.get("spk_model", None)
136
+ spk_kwargs = kwargs.get("spk_model_revision", None)
137
+ if spk_model is not None:
138
+ logging.info("Building SPK model.")
139
+ spk_kwargs = {
140
+ "model": spk_model,
141
+ "model_revision": spk_kwargs,
142
+ "device": kwargs["device"],
143
+ }
144
+ spk_model, spk_kwargs = self.build_model(**spk_kwargs)
145
+ self.cb_model = ClusterBackend().to(kwargs["device"])
146
+ spk_mode = kwargs.get("spk_mode", "punc_segment")
147
+ if spk_mode not in ["default", "vad_segment", "punc_segment"]:
148
+ logging.error(
149
+ "spk_mode should be one of default, vad_segment and punc_segment."
150
+ )
151
+ self.spk_mode = spk_mode
152
+
153
+ self.kwargs = kwargs
154
+ self.model = model
155
+ self.vad_model = vad_model
156
+ self.vad_kwargs = vad_kwargs
157
+ self.punc_model = punc_model
158
+ self.punc_kwargs = punc_kwargs
159
+ self.spk_model = spk_model
160
+ self.spk_kwargs = spk_kwargs
161
+ self.model_path = kwargs.get("model_path")
162
+
163
+ def build_model(self, **kwargs):
164
+ assert "model" in kwargs
165
+ if "model_conf" not in kwargs:
166
+ logging.info(
167
+ "download models from model hub: {}".format(
168
+ kwargs.get("model_hub", "ms")
169
+ )
170
+ )
171
+ kwargs = download_model(**kwargs)
172
+
173
+ set_all_random_seed(kwargs.get("seed", 0))
174
+
175
+ device = kwargs.get("device", "cuda")
176
+ if not torch.cuda.is_available() or kwargs.get("ngpu", 1) == 0:
177
+ device = "cpu"
178
+ kwargs["batch_size"] = 1
179
+ kwargs["device"] = device
180
+
181
+ if kwargs.get("ncpu", None):
182
+ torch.set_num_threads(kwargs.get("ncpu"))
183
+
184
+ # build tokenizer
185
+ tokenizer = kwargs.get("tokenizer", None)
186
+ if tokenizer is not None:
187
+ tokenizer_class = tables.tokenizer_classes.get(tokenizer)
188
+ tokenizer = tokenizer_class(**kwargs["tokenizer_conf"])
189
+ kwargs["tokenizer"] = tokenizer
190
+ kwargs["token_list"] = tokenizer.token_list
191
+ vocab_size = len(tokenizer.token_list)
192
+ else:
193
+ vocab_size = -1
194
+
195
+ # build frontend
196
+ frontend = kwargs.get("frontend", None)
197
+ if frontend is not None:
198
+ frontend_class = tables.frontend_classes.get(frontend)
199
+ frontend = frontend_class(**kwargs["frontend_conf"])
200
+ kwargs["frontend"] = frontend
201
+ kwargs["input_size"] = frontend.output_size()
202
+
203
+ # build model
204
+ model_class = tables.model_classes.get(kwargs["model"])
205
+ model = model_class(**kwargs, **kwargs["model_conf"], vocab_size=vocab_size)
206
+
207
+ model.to(device)
208
+
209
+ # init_param
210
+ init_param = kwargs.get("init_param", None)
211
+ if init_param is not None:
212
+ logging.info(f"Loading pretrained params from {init_param}")
213
+ load_pretrained_model(
214
+ model=model,
215
+ path=init_param,
216
+ ignore_init_mismatch=kwargs.get("ignore_init_mismatch", False),
217
+ oss_bucket=kwargs.get("oss_bucket", None),
218
+ scope_map=kwargs.get("scope_map", None),
219
+ excludes=kwargs.get("excludes", None),
220
+ )
221
+
222
+ return model, kwargs
223
+
224
+ def __call__(self, *args, **cfg):
225
+ kwargs = self.kwargs
226
+ kwargs.update(cfg)
227
+ res = self.model(*args, kwargs)
228
+ return res
229
+
230
+ def generate(self, input, input_len=None, **cfg):
231
+ if self.vad_model is None:
232
+ return self.inference(input, input_len=input_len, **cfg)
233
+
234
+ else:
235
+ return self.inference_with_vad(input, input_len=input_len, **cfg)
236
+
237
+ def inference(
238
+ self, input, input_len=None, model=None, kwargs=None, key=None, **cfg
239
+ ):
240
+ kwargs = self.kwargs if kwargs is None else kwargs
241
+ kwargs.update(cfg)
242
+ model = self.model if model is None else model
243
+ model = model.cuda()
244
+ model.eval()
245
+
246
+ batch_size = kwargs.get("batch_size", 1)
247
+ # if kwargs.get("device", "cpu") == "cpu":
248
+ # batch_size = 1
249
+
250
+ key_list, data_list = prepare_data_iterator(
251
+ input, input_len=input_len, data_type=kwargs.get("data_type", None), key=key
252
+ )
253
+
254
+ speed_stats = {}
255
+ asr_result_list = []
256
+ num_samples = len(data_list)
257
+ disable_pbar = kwargs.get("disable_pbar", False)
258
+ pbar = (
259
+ tqdm(colour="blue", total=num_samples, dynamic_ncols=True)
260
+ if not disable_pbar
261
+ else None
262
+ )
263
+ time_speech_total = 0.0
264
+ time_escape_total = 0.0
265
+ for beg_idx in range(0, num_samples, batch_size):
266
+ end_idx = min(num_samples, beg_idx + batch_size)
267
+ data_batch = data_list[beg_idx:end_idx]
268
+ key_batch = key_list[beg_idx:end_idx]
269
+ batch = {"data_in": data_batch, "key": key_batch}
270
+ if (end_idx - beg_idx) == 1 and kwargs.get(
271
+ "data_type", None
272
+ ) == "fbank": # fbank
273
+ batch["data_in"] = data_batch[0]
274
+ batch["data_lengths"] = input_len
275
+
276
+ time1 = time.perf_counter()
277
+ with torch.no_grad():
278
+ results, meta_data = model.inference(**batch, **kwargs)
279
+ time2 = time.perf_counter()
280
+
281
+ asr_result_list.extend(results)
282
+
283
+ # batch_data_time = time_per_frame_s * data_batch_i["speech_lengths"].sum().item()
284
+ batch_data_time = meta_data.get("batch_data_time", -1)
285
+ time_escape = time2 - time1
286
+ speed_stats["load_data"] = meta_data.get("load_data", 0.0)
287
+ speed_stats["extract_feat"] = meta_data.get("extract_feat", 0.0)
288
+ speed_stats["forward"] = f"{time_escape:0.3f}"
289
+ speed_stats["batch_size"] = f"{len(results)}"
290
+ speed_stats["time_cost"] = f"{(time_escape)}"
291
+ speed_stats["rtf"] = f"{(time_escape) / batch_data_time:0.3f}"
292
+ description = f"{speed_stats}, "
293
+ if pbar:
294
+ pbar.update(1)
295
+ pbar.set_description(description)
296
+ time_speech_total += batch_data_time
297
+ time_escape_total += time_escape
298
+
299
+ if pbar:
300
+ # pbar.update(1)
301
+ pbar.set_description(f"rtf_avg: {time_escape_total/time_speech_total:0.3f}")
302
+ torch.cuda.empty_cache()
303
+ return asr_result_list
304
+
305
+ def inference_with_vad(self, input, input_len=None, **cfg):
306
+
307
+ # step.1: compute the vad model
308
+ self.vad_kwargs.update(cfg)
309
+ beg_vad = time.time()
310
+ res = self.inference(
311
+ input,
312
+ input_len=input_len,
313
+ model=self.vad_model,
314
+ kwargs=self.vad_kwargs,
315
+ **cfg,
316
+ )
317
+ end_vad = time.time()
318
+ print(f"time cost vad: {end_vad - beg_vad:0.3f}")
319
+
320
+ # step.2 compute asr model
321
+ model = self.model
322
+ kwargs = self.kwargs
323
+ kwargs.update(cfg)
324
+ batch_size = int(kwargs.get("batch_size_s", 300)) * 1000
325
+ batch_size_threshold_ms = int(kwargs.get("batch_size_threshold_s", 60)) * 1000
326
+ kwargs["batch_size"] = batch_size
327
+
328
+ key_list, data_list = prepare_data_iterator(
329
+ input, input_len=input_len, data_type=kwargs.get("data_type", None)
330
+ )
331
+ results_ret_list = []
332
+ time_speech_total_all_samples = 1e-6
333
+
334
+ beg_total = time.time()
335
+ pbar_total = tqdm(colour="red", total=len(res), dynamic_ncols=True)
336
+ for i in range(len(res)):
337
+ key = res[i]["key"]
338
+ vadsegments = res[i]["value"]
339
+ input_i = data_list[i]
340
+ speech = load_audio_text_image_video(
341
+ input_i, fs=kwargs["frontend"].fs, audio_fs=kwargs.get("fs", 16000)
342
+ )
343
+ speech_lengths = len(speech)
344
+ n = len(vadsegments)
345
+ data_with_index = [(vadsegments[i], i) for i in range(n)]
346
+ sorted_data = sorted(data_with_index, key=lambda x: x[0][1] - x[0][0])
347
+ results_sorted = []
348
+
349
+ if not len(sorted_data):
350
+ logging.info("decoding, utt: {}, empty speech".format(key))
351
+ continue
352
+
353
+ if len(sorted_data) > 0 and len(sorted_data[0]) > 0:
354
+ batch_size = max(
355
+ batch_size, sorted_data[0][0][1] - sorted_data[0][0][0]
356
+ )
357
+
358
+ batch_size_ms_cum = 0
359
+ beg_idx = 0
360
+ beg_asr_total = time.time()
361
+ time_speech_total_per_sample = speech_lengths / 16000
362
+ time_speech_total_all_samples += time_speech_total_per_sample
363
+
364
+ all_segments = []
365
+ for j, _ in enumerate(range(0, n)):
366
+ # pbar_sample.update(1)
367
+ batch_size_ms_cum += sorted_data[j][0][1] - sorted_data[j][0][0]
368
+ if (
369
+ j < n - 1
370
+ and (
371
+ batch_size_ms_cum
372
+ + sorted_data[j + 1][0][1]
373
+ - sorted_data[j + 1][0][0]
374
+ )
375
+ < batch_size
376
+ and (sorted_data[j + 1][0][1] - sorted_data[j + 1][0][0])
377
+ < batch_size_threshold_ms
378
+ ):
379
+ continue
380
+ batch_size_ms_cum = 0
381
+ end_idx = j + 1
382
+ speech_j, speech_lengths_j = slice_padding_audio_samples(
383
+ speech, speech_lengths, sorted_data[beg_idx:end_idx]
384
+ )
385
+ results = self.inference(
386
+ speech_j,
387
+ input_len=None,
388
+ model=model,
389
+ kwargs=kwargs,
390
+ disable_pbar=True,
391
+ **cfg,
392
+ )
393
+ if self.spk_model is not None:
394
+ # compose vad segments: [[start_time_sec, end_time_sec, speech], [...]]
395
+ for _b in range(len(speech_j)):
396
+ vad_segments = [
397
+ [
398
+ sorted_data[beg_idx:end_idx][_b][0][0] / 1000.0,
399
+ sorted_data[beg_idx:end_idx][_b][0][1] / 1000.0,
400
+ np.array(speech_j[_b]),
401
+ ]
402
+ ]
403
+ segments = sv_chunk(vad_segments)
404
+ all_segments.extend(segments)
405
+ speech_b = [i[2] for i in segments]
406
+ spk_res = self.inference(
407
+ speech_b,
408
+ input_len=None,
409
+ model=self.spk_model,
410
+ kwargs=kwargs,
411
+ disable_pbar=True,
412
+ **cfg,
413
+ )
414
+ results[_b]["spk_embedding"] = spk_res[0]["spk_embedding"]
415
+ beg_idx = end_idx
416
+ if len(results) < 1:
417
+ continue
418
+ results_sorted.extend(results)
419
+
420
+ restored_data = [0] * n
421
+ for j in range(n):
422
+ index = sorted_data[j][1]
423
+ restored_data[index] = results_sorted[j]
424
+ result = {}
425
+
426
+ # results combine for texts, timestamps, speaker embeddings and others
427
+ # TODO: rewrite for clean code
428
+ for j in range(n):
429
+ for k, v in restored_data[j].items():
430
+ if k.startswith("timestamp"):
431
+ if k not in result:
432
+ result[k] = []
433
+ for t in restored_data[j][k]:
434
+ t[0] += vadsegments[j][0]
435
+ t[1] += vadsegments[j][0]
436
+ result[k].extend(restored_data[j][k])
437
+ elif k == "spk_embedding":
438
+ if k not in result:
439
+ result[k] = restored_data[j][k]
440
+ else:
441
+ result[k] = torch.cat(
442
+ [result[k], restored_data[j][k]], dim=0
443
+ )
444
+ elif "text" in k:
445
+ if k not in result:
446
+ result[k] = restored_data[j][k]
447
+ else:
448
+ result[k] += " " + restored_data[j][k]
449
+ else:
450
+ if k not in result:
451
+ result[k] = restored_data[j][k]
452
+ else:
453
+ result[k] += restored_data[j][k]
454
+
455
+ return_raw_text = kwargs.get("return_raw_text", False)
456
+ # step.3 compute punc model
457
+ if self.punc_model is not None:
458
+ self.punc_kwargs.update(cfg)
459
+ punc_res = self.inference(
460
+ result["text"],
461
+ model=self.punc_model,
462
+ kwargs=self.punc_kwargs,
463
+ disable_pbar=True,
464
+ **cfg,
465
+ )
466
+ raw_text = copy.copy(result["text"])
467
+ if return_raw_text:
468
+ result["raw_text"] = raw_text
469
+ result["text"] = punc_res[0]["text"]
470
+ else:
471
+ raw_text = None
472
+
473
+ # speaker embedding cluster after resorted
474
+ if self.spk_model is not None and kwargs.get("return_spk_res", True):
475
+ if raw_text is None:
476
+ logging.error("Missing punc_model, which is required by spk_model.")
477
+ all_segments = sorted(all_segments, key=lambda x: x[0])
478
+ spk_embedding = result["spk_embedding"]
479
+ labels = self.cb_model(
480
+ spk_embedding.cpu(), oracle_num=kwargs.get("preset_spk_num", None)
481
+ )
482
+ # del result['spk_embedding']
483
+ sv_output = postprocess(all_segments, None, labels, spk_embedding.cpu())
484
+ if self.spk_mode == "vad_segment": # recover sentence_list
485
+ sentence_list = []
486
+ for res, vadsegment in zip(restored_data, vadsegments):
487
+ if "timestamp" not in res:
488
+ logging.error(
489
+ "Only 'iic/speech_paraformer-large-vad-punc_asr_nat-zh-cn-16k-common-vocab8404-pytorch' \
490
+ and 'iic/speech_seaco_paraformer_large_asr_nat-zh-cn-16k-common-vocab8404-pytorch'\
491
+ can predict timestamp, and speaker diarization relies on timestamps."
492
+ )
493
+ sentence_list.append(
494
+ {
495
+ "start": vadsegment[0],
496
+ "end": vadsegment[1],
497
+ "sentence": res["text"],
498
+ "timestamp": res["timestamp"],
499
+ }
500
+ )
501
+ elif self.spk_mode == "punc_segment":
502
+ if "timestamp" not in result:
503
+ logging.error(
504
+ "Only 'iic/speech_paraformer-large-vad-punc_asr_nat-zh-cn-16k-common-vocab8404-pytorch' \
505
+ and 'iic/speech_seaco_paraformer_large_asr_nat-zh-cn-16k-common-vocab8404-pytorch'\
506
+ can predict timestamp, and speaker diarization relies on timestamps."
507
+ )
508
+ sentence_list = timestamp_sentence(
509
+ punc_res[0]["punc_array"],
510
+ result["timestamp"],
511
+ raw_text,
512
+ return_raw_text=return_raw_text,
513
+ )
514
+ distribute_spk(sentence_list, sv_output)
515
+ result["sentence_info"] = sentence_list
516
+ elif kwargs.get("sentence_timestamp", False):
517
+ sentence_list = timestamp_sentence(
518
+ punc_res[0]["punc_array"],
519
+ result["timestamp"],
520
+ raw_text,
521
+ return_raw_text=return_raw_text,
522
+ )
523
+ result["sentence_info"] = sentence_list
524
+ if "spk_embedding" in result:
525
+ del result["spk_embedding"]
526
+
527
+ result["key"] = key
528
+ results_ret_list.append(result)
529
+ end_asr_total = time.time()
530
+ time_escape_total_per_sample = end_asr_total - beg_asr_total
531
+ pbar_total.update(1)
532
+ pbar_total.set_description(
533
+ f"rtf_avg: {time_escape_total_per_sample / time_speech_total_per_sample:0.3f}, "
534
+ f"time_speech: {time_speech_total_per_sample: 0.3f}, "
535
+ f"time_escape: {time_escape_total_per_sample:0.3f}"
536
+ )
537
+
538
+ return results_ret_list
539
+
540
+ def infer_encoder(
541
+ self, input, input_len=None, model=None, kwargs=None, key=None, **cfg
542
+ ):
543
+ kwargs = self.kwargs if kwargs is None else kwargs
544
+ kwargs.update(cfg)
545
+ model = self.model if model is None else model
546
+ model = model.cuda()
547
+ model.eval()
548
+
549
+ batch_size = kwargs.get("batch_size", 1)
550
+
551
+ key_list, data_list = prepare_data_iterator(
552
+ input, input_len=input_len, data_type=kwargs.get("data_type", None), key=key
553
+ )
554
+
555
+ asr_result_list = []
556
+ num_samples = len(data_list)
557
+ for beg_idx in range(0, num_samples, batch_size):
558
+ end_idx = min(num_samples, beg_idx + batch_size)
559
+ data_batch = data_list[beg_idx:end_idx]
560
+ key_batch = key_list[beg_idx:end_idx]
561
+ batch = {"data_in": data_batch, "key": key_batch}
562
+ if (end_idx - beg_idx) == 1 and kwargs.get(
563
+ "data_type", None
564
+ ) == "fbank": # fbank
565
+ batch["data_in"] = data_batch[0]
566
+ batch["data_lengths"] = input_len
567
+
568
+ with torch.no_grad():
569
+ results, meta_data, cache = model.infer_encoder(**batch, **kwargs)
570
+ asr_result_list.extend(results)
571
+
572
+ torch.cuda.empty_cache()
573
+ return asr_result_list, cache
almeval/models/stepaudio/funasr_detach/auto/auto_tokenizer.py ADDED
@@ -0,0 +1,7 @@
 
 
 
 
 
 
 
 
1
+ class AutoTokenizer:
2
+ """
3
+ Undo
4
+ """
5
+
6
+ def __init__(self):
7
+ pass
almeval/models/stepaudio/funasr_detach/bin/__init__.py ADDED
File without changes
almeval/models/stepaudio/funasr_detach/bin/compute_audio_cmvn.py ADDED
@@ -0,0 +1,152 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import json
3
+ import numpy as np
4
+ import torch
5
+ import hydra
6
+ import logging
7
+ from omegaconf import DictConfig, OmegaConf
8
+
9
+ from funasr_detach.register import tables
10
+ from funasr_detach.download.download_from_hub import download_model
11
+ from funasr_detach.train_utils.set_all_random_seed import set_all_random_seed
12
+
13
+
14
+ @hydra.main(config_name=None, version_base=None)
15
+ def main_hydra(kwargs: DictConfig):
16
+ if kwargs.get("debug", False):
17
+ import pdb
18
+
19
+ pdb.set_trace()
20
+
21
+ assert "model" in kwargs
22
+ if "model_conf" not in kwargs:
23
+ logging.info(
24
+ "download models from model hub: {}".format(kwargs.get("model_hub", "ms"))
25
+ )
26
+ kwargs = download_model(is_training=kwargs.get("is_training", True), **kwargs)
27
+
28
+ main(**kwargs)
29
+
30
+
31
+ def main(**kwargs):
32
+ print(kwargs)
33
+ # set random seed
34
+ tables.print()
35
+ set_all_random_seed(kwargs.get("seed", 0))
36
+ torch.backends.cudnn.enabled = kwargs.get(
37
+ "cudnn_enabled", torch.backends.cudnn.enabled
38
+ )
39
+ torch.backends.cudnn.benchmark = kwargs.get(
40
+ "cudnn_benchmark", torch.backends.cudnn.benchmark
41
+ )
42
+ torch.backends.cudnn.deterministic = kwargs.get("cudnn_deterministic", True)
43
+
44
+ tokenizer = kwargs.get("tokenizer", None)
45
+
46
+ # build frontend if frontend is none None
47
+ frontend = kwargs.get("frontend", None)
48
+ if frontend is not None:
49
+ frontend_class = tables.frontend_classes.get(frontend)
50
+ frontend = frontend_class(**kwargs["frontend_conf"])
51
+ kwargs["frontend"] = frontend
52
+ kwargs["input_size"] = frontend.output_size()
53
+
54
+ # dataset
55
+ dataset_class = tables.dataset_classes.get(kwargs.get("dataset", "AudioDataset"))
56
+ dataset_train = dataset_class(
57
+ kwargs.get("train_data_set_list"),
58
+ frontend=frontend,
59
+ tokenizer=None,
60
+ is_training=False,
61
+ **kwargs.get("dataset_conf")
62
+ )
63
+
64
+ # dataloader
65
+ batch_sampler = kwargs["dataset_conf"].get(
66
+ "batch_sampler", "DynamicBatchLocalShuffleSampler"
67
+ )
68
+ batch_sampler_train = None
69
+ if batch_sampler is not None:
70
+ batch_sampler_class = tables.batch_sampler_classes.get(batch_sampler)
71
+ dataset_conf = kwargs.get("dataset_conf")
72
+ dataset_conf["batch_type"] = "example"
73
+ dataset_conf["batch_size"] = 1
74
+ batch_sampler_train = batch_sampler_class(
75
+ dataset_train, is_training=False, **dataset_conf
76
+ )
77
+
78
+ dataloader_train = torch.utils.data.DataLoader(
79
+ dataset_train,
80
+ collate_fn=dataset_train.collator,
81
+ batch_sampler=batch_sampler_train,
82
+ num_workers=int(kwargs.get("dataset_conf").get("num_workers", 4)),
83
+ pin_memory=True,
84
+ )
85
+
86
+ iter_stop = int(kwargs.get("scale", 1.0) * len(dataloader_train))
87
+
88
+ total_frames = 0
89
+ for batch_idx, batch in enumerate(dataloader_train):
90
+ if batch_idx >= iter_stop:
91
+ break
92
+
93
+ fbank = batch["speech"].numpy()[0, :, :]
94
+ if total_frames == 0:
95
+ mean_stats = np.sum(fbank, axis=0)
96
+ var_stats = np.sum(np.square(fbank), axis=0)
97
+ else:
98
+ mean_stats += np.sum(fbank, axis=0)
99
+ var_stats += np.sum(np.square(fbank), axis=0)
100
+ total_frames += fbank.shape[0]
101
+
102
+ cmvn_info = {
103
+ "mean_stats": list(mean_stats.tolist()),
104
+ "var_stats": list(var_stats.tolist()),
105
+ "total_frames": total_frames,
106
+ }
107
+ cmvn_file = kwargs.get("cmvn_file", "cmvn.json")
108
+ # import pdb;pdb.set_trace()
109
+ with open(cmvn_file, "w") as fout:
110
+ fout.write(json.dumps(cmvn_info))
111
+
112
+ mean = -1.0 * mean_stats / total_frames
113
+ var = 1.0 / np.sqrt(var_stats / total_frames - mean * mean)
114
+ dims = mean.shape[0]
115
+ am_mvn = os.path.dirname(cmvn_file) + "/am.mvn"
116
+ with open(am_mvn, "w") as fout:
117
+ fout.write(
118
+ "<Nnet>"
119
+ + "\n"
120
+ + "<Splice> "
121
+ + str(dims)
122
+ + " "
123
+ + str(dims)
124
+ + "\n"
125
+ + "[ 0 ]"
126
+ + "\n"
127
+ + "<AddShift> "
128
+ + str(dims)
129
+ + " "
130
+ + str(dims)
131
+ + "\n"
132
+ )
133
+ mean_str = (
134
+ str(list(mean)).replace(",", "").replace("[", "[ ").replace("]", " ]")
135
+ )
136
+ fout.write("<LearnRateCoef> 0 " + mean_str + "\n")
137
+ fout.write("<Rescale> " + str(dims) + " " + str(dims) + "\n")
138
+ var_str = str(list(var)).replace(",", "").replace("[", "[ ").replace("]", " ]")
139
+ fout.write("<LearnRateCoef> 0 " + var_str + "\n")
140
+ fout.write("</Nnet>" + "\n")
141
+
142
+
143
+ """
144
+ python funasr/bin/compute_audio_cmvn.py \
145
+ --config-path "/Users/zhifu/funasr1.0/examples/aishell/paraformer/conf" \
146
+ --config-name "train_asr_paraformer_conformer_12e_6d_2048_256.yaml" \
147
+ ++train_data_set_list="/Users/zhifu/funasr1.0/data/list/audio_datasets.jsonl" \
148
+ ++cmvn_file="/Users/zhifu/funasr1.0/data/list/cmvn.json" \
149
+ ++dataset_conf.num_workers=0
150
+ """
151
+ if __name__ == "__main__":
152
+ main_hydra()
almeval/models/stepaudio/funasr_detach/bin/inference.py ADDED
@@ -0,0 +1,33 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import hydra
2
+ import logging
3
+ from omegaconf import DictConfig, OmegaConf, ListConfig
4
+
5
+ from funasr_detach.auto.auto_model import AutoModel
6
+
7
+
8
+ @hydra.main(config_name=None, version_base=None)
9
+ def main_hydra(cfg: DictConfig):
10
+ def to_plain_list(cfg_item):
11
+ if isinstance(cfg_item, ListConfig):
12
+ return OmegaConf.to_container(cfg_item, resolve=True)
13
+ elif isinstance(cfg_item, DictConfig):
14
+ return {k: to_plain_list(v) for k, v in cfg_item.items()}
15
+ else:
16
+ return cfg_item
17
+
18
+ kwargs = to_plain_list(cfg)
19
+ log_level = getattr(logging, kwargs.get("log_level", "INFO").upper())
20
+
21
+ logging.basicConfig(level=log_level)
22
+
23
+ if kwargs.get("debug", False):
24
+ import pdb
25
+
26
+ pdb.set_trace()
27
+ model = AutoModel(**kwargs)
28
+ res = model.generate(input=kwargs["input"])
29
+ print(res)
30
+
31
+
32
+ if __name__ == "__main__":
33
+ main_hydra()
almeval/models/stepaudio/funasr_detach/bin/tokenize_text.py ADDED
@@ -0,0 +1,281 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ import argparse
3
+ from collections import Counter
4
+ import logging
5
+ from pathlib import Path
6
+ import sys
7
+ from typing import List
8
+ from typing import Optional
9
+
10
+
11
+ from funasr_detach.utils.cli_utils import get_commandline_args
12
+ from funasr_detach.tokenizer.build_tokenizer import build_tokenizer
13
+ from funasr_detach.tokenizer.cleaner import TextCleaner
14
+ from funasr_detach.tokenizer.phoneme_tokenizer import g2p_classes
15
+ from funasr_detach.utils.types import str2bool
16
+ from funasr_detach.utils.types import str_or_none
17
+
18
+
19
+ def field2slice(field: Optional[str]) -> slice:
20
+ """Convert field string to slice
21
+
22
+ Note that field string accepts 1-based integer.
23
+
24
+ Examples:
25
+ >>> field2slice("1-")
26
+ slice(0, None, None)
27
+ >>> field2slice("1-3")
28
+ slice(0, 3, None)
29
+ >>> field2slice("-3")
30
+ slice(None, 3, None)
31
+ """
32
+ field = field.strip()
33
+ try:
34
+ if "-" in field:
35
+ # e.g. "2-" or "2-5" or "-7"
36
+ s1, s2 = field.split("-", maxsplit=1)
37
+ if s1.strip() == "":
38
+ s1 = None
39
+ else:
40
+ s1 = int(s1)
41
+ if s1 == 0:
42
+ raise ValueError("1-based string")
43
+ if s2.strip() == "":
44
+ s2 = None
45
+ else:
46
+ s2 = int(s2)
47
+ else:
48
+ # e.g. "2"
49
+ s1 = int(field)
50
+ s2 = s1 + 1
51
+ if s1 == 0:
52
+ raise ValueError("must be 1 or more value")
53
+ except ValueError:
54
+ raise RuntimeError(f"Format error: e.g. '2-', '2-5', or '-5': {field}")
55
+
56
+ if s1 is None:
57
+ slic = slice(None, s2)
58
+ else:
59
+ # -1 because of 1-based integer following "cut" command
60
+ # e.g "1-3" -> slice(0, 3)
61
+ slic = slice(s1 - 1, s2)
62
+ return slic
63
+
64
+
65
+ def tokenize(
66
+ input: str,
67
+ output: str,
68
+ field: Optional[str],
69
+ delimiter: Optional[str],
70
+ token_type: str,
71
+ space_symbol: str,
72
+ non_linguistic_symbols: Optional[str],
73
+ bpemodel: Optional[str],
74
+ log_level: str,
75
+ write_vocabulary: bool,
76
+ vocabulary_size: int,
77
+ remove_non_linguistic_symbols: bool,
78
+ cutoff: int,
79
+ add_symbol: List[str],
80
+ cleaner: Optional[str],
81
+ g2p: Optional[str],
82
+ ):
83
+
84
+ logging.basicConfig(
85
+ level=log_level,
86
+ format="%(asctime)s (%(module)s:%(lineno)d) %(levelname)s: %(message)s",
87
+ )
88
+ if input == "-":
89
+ fin = sys.stdin
90
+ else:
91
+ fin = Path(input).open("r", encoding="utf-8")
92
+ if output == "-":
93
+ fout = sys.stdout
94
+ else:
95
+ p = Path(output)
96
+ p.parent.mkdir(parents=True, exist_ok=True)
97
+ fout = p.open("w", encoding="utf-8")
98
+
99
+ cleaner = TextCleaner(cleaner)
100
+ tokenizer = build_tokenizer(
101
+ token_type=token_type,
102
+ bpemodel=bpemodel,
103
+ delimiter=delimiter,
104
+ space_symbol=space_symbol,
105
+ non_linguistic_symbols=non_linguistic_symbols,
106
+ remove_non_linguistic_symbols=remove_non_linguistic_symbols,
107
+ g2p_type=g2p,
108
+ )
109
+
110
+ counter = Counter()
111
+ if field is not None:
112
+ field = field2slice(field)
113
+
114
+ for line in fin:
115
+ line = line.rstrip()
116
+ if field is not None:
117
+ # e.g. field="2-"
118
+ # uttidA hello world!! -> hello world!!
119
+ tokens = line.split(delimiter)
120
+ tokens = tokens[field]
121
+ if delimiter is None:
122
+ line = " ".join(tokens)
123
+ else:
124
+ line = delimiter.join(tokens)
125
+
126
+ line = cleaner(line)
127
+ tokens = tokenizer.text2tokens(line)
128
+ if not write_vocabulary:
129
+ fout.write(" ".join(tokens) + "\n")
130
+ else:
131
+ for t in tokens:
132
+ counter[t] += 1
133
+
134
+ if not write_vocabulary:
135
+ return
136
+
137
+ ## FIXME
138
+ ## del duplicate add_symbols in counter
139
+ for symbol_and_id in add_symbol:
140
+ # e.g symbol="<blank>:0"
141
+ try:
142
+ symbol, idx = symbol_and_id.split(":")
143
+ except ValueError:
144
+ raise RuntimeError(f"Format error: e.g. '<blank>:0': {symbol_and_id}")
145
+ symbol = symbol.strip()
146
+ if symbol in counter:
147
+ del counter[symbol]
148
+
149
+ # ======= write_vocabulary mode from here =======
150
+ # Sort by the number of occurrences in descending order
151
+ # and filter lower frequency words than cutoff value
152
+ words_and_counts = list(
153
+ filter(lambda x: x[1] > cutoff, sorted(counter.items(), key=lambda x: -x[1]))
154
+ )
155
+ # Restrict the vocabulary size
156
+ if vocabulary_size > 0:
157
+ if vocabulary_size < len(add_symbol):
158
+ raise RuntimeError(f"vocabulary_size is too small: {vocabulary_size}")
159
+ words_and_counts = words_and_counts[: vocabulary_size - len(add_symbol)]
160
+
161
+ # Parse the values of --add_symbol
162
+ for symbol_and_id in add_symbol:
163
+ # e.g symbol="<blank>:0"
164
+ try:
165
+ symbol, idx = symbol_and_id.split(":")
166
+ idx = int(idx)
167
+ except ValueError:
168
+ raise RuntimeError(f"Format error: e.g. '<blank>:0': {symbol_and_id}")
169
+ symbol = symbol.strip()
170
+
171
+ # e.g. idx=0 -> append as the first symbol
172
+ # e.g. idx=-1 -> append as the last symbol
173
+ if idx < 0:
174
+ idx = len(words_and_counts) + 1 + idx
175
+ words_and_counts.insert(idx, (symbol, None))
176
+
177
+ # Write words
178
+ for w, c in words_and_counts:
179
+ fout.write(w + "\n")
180
+
181
+ # Logging
182
+ total_count = sum(counter.values())
183
+ invocab_count = sum(c for w, c in words_and_counts if c is not None)
184
+ logging.info(f"OOV rate = {(total_count - invocab_count) / total_count * 100} %")
185
+
186
+
187
+ def get_parser() -> argparse.ArgumentParser:
188
+ parser = argparse.ArgumentParser(
189
+ description="Tokenize texts",
190
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
191
+ )
192
+ parser.add_argument(
193
+ "--log_level",
194
+ type=lambda x: x.upper(),
195
+ default="INFO",
196
+ choices=("CRITICAL", "ERROR", "WARNING", "INFO", "DEBUG", "NOTSET"),
197
+ help="The verbose level of logging",
198
+ )
199
+
200
+ parser.add_argument(
201
+ "--input", "-i", required=True, help="Input text. - indicates sys.stdin"
202
+ )
203
+ parser.add_argument(
204
+ "--output", "-o", required=True, help="Output text. - indicates sys.stdout"
205
+ )
206
+ parser.add_argument(
207
+ "--field",
208
+ "-f",
209
+ help="The target columns of the input text as 1-based integer. e.g 2-",
210
+ )
211
+ parser.add_argument(
212
+ "--token_type",
213
+ "-t",
214
+ default="char",
215
+ choices=["char", "bpe", "word", "phn"],
216
+ help="Token type",
217
+ )
218
+ parser.add_argument("--delimiter", "-d", default=None, help="The delimiter")
219
+ parser.add_argument("--space_symbol", default="<space>", help="The space symbol")
220
+ parser.add_argument("--bpemodel", default=None, help="The bpemodel file path")
221
+ parser.add_argument(
222
+ "--non_linguistic_symbols",
223
+ type=str_or_none,
224
+ help="non_linguistic_symbols file path",
225
+ )
226
+ parser.add_argument(
227
+ "--remove_non_linguistic_symbols",
228
+ type=str2bool,
229
+ default=False,
230
+ help="Remove non-language-symbols from tokens",
231
+ )
232
+ parser.add_argument(
233
+ "--cleaner",
234
+ type=str_or_none,
235
+ choices=[None, "tacotron", "jaconv", "vietnamese", "korean_cleaner"],
236
+ default=None,
237
+ help="Apply text cleaning",
238
+ )
239
+ parser.add_argument(
240
+ "--g2p",
241
+ type=str_or_none,
242
+ choices=g2p_classes,
243
+ default=None,
244
+ help="Specify g2p method if --token_type=phn",
245
+ )
246
+
247
+ group = parser.add_argument_group("write_vocabulary mode related")
248
+ group.add_argument(
249
+ "--write_vocabulary",
250
+ type=str2bool,
251
+ default=False,
252
+ help="Write tokens list instead of tokenized text per line",
253
+ )
254
+ group.add_argument("--vocabulary_size", type=int, default=0, help="Vocabulary size")
255
+ group.add_argument(
256
+ "--cutoff",
257
+ default=0,
258
+ type=int,
259
+ help="cut-off frequency used for write-vocabulary mode",
260
+ )
261
+ group.add_argument(
262
+ "--add_symbol",
263
+ type=str,
264
+ default=[],
265
+ action="append",
266
+ help="Append symbol e.g. --add_symbol '<blank>:0' --add_symbol '<unk>:1'",
267
+ )
268
+
269
+ return parser
270
+
271
+
272
+ def main(cmd=None):
273
+ print(get_commandline_args(), file=sys.stderr)
274
+ parser = get_parser()
275
+ args = parser.parse_args(cmd)
276
+ kwargs = vars(args)
277
+ tokenize(**kwargs)
278
+
279
+
280
+ if __name__ == "__main__":
281
+ main()
almeval/models/stepaudio/funasr_detach/bin/train.py ADDED
@@ -0,0 +1,227 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ # -*- encoding: utf-8 -*-
3
+
4
+ import os
5
+ import sys
6
+ import torch
7
+ import hydra
8
+ import logging
9
+ import argparse
10
+ from io import BytesIO
11
+ import torch.distributed as dist
12
+ from collections.abc import Sequence
13
+ from omegaconf import DictConfig, OmegaConf
14
+ from torch.nn.parallel import DistributedDataParallel as DDP
15
+ from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
16
+
17
+ from funasr_detach.register import tables
18
+ from funasr_detach.optimizers import optim_classes
19
+ from funasr_detach.train_utils.trainer import Trainer
20
+ from funasr_detach.schedulers import scheduler_classes
21
+ from funasr_detach.train_utils.initialize import initialize
22
+ from funasr_detach.download.download_from_hub import download_model
23
+ from funasr_detach.models.lora.utils import mark_only_lora_as_trainable
24
+ from funasr_detach.train_utils.set_all_random_seed import set_all_random_seed
25
+ from funasr_detach.train_utils.load_pretrained_model import load_pretrained_model
26
+
27
+ # from funasr_detach.tokenizer.build_tokenizer import build_tokenizer
28
+ # from funasr_detach.tokenizer.token_id_converter import TokenIDConverter
29
+ # from funasr_detach.tokenizer.funtoken import build_tokenizer
30
+
31
+
32
+ @hydra.main(config_name=None, version_base=None)
33
+ def main_hydra(kwargs: DictConfig):
34
+ if kwargs.get("debug", False):
35
+ import pdb
36
+
37
+ pdb.set_trace()
38
+
39
+ assert "model" in kwargs
40
+ if "model_conf" not in kwargs:
41
+ logging.info(
42
+ "download models from model hub: {}".format(kwargs.get("model_hub", "ms"))
43
+ )
44
+ kwargs = download_model(is_training=kwargs.get("is_training", True), **kwargs)
45
+
46
+ main(**kwargs)
47
+
48
+
49
+ def main(**kwargs):
50
+ print(kwargs)
51
+
52
+ # set random seed
53
+ set_all_random_seed(kwargs.get("seed", 0))
54
+ torch.backends.cudnn.enabled = kwargs.get(
55
+ "cudnn_enabled", torch.backends.cudnn.enabled
56
+ )
57
+ torch.backends.cudnn.benchmark = kwargs.get(
58
+ "cudnn_benchmark", torch.backends.cudnn.benchmark
59
+ )
60
+ torch.backends.cudnn.deterministic = kwargs.get("cudnn_deterministic", True)
61
+
62
+ local_rank = int(os.environ.get("LOCAL_RANK", 0))
63
+ if local_rank == 0:
64
+ tables.print()
65
+ # Check if we are using DDP or FSDP
66
+ use_ddp = "WORLD_SIZE" in os.environ and int(os.environ["WORLD_SIZE"]) > 1
67
+ use_fsdp = kwargs.get("use_fsdp", None)
68
+ if use_ddp or use_fsdp:
69
+ dist.init_process_group(
70
+ backend=kwargs.get("backend", "nccl"), init_method="env://"
71
+ )
72
+ torch.cuda.set_device(local_rank)
73
+
74
+ # save config.yaml
75
+ if (
76
+ (use_ddp or use_fsdp)
77
+ and dist.get_rank() == 0
78
+ or not (use_ddp or use_fsdp)
79
+ and local_rank == 0
80
+ ):
81
+ os.makedirs(kwargs.get("output_dir", "./"), exist_ok=True)
82
+ yaml_file = os.path.join(kwargs.get("output_dir", "./"), "config.yaml")
83
+ OmegaConf.save(config=kwargs, f=yaml_file)
84
+ logging.info("config.yaml is saved to: %s", yaml_file)
85
+
86
+ tokenizer = kwargs.get("tokenizer", None)
87
+ if tokenizer is not None:
88
+ tokenizer_class = tables.tokenizer_classes.get(tokenizer)
89
+ tokenizer = tokenizer_class(**kwargs["tokenizer_conf"])
90
+ kwargs["tokenizer"] = tokenizer
91
+
92
+ # build frontend if frontend is none None
93
+ frontend = kwargs.get("frontend", None)
94
+ if frontend is not None:
95
+ frontend_class = tables.frontend_classes.get(frontend)
96
+ frontend = frontend_class(**kwargs["frontend_conf"])
97
+ kwargs["frontend"] = frontend
98
+ kwargs["input_size"] = frontend.output_size()
99
+
100
+ # build model
101
+ model_class = tables.model_classes.get(kwargs["model"])
102
+ model = model_class(
103
+ **kwargs, **kwargs["model_conf"], vocab_size=len(tokenizer.token_list)
104
+ )
105
+
106
+ # init_param
107
+ init_param = kwargs.get("init_param", None)
108
+ if init_param is not None:
109
+ if not isinstance(init_param, (list, tuple)):
110
+ init_param = (init_param,)
111
+ logging.info("init_param is not None: %s", init_param)
112
+ for p in init_param:
113
+ logging.info(f"Loading pretrained params from {p}")
114
+ load_pretrained_model(
115
+ model=model,
116
+ path=p,
117
+ ignore_init_mismatch=kwargs.get("ignore_init_mismatch", True),
118
+ oss_bucket=kwargs.get("oss_bucket", None),
119
+ scope_map=kwargs.get("scope_map", None),
120
+ excludes=kwargs.get("excludes", None),
121
+ )
122
+ else:
123
+ initialize(model, kwargs.get("init", "kaiming_normal"))
124
+
125
+ # freeze_param
126
+ freeze_param = kwargs.get("freeze_param", None)
127
+ if freeze_param is not None:
128
+ freeze_param = eval(freeze_param)
129
+ if isinstance(freeze_param, Sequence):
130
+ freeze_param = (freeze_param,)
131
+ logging.info("freeze_param is not None: %s", freeze_param)
132
+ for t in freeze_param:
133
+ for k, p in model.named_parameters():
134
+ if k.startswith(t + ".") or k == t:
135
+ logging.info(f"Setting {k}.requires_grad = False")
136
+ p.requires_grad = False
137
+
138
+ if use_ddp:
139
+ model = model.cuda(local_rank)
140
+ model = DDP(
141
+ model,
142
+ device_ids=[local_rank],
143
+ find_unused_parameters=kwargs.get("train_conf", {}).get(
144
+ "find_unused_parameters", False
145
+ ),
146
+ )
147
+ elif use_fsdp:
148
+ model = FSDP(model).cuda(local_rank)
149
+ else:
150
+ model = model.to(device=kwargs.get("device", "cuda"))
151
+
152
+ # optim
153
+ optim = kwargs.get("optim", "adam")
154
+ assert optim in optim_classes
155
+ optim_class = optim_classes.get(optim)
156
+ optim = optim_class(model.parameters(), **kwargs.get("optim_conf"))
157
+
158
+ # scheduler
159
+ scheduler = kwargs.get("scheduler", "warmuplr")
160
+ assert scheduler in scheduler_classes
161
+ scheduler_class = scheduler_classes.get(scheduler)
162
+ scheduler = scheduler_class(optim, **kwargs.get("scheduler_conf"))
163
+
164
+ # dataset
165
+ dataset_class = tables.dataset_classes.get(kwargs.get("dataset", "AudioDataset"))
166
+ dataset_tr = dataset_class(
167
+ kwargs.get("train_data_set_list"),
168
+ frontend=frontend,
169
+ tokenizer=tokenizer,
170
+ is_training=True,
171
+ **kwargs.get("dataset_conf"),
172
+ )
173
+ dataset_val = dataset_class(
174
+ kwargs.get("valid_data_set_list"),
175
+ frontend=frontend,
176
+ tokenizer=tokenizer,
177
+ is_training=False,
178
+ **kwargs.get("dataset_conf"),
179
+ )
180
+
181
+ # dataloader
182
+ batch_sampler = kwargs["dataset_conf"].get(
183
+ "batch_sampler", "DynamicBatchLocalShuffleSampler"
184
+ )
185
+ batch_sampler_val = None
186
+ if batch_sampler is not None:
187
+ batch_sampler_class = tables.batch_sampler_classes.get(batch_sampler)
188
+ batch_sampler = batch_sampler_class(dataset_tr, **kwargs.get("dataset_conf"))
189
+ batch_sampler_val = batch_sampler_class(
190
+ dataset_val, is_training=False, **kwargs.get("dataset_conf")
191
+ )
192
+ dataloader_tr = torch.utils.data.DataLoader(
193
+ dataset_tr,
194
+ collate_fn=dataset_tr.collator,
195
+ batch_sampler=batch_sampler,
196
+ num_workers=kwargs.get("dataset_conf").get("num_workers", 4),
197
+ pin_memory=True,
198
+ )
199
+
200
+ dataloader_val = torch.utils.data.DataLoader(
201
+ dataset_val,
202
+ collate_fn=dataset_val.collator,
203
+ batch_sampler=batch_sampler_val,
204
+ num_workers=kwargs.get("dataset_conf").get("num_workers", 4),
205
+ pin_memory=True,
206
+ )
207
+ trainer = Trainer(
208
+ model=model,
209
+ optim=optim,
210
+ scheduler=scheduler,
211
+ dataloader_train=dataloader_tr,
212
+ dataloader_val=dataloader_val,
213
+ local_rank=local_rank,
214
+ use_ddp=use_ddp,
215
+ use_fsdp=use_fsdp,
216
+ output_dir=kwargs.get("output_dir", "./exp"),
217
+ resume=kwargs.get("resume", True),
218
+ **kwargs.get("train_conf"),
219
+ )
220
+ trainer.run()
221
+
222
+ if use_ddp or use_fsdp:
223
+ torch.distributed.destroy_process_group()
224
+
225
+
226
+ if __name__ == "__main__":
227
+ main_hydra()
almeval/models/stepaudio/funasr_detach/datasets/__init__.py ADDED
File without changes
almeval/models/stepaudio/funasr_detach/datasets/audio_datasets/__init__.py ADDED
File without changes
almeval/models/stepaudio/funasr_detach/datasets/audio_datasets/datasets.py ADDED
@@ -0,0 +1,112 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+
3
+ from funasr_detach.register import tables
4
+ from funasr_detach.utils.load_utils import extract_fbank, load_audio_text_image_video
5
+
6
+
7
+ @tables.register("dataset_classes", "AudioDataset")
8
+ class AudioDataset(torch.utils.data.Dataset):
9
+ """
10
+ AudioDataset
11
+ """
12
+
13
+ def __init__(
14
+ self,
15
+ path,
16
+ index_ds: str = None,
17
+ frontend=None,
18
+ tokenizer=None,
19
+ int_pad_value: int = -1,
20
+ float_pad_value: float = 0.0,
21
+ **kwargs
22
+ ):
23
+ super().__init__()
24
+ index_ds_class = tables.index_ds_classes.get(index_ds)
25
+ self.index_ds = index_ds_class(path, **kwargs)
26
+ preprocessor_speech = kwargs.get("preprocessor_speech", None)
27
+ if preprocessor_speech:
28
+ preprocessor_speech_class = tables.preprocessor_classes.get(
29
+ preprocessor_speech
30
+ )
31
+ preprocessor_speech = preprocessor_speech_class(
32
+ **kwargs.get("preprocessor_speech_conf")
33
+ )
34
+ self.preprocessor_speech = preprocessor_speech
35
+ preprocessor_text = kwargs.get("preprocessor_text", None)
36
+ if preprocessor_text:
37
+ preprocessor_text_class = tables.preprocessor_classes.get(preprocessor_text)
38
+ preprocessor_text = preprocessor_text_class(
39
+ **kwargs.get("preprocessor_text_conf")
40
+ )
41
+ self.preprocessor_text = preprocessor_text
42
+
43
+ self.frontend = frontend
44
+ self.fs = 16000 if frontend is None else frontend.fs
45
+ self.data_type = "sound"
46
+ self.tokenizer = tokenizer
47
+
48
+ self.int_pad_value = int_pad_value
49
+ self.float_pad_value = float_pad_value
50
+
51
+ def get_source_len(self, index):
52
+ item = self.index_ds[index]
53
+ return self.index_ds.get_source_len(item)
54
+
55
+ def get_target_len(self, index):
56
+ item = self.index_ds[index]
57
+ return self.index_ds.get_target_len(item)
58
+
59
+ def __len__(self):
60
+ return len(self.index_ds)
61
+
62
+ def __getitem__(self, index):
63
+ item = self.index_ds[index]
64
+ # import pdb;
65
+ # pdb.set_trace()
66
+ source = item["source"]
67
+ data_src = load_audio_text_image_video(source, fs=self.fs)
68
+ if self.preprocessor_speech:
69
+ data_src = self.preprocessor_speech(data_src, fs=self.fs)
70
+ speech, speech_lengths = extract_fbank(
71
+ data_src, data_type=self.data_type, frontend=self.frontend, is_final=True
72
+ ) # speech: [b, T, d]
73
+
74
+ target = item["target"]
75
+ if self.preprocessor_text:
76
+ target = self.preprocessor_text(target)
77
+ if self.tokenizer:
78
+ ids = self.tokenizer.encode(target)
79
+ text = torch.tensor(ids, dtype=torch.int64)
80
+ else:
81
+ ids = target
82
+ text = ids
83
+ ids_lengths = len(ids)
84
+ text_lengths = torch.tensor([ids_lengths], dtype=torch.int32)
85
+
86
+ return {
87
+ "speech": speech[0, :, :],
88
+ "speech_lengths": speech_lengths,
89
+ "text": text,
90
+ "text_lengths": text_lengths,
91
+ }
92
+
93
+ def collator(self, samples: list = None):
94
+ outputs = {}
95
+ for sample in samples:
96
+ for key in sample.keys():
97
+ if key not in outputs:
98
+ outputs[key] = []
99
+ outputs[key].append(sample[key])
100
+
101
+ for key, data_list in outputs.items():
102
+ if isinstance(data_list[0], torch.Tensor):
103
+ if data_list[0].dtype == torch.int64:
104
+
105
+ pad_value = self.int_pad_value
106
+ else:
107
+ pad_value = self.float_pad_value
108
+
109
+ outputs[key] = torch.nn.utils.rnn.pad_sequence(
110
+ data_list, batch_first=True, padding_value=pad_value
111
+ )
112
+ return outputs
almeval/models/stepaudio/funasr_detach/datasets/audio_datasets/index_ds.py ADDED
@@ -0,0 +1,150 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import json
3
+ import torch
4
+ import logging
5
+ import concurrent.futures
6
+ import librosa
7
+ import torch.distributed as dist
8
+
9
+ from funasr_detach.register import tables
10
+
11
+
12
+ @tables.register("index_ds_classes", "IndexDSJsonlRankSplit")
13
+ class IndexDSJsonlRankSplit(torch.utils.data.Dataset):
14
+
15
+ def __init__(self, path):
16
+ super().__init__()
17
+
18
+ contents = []
19
+ with open(path, encoding="utf-8") as fin:
20
+ for line in fin:
21
+ data = json.loads(line.strip())
22
+ if "text" in data: # for sft
23
+ self.contents.append(data["text"])
24
+ if "source" in data: # for speech lab pretrain
25
+ prompt = data["prompt"]
26
+ source = data["source"]
27
+ target = data["target"]
28
+ source_len = data["source_len"]
29
+ target_len = data["target_len"]
30
+
31
+ contents.append(
32
+ {
33
+ "source": source,
34
+ "prompt": prompt,
35
+ "target": target,
36
+ "source_len": source_len,
37
+ "target_len": target_len,
38
+ }
39
+ )
40
+
41
+ self.contents = []
42
+ total_num = len(contents)
43
+ try:
44
+ rank = dist.get_rank()
45
+ world_size = dist.get_world_size()
46
+ except:
47
+ rank = 0
48
+ world_size = 1
49
+ logging.warning("distributed is not initialized, only single shard")
50
+ num_per_rank = total_num // world_size
51
+
52
+ # rank = 0
53
+ # import ipdb; ipdb.set_trace()
54
+ self.contents = contents[rank * num_per_rank : (rank + 1) * num_per_rank]
55
+
56
+ logging.info(
57
+ "in rank: {}, num of samplers: {}, total_num of samplers across ranks: {}".format(
58
+ rank, len(self.contents), len(contents)
59
+ )
60
+ )
61
+
62
+ def __len__(self):
63
+ return len(self.contents)
64
+
65
+ def __getitem__(self, index):
66
+ try:
67
+ data = self.contents[index]
68
+ except:
69
+ print(index)
70
+ return data
71
+
72
+ def get_source_len(self, data_dict):
73
+ return data_dict["source_len"]
74
+
75
+ def get_target_len(self, data_dict):
76
+
77
+ return data_dict["target_len"] if "target_len" in data_dict else 0
78
+
79
+
80
+ @tables.register("index_ds_classes", "IndexDSJsonl")
81
+ @tables.register("index_ds_classes", "IndexDSJsonlRankFull")
82
+ class IndexDSJsonlRankFull(torch.utils.data.Dataset):
83
+
84
+ def __init__(self, path: str, **kwargs):
85
+ super().__init__()
86
+
87
+ if isinstance(path, (list, tuple)): # wav.scp, text.txt/text.trans
88
+ from funasr_detach.datasets.audio_datasets.scp2jsonl import (
89
+ gen_jsonl_from_wav_text_list,
90
+ )
91
+
92
+ jsonl_outdir = os.path.dirname(path[0])
93
+ jsonl_name = (
94
+ "datalist_train.jsonl"
95
+ if kwargs.get("is_training", True)
96
+ else "datalist_val.jsonl"
97
+ )
98
+ jsonl_file_out = os.path.join(jsonl_outdir, jsonl_name)
99
+ if not os.path.exists(jsonl_file_out):
100
+ print(f"datalist is: {path}, generate jsonl from it")
101
+ gen_jsonl_from_wav_text_list(
102
+ path, jsonl_file_out=jsonl_file_out, **kwargs
103
+ )
104
+ path = jsonl_file_out
105
+
106
+ contents = []
107
+ with open(path, encoding="utf-8") as fin:
108
+ for line in fin:
109
+ data = json.loads(line.strip())
110
+ if "text" in data: # for sft
111
+ self.contents.append(data["text"])
112
+ if "source" in data: # for speech lab pretrain
113
+ prompt = data.get("prompt", "<ASR>")
114
+ source = data["source"]
115
+ target = data["target"]
116
+ source_len = data.get("source_len", 1)
117
+ target_len = data.get("target_len", 0)
118
+
119
+ contents.append(
120
+ {
121
+ "source": source,
122
+ "prompt": prompt,
123
+ "target": target,
124
+ "source_len": source_len,
125
+ "target_len": target_len,
126
+ }
127
+ )
128
+
129
+ self.contents = contents
130
+
131
+ logging.info(
132
+ "total_num of samplers across ranks: {}".format(len(self.contents))
133
+ )
134
+
135
+ def __len__(self):
136
+ return len(self.contents)
137
+
138
+ def __getitem__(self, index):
139
+ try:
140
+ data = self.contents[index]
141
+ except:
142
+ print(index)
143
+ return data
144
+
145
+ def get_source_len(self, data_dict):
146
+ return data_dict.get("source_len", 1)
147
+
148
+ def get_target_len(self, data_dict):
149
+
150
+ return data_dict.get("target_len", 0)