Add files using upload-large-folder tool
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- .gitattributes +34 -0
- almeval/models/glm4voice/resources/architecture.jpeg +3 -0
- almeval/models/glm4voice/resources/web_demo.png +3 -0
- almeval/models/kimi_audio/assets/kimia_framework.png +3 -0
- almeval/models/kimi_audio/assets/kimia_logo.png +3 -0
- almeval/models/kimi_audio/assets/kimia_radar_chart.png +3 -0
- almeval/models/kimi_audio/assets/kimia_report.pdf +3 -0
- almeval/models/kimi_audio/kimia_infer/models/tokenizer/glm4/resources/architecture.jpeg +3 -0
- almeval/models/kimi_audio/kimia_infer/models/tokenizer/glm4/resources/web_demo.png +3 -0
- almeval/models/kimi_audio/kimia_infer/models/tokenizer/whisper_Lv3/mel_filters.npz +3 -0
- almeval/models/kimi_audio/test_audios/asr_example.wav +3 -0
- almeval/models/kimi_audio/test_audios/qa_example.wav +3 -0
- almeval/models/stepaudio/Dockerfile-vllm +35 -0
- almeval/models/stepaudio/README_JP.md +674 -0
- almeval/models/stepaudio/assets/Step-Audio.pdf +3 -0
- almeval/models/stepaudio/assets/architecture.png +3 -0
- almeval/models/stepaudio/assets/logo.png +0 -0
- almeval/models/stepaudio/assets/pipeline.png +3 -0
- almeval/models/stepaudio/assets/rlhf.png +3 -0
- almeval/models/stepaudio/assets/stepeval_radar_chart.png +3 -0
- almeval/models/stepaudio/assets/yuewen.jpeg +0 -0
- almeval/models/stepaudio/examples/clone_wav_lixueqin.wav +3 -0
- almeval/models/stepaudio/examples/clone_wav_yuqian.wav +3 -0
- almeval/models/stepaudio/examples/emotional_control1.wav +3 -0
- almeval/models/stepaudio/examples/emotional_control2.wav +3 -0
- almeval/models/stepaudio/examples/multilingual1.wav +3 -0
- almeval/models/stepaudio/examples/multilingual2.wav +3 -0
- almeval/models/stepaudio/examples/multilingual_singing.wav +3 -0
- almeval/models/stepaudio/examples/prompt_wav_lixueqin.wav +3 -0
- almeval/models/stepaudio/examples/prompt_wav_yuqian.wav +3 -0
- almeval/models/stepaudio/examples/prompt_wav_zhaobenshan.wav +3 -0
- almeval/models/stepaudio/examples/rap.wav +3 -0
- almeval/models/stepaudio/examples/singing.wav +3 -0
- almeval/models/stepaudio/examples/speed_control1.wav +3 -0
- almeval/models/stepaudio/examples/speed_control2.wav +3 -0
- almeval/models/stepaudio/examples/tone_control.wav +3 -0
- almeval/models/stepaudio/funasr_detach/__init__.py +38 -0
- almeval/models/stepaudio/funasr_detach/auto/__init__.py +0 -0
- almeval/models/stepaudio/funasr_detach/auto/auto_frontend.py +90 -0
- almeval/models/stepaudio/funasr_detach/auto/auto_model.py +573 -0
- almeval/models/stepaudio/funasr_detach/auto/auto_tokenizer.py +7 -0
- almeval/models/stepaudio/funasr_detach/bin/__init__.py +0 -0
- almeval/models/stepaudio/funasr_detach/bin/compute_audio_cmvn.py +152 -0
- almeval/models/stepaudio/funasr_detach/bin/inference.py +33 -0
- almeval/models/stepaudio/funasr_detach/bin/tokenize_text.py +281 -0
- almeval/models/stepaudio/funasr_detach/bin/train.py +227 -0
- almeval/models/stepaudio/funasr_detach/datasets/__init__.py +0 -0
- almeval/models/stepaudio/funasr_detach/datasets/audio_datasets/__init__.py +0 -0
- almeval/models/stepaudio/funasr_detach/datasets/audio_datasets/datasets.py +112 -0
- 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
|
almeval/models/glm4voice/resources/web_demo.png
ADDED
|
Git LFS Details
|
almeval/models/kimi_audio/assets/kimia_framework.png
ADDED
|
Git LFS Details
|
almeval/models/kimi_audio/assets/kimia_logo.png
ADDED
|
Git LFS Details
|
almeval/models/kimi_audio/assets/kimia_radar_chart.png
ADDED
|
Git LFS Details
|
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
|
almeval/models/kimi_audio/kimia_infer/models/tokenizer/glm4/resources/web_demo.png
ADDED
|
Git LFS Details
|
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>   |  <a href="README.md">English</a>  |   日本語 
|
| 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>  
|
| 12 |
+
<a href="https://x.com/StepFun_ai"><img src="https://img.shields.io/static/v1?label=X.com&message=Web&color=blue"></a>  
|
| 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>  
|
| 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>  
|
| 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>  
|
| 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>  
|
| 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 |
+

|
| 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 |
+

|
| 69 |
+
|
| 70 |
+
### 2.5 後トレーニングの詳細
|
| 71 |
+
後トレーニングフェーズでは、自動音声認識(ASR)およびテキストから音声への変換(TTS)のタスク固有の監督付き微調整(SFT)を実施しました。音声入力テキスト出力(AQTA)タスクについては、多様な高品質データセットを使用してSFTを実施し、人間のフィードバックからの強化学習(RLHF)を組み合わせて応答品質を向上させ、感情表現、音声速度、方言、および韻律の細かい制御を可能にしました。
|
| 72 |
+

|
| 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 (%) ↓</th>
|
| 335 |
+
<th style="text-align:center">WER (%) ↓</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 (%) ↓</th>
|
| 369 |
+
<th style="text-align:center">SS ↑</th>
|
| 370 |
+
<th style="text-align:center">WER (%) ↓</th>
|
| 371 |
+
<th style="text-align:center">SS ↑</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 (%) ↓</th>
|
| 445 |
+
<th style="text-align:center">SS ↑</th>
|
| 446 |
+
<th style="text-align:center">WER (%) ↓</th>
|
| 447 |
+
<th style="text-align:center">SS ↑</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">事実性(% ↑)</th>
|
| 487 |
+
<th style="text-align:center">関連性(% ↑)</th>
|
| 488 |
+
<th style="text-align:center">チャットスコア ↑</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
|
almeval/models/stepaudio/assets/logo.png
ADDED
|
almeval/models/stepaudio/assets/pipeline.png
ADDED
|
Git LFS Details
|
almeval/models/stepaudio/assets/rlhf.png
ADDED
|
Git LFS Details
|
almeval/models/stepaudio/assets/stepeval_radar_chart.png
ADDED
|
Git LFS Details
|
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)
|