Initial commit

This commit is contained in:
王新升
2026-02-06 20:31:14 +08:00
parent a0b51be095
commit c589bcb837
145 changed files with 28773 additions and 0 deletions
+38
View File
@@ -0,0 +1,38 @@
# Byte-compiled / optimized / DLL files
__pycache__/
dev/
results/
wandb/
.ipynb_checkpoints/
.vscode/
.cache
local/
outputs/
*.pt
*.ckpt
# Logs
logs/
*.log
results/
runs/
dev*
local/
generated/
.DS_Store
pretrained_models/
*.err
*.out
# Dev
dev/
# Data
data/
outputs/
deploy/
.gradio/
+179
View File
@@ -0,0 +1,179 @@
<div align="center">
<h1>🎤 SoulX-Singer</h1>
<p>
Official inference code for<br>
<b><em>SoulX-Singer: Towards High-Quality Zero-Shot Singing Voice Synthesis</em></b>
</p>
<p>
<img src="assets/soulx-logo.png" alt="SoulX-Logo" style="width:200px; height:68px;">
</p>
<p>
<a href="https://soul-ailab.github.io/soulx-singer/"><img src="https://img.shields.io/badge/Demo-Page-lightgrey" alt="Demo Page"></a>
<a href="https://github.com/Soul-AILab/SoulX-Singer"><img src="https://img.shields.io/badge/Github-Page-green" alt="GitHub"></a>
<a href="assets/technical-report.pdf"><img src="https://img.shields.io/badge/Report-Github-red" alt="Technical Report"></a>
<a href="https://github.com/Soul-AILab/SoulX-Singer"><img src="https://img.shields.io/badge/License-Apache%202.0-blue" alt="License"></a>
</p>
</div>
---
## 🎵 Overview
**SoulX-Singer** is a high-fidelity, zero-shot singing voice synthesis model that enables users to generate realistic singing voices for unseen singers.
It supports **melody-conditioned (F0 contour)** and **score-conditioned (MIDI notes)** control for precise pitch, rhythm, and expression.
---
## ✨ Key Features
- **🎤 Zero-Shot Singing** – Generate high-fidelity voices for unseen singers, no fine-tuning needed.
- **🎵 Flexible Control Modes** – Melody (F0) and Score (MIDI) conditioning.
- **📚 Large-Scale Dataset** – 42,000+ hours of aligned vocals, lyrics, notes across Mandarin, English, Cantonese.
- **🧑‍🎤 Timbre Cloning** – Preserve singer identity across languages, styles, and edited lyrics.
- **✏️ Singing Voice Editing** – Modify lyrics while keeping natural prosody.
- **🌐 Cross-Lingual Synthesis** – High-fidelity synthesis by disentangling timbre from content.
---
<p align="center">
<img src="assets/performance_radar.png" width="80%" alt="Performance Radar"/>
</p>
---
## 🎬 Demo Examples
<div align="center">
<a href="https://github.com/user-attachments/assets/13306f10-3a29-46ba-bcef-d6308d05cbcc">Demo 1</a><br><br>
<a href="https://github.com/user-attachments/assets/2eb260fe-6f0b-408c-aab8-5b81ddddb284">Demo 2</a>
</div>
---
## 📰 News
- **[2026-02-06]** SoulX-Singer inference code and models released.
---
## 🚀 Quick Start
### 1. Clone Repository
```bash
git clone https://github.com/Soul-AILab/SoulX-Singer.git
cd SoulX-Singer
```
### 2. Set Up Environment
**1. Install Conda** (if not already installed): https://docs.conda.io/en/latest/miniconda.html
**2. Create and activate a Conda environment:**
```
conda create -n soulxsinger -y python=3.10
conda activate soulxsinger
```
**3. Install dependencies:**
```
pip install -r requirements.txt
```
⚠️ If you are in mainland China, use a PyPI mirror:
```
pip install -r requirements.txt -i https://mirrors.aliyun.com/pypi/simple/ --trusted-host=mirrors.aliyun.com
```
---
### 3. Download Pretrained Models
Install Hugging Face Hub if needed:
```
pip install -U huggingface_hub
```
Download the SVS model and preprocessing models:
```sh
pip install -U huggingface_hub
# Download the SoulX-Singer SVS model
hf download Soul-AILab/SoulX-Singer --local-dir pretrained_models/SoulX-Singer
# Download models required for preprocessing
hf download Soul-AILab/SoulX-Singer-Preprocess --local-dir pretrained_models/SoulX-Singer-Preprocess
```
### 4. Run the Demo
Run the inference demo:
``` sh
bash example/infer.sh
```
This script relies on metadata generated from the preprocessing pipeline, including vocal separation and transcription. Users should follow the steps in [preprocess](preprocess/README.md) to prepare the necessary metadata before running the demo with their own data.
## 🚧 Roadmap
- [ ] 🖥️ Web-based UI for easy and interactive inference
- [ ] 🌐 Online demo deployment on Hugging Face Spaces
- [ ] 📊 Release the SoulX-Singer-Eval benchmark
- [ ] 📚 Comprehensive tutorials and usage documentation
## 🙏 Acknowledgements
Special thanks to the following open-source projects:
- [F5-TTS](https://github.com/SWivid/F5-TTS)
- [Amphion](https://github.com/open-mmlab/Amphion/tree/main)
- [Music Source Separation Training](https://github.com/ZFTurbo/Music-Source-Separation-Training)
- [Lead Vocal Separation](https://huggingface.co/becruily/mel-band-roformer-karaoke)
- [Vocal Dereverberation](https://huggingface.co/anvuew/dereverb_mel_band_roformer)
- [RMVPE](https://github.com/Dream-High/RMVPE)
[Paraformer](https://modelscope.cn/models/iic/speech_seaco_paraformer_large_asr_nat-zh-cn-16k-common-vocab8404-pytorch)
- [Parakeet-tdt-0.6b-v2](https://huggingface.co/nvidia/parakeet-tdt-0.6b-v2)
- [ROSVOT](https://github.com/RickyL-2000/ROSVOT)
## 📄 License
We use the Apache 2.0 license. Researchers and developers are free to use the codes and model weights of our SoulX-Singer. Check the license at [LICENSE](LICENSE) for more details.
## ⚠️ Usage Disclaimer
SoulX-Singer is intended for academic research, educational purposes, and legitimate applications such as personalized singing synthesis and assistive technologies.
Please note:
- 🎤 Respect intellectual property, privacy, and personal consent when generating singing content.
- 🚫 Do not use the model to impersonate individuals without authorization or to create deceptive audio.
- ⚠️ The developers assume no liability for any misuse of this model.
We advocate for the responsible development and use of AI and encourage the community to uphold safety and ethical principles. For ethics or misuse concerns, please contact us.
## 📬 Contact Us
We welcome your feedback, questions, and collaboration:
- **Email**: qianjiale@soulapp.cn | menghao@soulapp.cn | wangxinsheng@soulapp.cn
- **Join discussions**: WeChat or Soul APP groups for technical discussions and updates:
<p align="center">
<!-- <em>Due to group limits, if you can't scan the QR code, please add my WeChat for group access -->
<!-- : <strong>Tiamo James</strong></em> -->
<br>
<span style="display: inline-block; margin-right: 10px;">
<img src="assets/soul_wechat01.jpg" width="500" alt="WeChat Group QR Code"/>
</span>
<!-- <span style="display: inline-block;">
<img src="assets/wechat_tiamo.jpg" width="300" alt="WeChat QR Code"/>
</span> -->
</p>
Binary file not shown.

After

Width:  |  Height:  |  Size: 134 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 815 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 621 KiB

Binary file not shown.
+147
View File
@@ -0,0 +1,147 @@
import os
import torch
import json
import argparse
from tqdm import tqdm
import numpy as np
import soundfile as sf
from collections import OrderedDict
from omegaconf import DictConfig
from soulxsinger.utils.file_utils import load_config
from soulxsinger.models.soulxsinger import SoulXSinger
from soulxsinger.utils.data_processor import DataProcessor
def build_model(
model_path: str,
config: DictConfig,
device: str = "cuda",
):
"""
Build the model from the pre-trained model path and model configuration.
Args:
model_path (str): Path to the checkpoint file.
config (DictConfig): Model configuration.
device (str, optional): Device to use. Defaults to "cuda".
Returns:
Tuple[torch.nn.Module, torch.nn.Module]: The initialized model and vocoder.
"""
if not os.path.isfile(model_path):
raise FileNotFoundError(
f"Model checkpoint not found: {model_path}. "
"Please download the pretrained model and place it at the path, or set --model_path."
)
model = SoulXSinger(config).to(device)
print("Model initialized.")
print("Model parameters:", sum(p.numel() for p in model.parameters()) / 1e6, "M")
checkpoint = torch.load(model_path, weights_only=False, map_location=device)
if "state_dict" not in checkpoint:
raise KeyError(
f"Checkpoint at {model_path} has no 'state_dict' key. "
"Expected a checkpoint saved with model.state_dict()."
)
model.load_state_dict(checkpoint["state_dict"], strict=True)
model.eval()
model.to(device)
print("Model checkpoint loaded.")
return model
def process(args, config, model: torch.nn.Module):
"""Run the full inference pipeline given a data_processor and model.
"""
if args.control not in ("melody", "score"):
raise ValueError(f"control must be 'melody' or 'score', got: {args.control}")
print(f"prompt_metadata_path: {args.prompt_metadata_path}")
print(f"target_metadata_path: {args.target_metadata_path}")
os.makedirs(args.save_dir, exist_ok=True)
data_processor = DataProcessor(
hop_size=config.audio.hop_size,
sample_rate=config.audio.sample_rate,
phoneset_path=args.phoneset_path,
device=args.device,
)
with open(args.prompt_metadata_path, "r", encoding="utf-8") as f:
prompt_meta_list = json.load(f)
if not prompt_meta_list:
raise ValueError("Prompt metadata is empty. Please run preprocess on prompt audio first.")
prompt_meta = prompt_meta_list[0] # load the first segment as the prompt
with open(args.target_metadata_path, "r", encoding="utf-8") as f:
target_meta_list = json.load(f)
infer_prompt_data = data_processor.process(prompt_meta, args.prompt_wav_path)
assert len(target_meta_list) > 0, "No target segments found in the target metadata."
generated_len = int(target_meta_list[-1]["time"][1] / 1000 * config.audio.sample_rate)
generated_merged = np.zeros(generated_len, dtype=np.float32)
for idx, target_meta in enumerate(
tqdm(target_meta_list, total=len(target_meta_list), desc="Inferring segments"),
):
start_sample_idx = int(target_meta["time"][0] / 1000 * config.audio.sample_rate)
end_sample_idx = int(target_meta["time"][1] / 1000 * config.audio.sample_rate)
infer_target_data = data_processor.process(target_meta, None)
infer_data = {
"prompt": infer_prompt_data,
"target": infer_target_data,
}
with torch.no_grad():
generated_audio = model.infer(
infer_data,
auto_shift=args.auto_shift,
pitch_shift=args.pitch_shift,
n_steps=config.infer.n_steps,
cfg=config.infer.cfg,
control=args.control,
)
generated_audio = generated_audio.squeeze().cpu().numpy()
generated_merged[start_sample_idx : start_sample_idx + generated_audio.shape[0]] = generated_audio
merged_path = os.path.join(args.save_dir, "generated.wav")
sf.write(merged_path, generated_merged, 24000)
print(f"Generated audio saved to {merged_path}")
def main(args, config):
model = build_model(
model_path=args.model_path,
config=config,
device=args.device,
)
process(args, config, model)
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("--device", type=str, default="cuda")
parser.add_argument("--model_path", type=str, default='pretrained_models/soulx-singer/model.pt')
parser.add_argument("--config", type=str, default='soulxsinger/config/soulxsinger.yaml')
parser.add_argument("--prompt_wav_path", type=str, default='example/audio/zh_prompt.wav')
parser.add_argument("--prompt_metadata_path", type=str, default='example/metadata/zh_prompt.json')
parser.add_argument("--target_metadata_path", type=str, default='example/metadata/zh_target.json')
parser.add_argument("--phoneset_path", type=str, default='soulxsinger/utils/phoneme/phone_set.json')
parser.add_argument("--save_dir", type=str, default='outputs')
parser.add_argument("--auto_shift", action="store_true")
parser.add_argument("--pitch_shift", type=int, default=0)
parser.add_argument(
"--control",
type=str,
default="melody",
choices=["melody", "score"],
help="Control mode: melody or score only",
)
args = parser.parse_args()
config = load_config(args.config)
main(args, config)
+16
View File
@@ -0,0 +1,16 @@
[
{
"index": "vocal_5220_10280",
"language": "English",
"time": [
5220,
10280
],
"duration": "0.24 0.36 0.30 0.78 0.24 0.56 0.19 0.53 0.36 0.20 0.32 0.57 0.19 0.22",
"text": "<SP> Ooh Ooh <SP> I wish nothing nothing more more the best best <SP>",
"phoneme": "<SP> en_UW1 en_UW1 <SP> en_AY1 en_W-IH1-SH en_N-AH1-TH-IH0-NG en_N-AH1-TH-IH0-NG en_M-AO1-R en_M-AO1-R en_DH-AH0 en_B-EH1-S-T en_B-EH1-S-T <SP>",
"note_pitch": "0 63 65 0 65 67 68 62 62 64 67 67 65 0",
"note_type": "1 2 3 1 2 2 2 3 2 3 2 2 3 1",
"f0": "0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 345.2 343.1 341.6 339.8 337.8 331.9 319.5 312.1 310.8 312.6 315.1 316.1 315.3 314.6 315.3 317.9 322.0 329.6 337.5 344.7 347.5 347.2 344.3 339.5 338.2 341.7 342.8 342.2 340.7 343.0 342.9 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 347.0 345.3 348.7 350.2 350.9 350.3 344.7 340.3 338.3 338.0 342.8 347.4 348.3 346.7 343.4 339.5 340.5 345.2 350.4 357.7 367.3 376.6 385.9 392.6 393.6 389.9 384.7 381.8 382.0 383.0 380.6 373.5 367.9 377.0 385.4 391.4 393.8 395.6 396.1 397.2 399.8 406.0 413.5 416.1 416.0 414.4 413.5 412.9 415.5 418.9 417.5 408.8 389.2 373.9 0.0 0.0 0.0 288.5 286.0 284.2 285.6 288.9 291.3 293.5 294.5 295.2 297.8 299.5 301.0 303.0 305.9 306.8 306.0 304.4 301.8 301.0 300.8 301.8 310.2 309.8 308.2 305.9 303.6 301.5 299.3 298.5 300.0 302.1 303.5 303.6 302.2 299.7 297.5 296.3 296.4 296.8 298.6 302.6 311.8 322.0 333.8 349.0 368.8 393.3 407.1 410.7 407.0 402.3 401.2 401.7 403.9 405.7 403.5 396.8 387.4 378.6 377.8 381.4 384.0 384.7 383.5 382.5 380.8 377.3 378.4 383.5 390.0 392.7 390.5 387.6 385.3 382.7 381.0 382.8 383.9 382.2 379.6 379.3 380.2 383.1 386.0 386.5 385.4 384.3 383.7 384.4 386.2 388.2 388.5 385.0 378.6 360.4 333.7 328.2 332.4 340.2 348.9 339.6 334.9 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0"
}
]
Binary file not shown.
+16
View File
@@ -0,0 +1,16 @@
[
{
"index": "vocal_0_6900",
"language": "English",
"time": [
0,
6900
],
"duration": "0.16 0.24 0.32 0.15 0.17 0.24 0.15 0.44 0.29 0.32 0.24 0.32 0.22 0.18 0.24 0.25 1.01 0.26 0.48 0.29 0.79 0.14",
"text": "<SP> Who says you're you're not pretty <SP> pretty <SP> Who says you're you're not beautiful beautiful <SP> Who says says <SP>",
"phoneme": "<SP> en_HH-UW1 en_S-EH1-Z en_Y-UH1-R en_Y-UH1-R en_N-AA1-T en_P-R-IH1-T-IY0 <SP> en_P-R-IH1-T-IY0 <SP> en_HH-UW1 en_S-EH1-Z en_Y-UH1-R en_Y-UH1-R en_N-AA1-T en_B-Y-UW1-T-AH0-F-AH0-L en_B-Y-UW1-T-AH0-F-AH0-L <SP> en_HH-UW1 en_S-EH1-Z en_S-EH1-Z <SP>",
"note_pitch": "0 68 67 65 63 63 66 67 70 66 68 67 65 63 63 67 65 63 65 61 58 0",
"note_type": "1 2 2 2 3 2 2 1 3 1 2 2 2 3 2 2 3 1 2 2 3 1",
"f0": "0.0 0.0 382.7 387.7 385.9 379.8 376.0 380.9 390.1 403.2 415.3 423.6 421.6 402.6 385.2 381.1 0.0 0.0 425.8 419.0 409.6 397.8 392.2 389.0 388.5 391.4 389.1 381.4 375.9 0.0 0.0 0.0 0.0 359.0 354.7 353.8 353.7 354.7 353.1 351.1 350.4 349.0 348.9 346.3 337.4 328.0 312.8 303.1 298.4 296.0 298.9 302.0 306.3 307.9 307.3 307.5 307.3 302.9 301.8 0.0 0.0 0.0 0.0 0.0 343.7 364.3 375.9 368.5 358.1 359.1 365.9 378.4 393.1 406.0 412.5 410.9 407.0 404.1 403.5 403.4 401.5 399.4 397.7 395.4 394.4 394.8 395.5 396.5 397.5 400.8 407.9 415.1 417.8 453.1 472.2 481.0 482.3 481.9 480.8 478.7 477.4 476.8 474.8 467.5 446.0 390.4 382.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 374.3 375.5 370.9 370.7 373.1 378.3 392.2 407.6 418.5 423.9 423.2 415.9 395.4 0.0 0.0 0.0 421.5 416.2 405.3 391.0 383.1 380.8 383.0 388.3 388.8 378.3 371.7 0.0 0.0 0.0 371.4 365.1 362.7 358.5 353.0 352.0 353.5 356.1 356.4 353.6 348.3 341.1 330.6 317.7 303.8 293.3 296.5 297.7 301.4 305.3 308.8 308.8 308.2 308.2 306.3 305.6 285.0 269.8 265.6 280.0 304.4 331.2 351.0 357.9 364.2 370.6 381.1 392.9 399.0 399.1 395.0 389.5 379.9 363.0 338.9 318.5 305.6 300.3 299.6 296.3 292.2 0.0 0.0 0.0 0.0 0.0 0.0 0.0 309.6 322.1 329.8 331.2 332.1 332.6 332.4 335.4 340.7 345.0 347.2 346.2 342.6 339.6 337.4 338.3 340.9 342.6 344.0 344.6 344.0 344.2 343.6 341.9 338.8 336.7 337.6 341.1 347.0 350.4 343.0 326.6 330.8 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 310.0 315.7 317.7 317.4 316.5 314.4 314.1 322.5 336.3 350.0 354.2 352.9 350.5 348.4 347.1 347.3 348.4 349.8 349.8 350.4 350.3 323.9 324.7 0.0 0.0 0.0 0.0 0.0 0.0 0.0 297.3 289.8 279.5 275.1 276.1 276.1 274.9 275.4 274.6 271.8 268.6 264.0 258.3 251.7 244.3 239.9 236.1 233.7 234.0 236.0 237.3 236.9 235.2 233.5 231.7 231.0 232.1 233.6 235.4 236.2 236.7 235.8 234.1 232.2 231.3 232.6 233.5 235.2 236.0 232.3 228.8 229.6 233.8 241.3 239.4 226.3 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0"
}
]
Binary file not shown.
File diff suppressed because one or more lines are too long
Binary file not shown.
+16
View File
@@ -0,0 +1,16 @@
[
{
"index": "vocal_420_14370",
"language": "Cantonese",
"time": [
420,
14370
],
"duration": "0.31 0.26 0.28 0.26 0.40 0.20 0.42 0.24 0.36 0.24 0.32 0.26 0.94 0.32 0.24 0.30 0.34 0.22 0.34 0.90 0.22 0.36 0.32 0.30 0.22 0.36 0.22 0.32 0.34 0.20 0.40 0.24 0.30 0.38 0.22 0.32 0.28 0.36 0.24 0.34 0.26 0.60",
"text": "<SP> 我 的 心 情 又 像 真 该 等 被 揭 开 嘴 巴 却 再 仰 千 台 人 潮 内 越 文 静 越 变 得 不 受 理 睬 睬 自 己 己 要 交 出 意 外",
"phoneme": "<SP> yue_ngo5 yue_dik1 yue_sam1 yue_cing4 yue_jau6 yue_zoeng6 yue_zan1 yue_goi1 yue_dang2 yue_bei6 yue_kit3 yue_hoi1 yue_zeoi2 yue_baa1 yue_koek3 yue_zoi3 yue_joeng5 yue_cin1 yue_toi4 yue_jan4 yue_ciu4 yue_noi6 yue_jyut6 yue_man4 yue_zing6 yue_jyut6 yue_bin3 yue_dak1 yue_bat1 yue_sau6 yue_lei5 yue_coi2 yue_coi2 yue_zi6 yue_gei2 yue_gei2 yue_jiu3 yue_gaau1 yue_ceot1 yue_ji3 yue_ngoi6",
"note_pitch": "0 52 57 59 55 57 59 62 60 58 54 57 59 59 57 55 54 53 57 51 50 54 57 58 54 57 59 61 64 59 54 54 57 59 51 56 58 57 56 56 55 52",
"note_type": "1 2 2 2 2 2 2 2 2 2 2 2 2 2 2 2 2 2 2 2 2 2 2 2 2 2 2 2 2 2 2 2 2 3 2 2 3 2 2 2 2 2",
"f0": "0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 74.8 74.8 82.4 85.7 81.0 76.2 78.8 111.9 129.0 146.8 160.6 175.0 182.1 172.9 163.9 190.3 214.7 218.4 221.5 223.5 220.2 209.1 173.8 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 254.5 248.7 245.1 243.3 239.8 238.8 241.2 244.0 245.7 245.6 239.8 216.8 191.2 0.0 0.0 170.9 179.4 189.0 192.6 193.3 192.8 193.5 194.1 194.1 194.4 195.4 197.4 199.7 202.2 205.8 210.6 212.3 214.4 216.5 219.4 221.9 222.4 222.4 222.8 222.9 217.1 189.6 175.4 0.0 0.0 255.5 251.2 246.5 247.3 248.9 249.1 249.4 251.9 253.6 250.2 247.1 246.6 239.1 193.8 191.0 0.0 295.9 301.1 302.8 301.5 296.6 287.7 286.4 290.4 294.2 297.1 297.3 294.7 287.9 273.4 221.0 262.3 265.3 259.7 255.9 254.7 255.1 256.2 257.4 259.3 260.8 261.3 249.3 209.1 194.0 236.7 224.9 210.2 202.6 197.9 201.4 210.6 220.6 230.1 237.9 242.4 242.7 241.0 234.7 220.3 190.9 179.0 185.6 182.3 178.3 177.1 179.0 181.3 182.4 184.5 187.6 185.9 172.3 161.7 167.7 0.0 0.0 206.9 210.4 213.0 214.2 215.6 217.3 217.1 194.7 181.9 184.5 182.3 171.8 155.9 161.4 167.7 195.8 235.3 245.0 245.8 241.9 237.0 232.3 231.3 234.9 241.8 250.9 253.2 248.9 238.0 226.0 220.5 224.9 236.6 250.9 259.6 262.0 259.1 252.5 246.5 241.8 236.0 228.8 216.3 210.5 211.1 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 209.1 197.4 191.2 191.2 200.4 222.5 235.6 248.0 252.9 249.4 243.0 235.6 216.1 182.1 175.8 213.4 214.1 214.7 215.1 215.2 215.9 215.7 214.8 211.6 202.2 187.0 183.6 0.0 0.0 0.0 191.3 192.1 191.4 192.4 193.3 192.6 188.9 171.2 0.0 0.0 0.0 0.0 0.0 0.0 192.1 186.5 181.8 178.2 176.1 176.8 178.2 179.4 179.7 180.4 183.3 184.4 182.7 179.0 176.1 174.6 168.7 166.5 167.2 168.8 171.8 175.6 185.4 194.7 197.2 192.6 181.9 0.0 0.0 0.0 192.6 204.8 217.5 220.0 220.4 218.7 216.0 214.2 216.0 216.8 213.0 200.5 183.9 0.0 0.0 158.5 156.8 159.6 161.4 160.7 159.4 159.5 159.3 157.1 152.4 152.1 156.4 162.9 167.3 166.4 160.3 149.6 146.1 149.5 155.4 161.3 164.1 163.3 161.3 154.8 148.5 144.6 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 99.4 97.6 102.0 115.3 131.4 142.5 150.1 155.1 159.7 162.5 161.1 152.3 0.0 0.0 0.0 0.0 174.2 183.7 189.8 191.6 190.4 189.1 189.5 190.4 191.9 192.1 192.5 196.2 203.3 212.2 214.1 216.7 217.4 217.2 216.8 216.4 214.9 214.7 216.3 218.9 220.4 221.5 224.6 231.2 237.9 242.2 242.6 240.3 239.0 240.4 241.7 242.3 238.4 184.9 178.6 194.1 203.6 194.7 185.4 180.5 181.9 185.4 188.7 191.6 193.9 194.6 192.2 191.0 189.5 186.7 181.0 0.0 0.0 219.2 217.8 216.8 215.3 213.3 212.0 213.6 215.2 215.0 214.7 215.8 218.2 221.0 224.4 229.9 237.5 244.9 246.6 246.6 248.0 250.2 249.8 193.6 186.3 193.0 0.0 0.0 0.0 288.3 289.9 287.8 287.2 287.7 285.0 283.4 281.9 280.1 283.4 287.8 290.9 291.6 287.6 240.0 238.5 0.0 333.3 328.9 325.3 323.1 321.6 317.4 298.0 274.9 0.0 0.0 0.0 0.0 0.0 0.0 243.7 248.7 245.5 243.3 243.8 246.2 244.5 226.2 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 197.1 189.0 184.3 182.5 184.8 188.1 188.8 190.3 190.4 189.5 188.7 188.0 189.8 188.2 185.0 182.2 179.7 178.1 178.7 188.8 208.5 217.2 220.0 218.6 212.8 191.9 167.3 0.0 0.0 0.0 183.9 201.1 210.5 210.3 209.5 206.7 206.5 212.0 221.1 235.9 250.8 253.0 247.7 238.4 229.7 229.1 235.1 244.7 251.7 253.5 251.4 246.7 242.4 234.6 209.5 180.1 173.1 0.0 0.0 0.0 147.9 147.5 155.4 159.1 160.2 161.4 162.0 161.7 159.0 152.4 139.6 121.5 126.4 0.0 0.0 197.2 200.2 196.8 195.8 200.3 203.2 203.8 202.5 203.7 212.8 220.0 225.8 231.6 234.1 231.7 228.3 225.4 226.0 229.7 234.3 237.3 238.2 236.9 232.9 227.1 220.3 215.1 211.3 205.9 210.1 216.0 217.6 218.4 218.9 219.3 218.5 217.1 216.9 217.7 216.4 212.1 196.3 171.7 171.3 210.9 203.0 194.9 194.8 199.1 205.4 210.5 214.9 219.6 224.8 225.1 221.4 212.4 198.5 0.0 0.0 204.9 204.1 208.9 212.7 212.8 213.2 214.9 214.4 208.4 189.0 159.7 160.4 193.1 198.9 196.0 192.8 193.0 195.3 195.2 194.5 193.6 193.9 192.5 192.3 192.2 183.0 164.8 150.0 147.3 150.7 155.9 160.8 163.6 164.8 161.8 156.6 155.0 160.5 166.2 167.3 165.1 162.0 154.3 142.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0"
}
]
Binary file not shown.
+16
View File
@@ -0,0 +1,16 @@
[
{
"index": "vocal_320_10687",
"language": "Mandarin",
"time": [
320,
10687
],
"duration": "0.23 0.34 0.26 0.70 0.52 0.46 0.36 0.44 0.14 0.24 0.64 0.47 0.51 1.10 0.28 0.38 0.32 0.32 0.38 0.32 0.31 0.19 1.45",
"text": "<SP> 除 了 想 你 你 <SP> 除 了 了 爱 你 你 <SP> 我 什 么 什 么 都 愿 愿 意",
"phoneme": "<SP> zh_chu2 zh_le5 zh_xiang3 zh_ni3 zh_ni3 <SP> zh_chu2 zh_le5 zh_le5 zh_ai4 zh_ni3 zh_ni3 <SP> zh_wo3 zh_shen2 zh_me5 zh_shen2 zh_me5 zh_dou1 zh_yuan4 zh_yuan4 zh_yi4",
"note_pitch": "0 62 65 67 67 69 0 67 69 67 65 67 69 67 67 66 64 64 60 60 65 67 0",
"note_type": "1 2 2 2 2 3 1 2 2 3 2 2 3 1 2 2 2 2 2 2 2 3 2",
"f0": "0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 294.1 288.2 290.6 294.6 295.7 292.8 291.5 294.4 295.4 294.7 293.5 292.2 294.4 295.8 293.2 297.7 320.1 338.1 348.3 348.7 344.4 342.6 346.2 354.8 356.9 353.1 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 401.1 403.8 405.0 402.1 398.4 395.6 393.3 392.0 391.4 390.3 390.2 390.3 390.3 391.3 392.8 393.6 391.7 390.5 391.4 391.6 391.5 393.5 393.8 390.8 387.4 387.8 389.3 390.8 392.2 391.6 390.2 389.8 389.1 388.4 390.0 395.5 397.2 396.7 395.5 395.1 394.6 394.9 395.6 395.2 394.6 395.4 395.9 394.0 391.7 390.7 391.7 392.6 391.6 395.7 405.8 441.7 462.5 463.8 450.2 430.9 414.5 415.2 426.7 439.8 454.4 462.9 447.8 422.5 400.6 403.2 423.5 451.3 482.3 492.8 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 435.9 414.3 406.2 402.8 398.2 396.7 396.7 395.2 398.3 416.9 441.6 446.9 441.5 437.3 435.3 432.9 431.6 432.3 434.8 436.7 433.8 422.8 406.7 390.9 382.0 381.7 384.7 382.7 368.1 357.4 355.6 355.1 352.4 348.1 346.4 348.7 351.6 355.0 354.4 351.7 349.9 349.2 348.1 346.0 345.4 344.2 344.4 345.5 346.6 349.0 349.7 349.1 349.5 349.6 349.7 349.4 349.5 352.2 354.5 355.7 355.6 356.9 359.4 361.5 363.8 360.4 354.3 357.2 363.9 372.4 382.6 399.1 402.9 400.6 395.5 390.5 388.9 390.2 391.1 391.9 391.5 390.4 390.2 391.3 391.6 391.0 388.6 386.2 389.1 403.8 430.2 441.8 449.5 448.1 443.2 438.3 432.9 430.7 434.6 442.0 447.6 446.0 440.3 434.7 431.1 435.7 442.3 445.9 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 423.8 402.1 398.5 397.6 393.4 390.5 391.4 402.5 427.6 442.5 442.5 435.1 430.1 430.8 439.9 447.4 442.2 426.1 412.3 399.8 391.8 389.5 388.5 387.8 386.4 384.7 384.5 387.9 391.2 391.8 392.9 393.8 392.0 392.0 395.4 398.1 398.2 396.3 393.1 391.1 388.9 386.6 383.1 381.5 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 357.8 367.6 370.9 365.4 359.3 355.7 358.1 372.7 396.8 404.0 398.4 392.6 389.2 388.8 383.6 362.3 341.1 325.1 326.3 327.9 331.3 333.3 326.0 319.4 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 356.4 353.5 349.0 346.5 338.4 328.4 323.7 326.2 334.2 338.4 331.6 309.8 280.6 256.4 252.8 256.1 258.3 256.9 257.6 261.1 259.7 258.7 258.6 260.2 262.2 262.4 262.9 264.2 263.4 260.2 256.4 253.2 238.7 223.3 231.6 257.7 258.1 258.4 258.8 258.0 256.8 255.8 254.0 255.1 258.1 261.9 263.7 262.6 256.7 253.1 250.2 246.7 258.7 294.4 327.1 342.6 346.0 344.1 341.2 342.3 345.2 350.1 364.7 382.5 396.1 396.1 389.2 381.8 381.1 387.2 397.0 399.3 390.7 374.3 360.1 350.6 346.6 347.4 350.7 354.4 354.3 351.7 349.8 348.4 346.9 347.0 348.4 349.5 351.0 352.3 353.6 353.3 350.5 348.3 345.5 344.3 344.4 347.0 350.6 352.0 351.0 350.8 350.1 347.6 345.7 347.0 350.3 351.7 350.7 348.5 346.9 347.7 349.0 349.1 348.5 346.8 346.3 348.4 349.0 349.2 351.1 349.6 348.3 350.5 351.1 348.0 347.6 349.1 351.3 356.0 361.3 360.6 354.0 341.0 316.2 302.9 302.1 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0"
}
]
Binary file not shown.
+16
View File
@@ -0,0 +1,16 @@
[
{
"index": "vocal_0_6710",
"language": "Mandarin",
"time": [
0,
6710
],
"duration": "0.13 0.26 0.24 0.22 0.24 0.33 0.13 0.24 0.22 0.46 0.69 0.84 0.26 0.30 0.16 0.26 0.26 0.20 0.32 0.94",
"text": "<SP> 像 我 这 样 懦 懦 弱 的 人 人 <SP> 凡 事 都 要 留 留 几 分",
"phoneme": "<SP> zh_xiang4 zh_wo3 zh_zhe4 zh_yang4 zh_nuo4 zh_nuo4 zh_ruo4 zh_de5 zh_ren2 zh_ren2 <SP> zh_fan2 zh_shi4 zh_dou1 zh_yao4 zh_liu2 zh_liu2 zh_ji3 zh_fen1",
"note_pitch": "0 50 53 55 53 56 54 53 50 51 53 0 51 53 55 53 54 56 51 53",
"note_type": "1 2 2 2 2 2 3 2 2 2 3 1 2 2 2 2 2 3 2 2",
"f0": "0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 132.7 137.2 144.9 147.0 148.0 148.6 148.9 148.5 147.9 147.4 148.1 149.1 154.3 166.4 173.1 175.0 176.3 175.9 172.7 173.6 175.9 172.9 159.1 165.8 0.0 214.2 213.4 210.2 201.5 198.1 197.1 197.2 197.8 200.8 206.4 206.1 200.5 189.9 180.6 172.0 170.8 171.4 176.2 180.9 182.3 182.2 181.0 180.5 183.4 192.6 211.8 220.6 223.7 219.3 211.9 207.0 203.6 202.6 204.2 204.4 204.1 202.1 198.0 192.9 185.9 177.9 174.1 174.6 174.5 173.8 173.6 172.4 168.3 168.3 172.7 173.2 171.9 170.7 170.2 169.9 170.6 173.1 172.4 164.2 148.2 147.8 152.3 148.5 143.8 145.6 149.2 149.9 150.1 152.5 153.6 154.7 156.1 155.0 152.4 152.1 153.7 155.3 156.4 156.8 157.3 157.7 157.1 156.8 157.8 158.9 157.9 157.5 157.1 157.0 159.1 162.0 167.7 172.1 174.9 176.2 174.5 172.0 170.9 171.3 172.5 173.1 173.5 173.1 174.1 174.6 175.2 176.7 177.2 177.3 176.9 175.9 174.1 172.4 174.1 174.8 171.8 172.1 176.5 177.3 176.0 179.4 179.9 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 137.5 138.5 145.8 151.0 153.3 154.3 156.2 158.5 161.0 158.1 148.6 151.3 157.0 163.5 172.7 176.4 176.1 175.8 175.6 174.7 171.4 169.3 161.7 194.2 199.3 199.7 201.0 202.7 201.4 199.4 198.5 196.6 194.0 190.6 186.1 181.9 179.4 177.4 177.1 176.3 175.8 175.5 174.8 173.6 172.4 170.8 165.0 160.6 175.9 179.3 179.8 180.4 180.3 178.8 178.6 181.5 185.0 190.7 198.3 206.3 210.6 210.9 207.4 203.5 203.4 204.6 203.5 195.6 182.5 0.0 0.0 0.0 0.0 144.7 144.1 146.6 150.2 151.9 153.4 155.0 156.0 155.9 155.4 155.2 153.5 147.0 144.8 0.0 0.0 0.0 0.0 0.0 0.0 181.1 178.9 178.0 177.5 176.0 173.2 172.9 172.9 174.1 176.2 177.3 178.7 178.8 176.0 175.2 175.1 176.3 178.1 177.6 177.4 177.9 177.6 177.3 177.3 177.5 176.6 175.7 176.6 177.5 177.2 175.9 174.8 173.5 174.0 175.7 177.4 177.8 174.7"
}
]
Binary file not shown.
+28
View File
@@ -0,0 +1,28 @@
#!/bin/bash
script_dir=$(dirname "$(realpath "$0")")
root_dir=$(dirname "$script_dir")
cd $root_dir || exit
export PYTHONPATH=$root_dir:$PYTHONPATH
model_path=pretrained_models/SoulX-Singer/model.pt
config=soulxsinger/config/soulxsinger.yaml
prompt_wav_path=example/audio/zh_prompt.mp3
prompt_metadata_path=example/audio/zh_prompt.json
target_metadata_path=example/audio/music.json
phoneset_path=soulxsinger/utils/phoneme/phone_set.json
save_dir=example/generated/music
control=score # melody or score
python -m cli.inference \
--device cuda \
--model_path $model_path \
--config $config \
--prompt_wav_path $prompt_wav_path \
--prompt_metadata_path $prompt_metadata_path \
--target_metadata_path $target_metadata_path \
--phoneset_path $phoneset_path \
--save_dir $save_dir \
--auto_shift \
--pitch_shift 0
+41
View File
@@ -0,0 +1,41 @@
#!/bin/bash
script_dir=$(dirname "$(realpath "$0")")
root_dir=$(dirname "$script_dir")
cd $root_dir || exit
export PYTHONPATH=$root_dir:$PYTHONPATH
device=cuda
####### Run Prompt Annotation #######
audio_path=example/audio/zh_prompt.mp3
save_dir=example/transcriptions/zh_prompt
language=Mandarin
vocal_sep=False
max_merge_duration=30000
python -m preprocess.pipeline \
--audio_path $audio_path \
--save_dir $save_dir \
--language $language \
--device $device \
--vocal_sep $vocal_sep \
--max_merge_duration $max_merge_duration
####### Run Target Annotation #######
audio_path=example/audio/music.mp3
save_dir=example/transcriptions/music
language=Mandarin
vocal_sep=True
max_merge_duration=60000
python -m preprocess.pipeline \
--audio_path $audio_path \
--save_dir $save_dir \
--language $language \
--device $device \
--vocal_sep $vocal_sep \
--max_merge_duration $max_merge_duration
+125
View File
@@ -0,0 +1,125 @@
# 🎵 SoulX-Singer-Preprocess
This part offers a comprehensive **singing transcription and editing toolkit** for real-world music audio. It provides the pipeline from vocal extraction to high-level annotation optimized for SVS dataset construction. By integrating state-of-the-art models, it transforms raw audio into structured singing data and supports the **customizable creation and editing of lyric-aligned MIDI scores**.
## ✨ Features
The toolkit includes the following core modules:
- 🎤 **Clean Dry Vocal Extraction**
Extracts the lead vocal track from polyphonic music audio and dereverberation.
- 📝 **Lyrics Transcription**
Automatically transcribes lyrics from clean vocal.
- 🎶 **Note Transcription**
Converts singing voice into note-level representations for SVS.
- 🎼 **MIDI Editor**
Supports customizable creation and editing of MIDI scores integrated with lyrics.
## 📁 Data Preparation
To ensure the data processing pipeline runs correctly, please verify that all required checkpoints are correctly placed in `pretrained_models/SoulX-Singer-Preprocess`
Before running the pipeline, prepare the following inputs:
- **Prompt audio**
Reference audio that provides timbre and style
- **Target audio**
Original vocal or music audio to be processed and transcribed.
Configure the corresponding parameters in:
```
example/preprocess.sh
```
Typical configuration includes:
- Input / output paths
- Module enable switches
## 🚀 Usage
After configuring `preprocess.sh`, run the transcription pipeline with:
```bash
bash example/preprocess.sh
```
The script will automatically execute the following steps:
1. **Vocal separation and dereverberation**
2. **F0 extraction and voice activity detection (VAD)**
3. **Lyrics transcription**
4. **Note transcription**
---
After the pipeline completes, you will obtain **SoulX-Singer–style metadata** that can be directly used for Singing Voice Synthesis (SVS).
⚠️ **Important Note**
Transcription errors—especially in **lyrics** and **note annotations**—can significantly affect the final SVS quality. We **strongly recommend manually reviewing and correcting** the generated metadata before inference.
To support this, we provide a **MIDI Editor** for editing lyrics, phoneme alignment, note pitches, and durations. The workflow is:
**Export metadata to MIDI** → edit in the MIDI Editor → **Import edited MIDI back to metadata** for SVS.
---
#### Step 1: Metadata → MIDI (for editing)
Convert SoulX-Singer metadata to a MIDI file so you can open it in the MIDI Editor:
```bash
preprocess_root=example/transcriptions/music
python -m preprocess.tools.midi_parser \
--meta2midi \
--meta "${preprocess_root}/metadata.json" \
--midi "${preprocess_root}/vocal.mid"
```
#### Step 2: Edit in the MIDI Editor
Open the MIDI Editor (see [MIDI Editor Tutorial](tools/midi_editor/README.md)), load `vocal.mid`, and correct lyrics, pitches, or durations as needed. Save the result as e.g. `vocal_edited.mid`.
#### Step 3: MIDI → Metadata (for SoulX-Singer inference)
Convert the edited MIDI back into SoulX-Singer-style metadata (and cut wavs) for SVS:
```bash
python -m preprocess.tools.midi_parser \
--midi2meta \
--midi "${preprocess_root}/vocal_edited.mid" \
--meta "${preprocess_root}/edit_metadata.json" \
--vocal "${preprocess_root}/vocal.wav" \
```
Use `edit_metadata.json` (and the wavs under `edit_cut_wavs`) as the target metadata in your inference pipeline.
## 🔗 References & Dependencies
This project builds upon the following excellent open-source works:
### 🎧 Vocal Separation & Dereverberation
- [Music Source Separation Training](https://github.com/ZFTurbo/Music-Source-Separation-Training)
- [Lead Vocal Separation](https://huggingface.co/becruily/mel-band-roformer-karaoke)
- [Vocal Dereverberation](https://huggingface.co/anvuew/dereverb_mel_band_roformer)
### 🎼 F0 Extraction
- [RMVPE](https://github.com/Dream-High/RMVPE)
### 📝 Lyrics Transcription (ASR)
- [Paraformer](https://modelscope.cn/models/iic/speech_seaco_paraformer_large_asr_nat-zh-cn-16k-common-vocab8404-pytorch)
- [Parakeet-tdt-0.6b-v2](https://huggingface.co/nvidia/parakeet-tdt-0.6b-v2)
### 🎶 Note Transcription
- [ROSVOT](https://github.com/RickyL-2000/ROSVOT)
We sincerely thank the authors of these repositories for their exceptional open-source contributions, which have been fundamental to the development of this toolkit.
+146
View File
@@ -0,0 +1,146 @@
import json
import shutil
import soundfile as sf
from pathlib import Path
import librosa
from preprocess.utils import convert_metadata, merge_short_segments
from preprocess.tools import (
F0Extractor,
VocalDetector,
VocalSeparator,
NoteTranscriber,
LyricTranscriber,
)
class PreprocessPipeline:
def __init__(self, device: str, language: str, save_dir: str, vocal_sep: bool = True, max_merge_duration: int = 60000):
self.device = device
self.language = language
self.save_dir = save_dir
self.vocal_sep = vocal_sep
self.max_merge_duration = max_merge_duration
if vocal_sep:
self.vocal_separator = VocalSeparator(
sep_model_path="pretrained_models/SoulX-Singer-Preprocess/mel-band-roformer-karaoke/mel_band_roformer_karaoke_becruily.ckpt",
sep_config_path="pretrained_models/SoulX-Singer-Preprocess/mel-band-roformer-karaoke/config_karaoke_becruily.yaml",
der_model_path="pretrained_models/SoulX-Singer-Preprocess/dereverb_mel_band_roformer/dereverb_mel_band_roformer_anvuew_sdr_19.1729.ckpt",
der_config_path="pretrained_models/SoulX-Singer-Preprocess/dereverb_mel_band_roformer/dereverb_mel_band_roformer_anvuew.yaml",
device=device
)
else:
self.vocal_separator = None
self.f0_extractor = F0Extractor(
model_path="pretrained_models/SoulX-Singer-Preprocess/rmvpe/rmvpe.pt",
device=device,
)
self.vocal_detector = VocalDetector(
cut_wavs_output_dir= f"{save_dir}/cut_wavs",
)
self.lyric_transcriber = LyricTranscriber(
zh_model_path="pretrained_models/SoulX-Singer-Preprocess/speech_seaco_paraformer_large_asr_nat-zh-cn-16k-common-vocab8404-pytorch",
en_model_path="pretrained_models/SoulX-Singer-Preprocess/parakeet-tdt-0.6b-v2/parakeet-tdt-0.6b-v2.nemo",
device=device
)
self.note_transcriber = NoteTranscriber(
rosvot_model_path="pretrained_models/SoulX-Singer-Preprocess/rosvot/rosvot/model.pt",
rwbd_model_path="pretrained_models/SoulX-Singer-Preprocess/rosvot/rwbd/model.pt",
device=device
)
def run(
self,
audio_path: str,
vocal_sep: bool = True,
max_merge_duration: int = 60000,
language: str = "Mandarin"
) -> None:
vocal_sep = self.vocal_sep if vocal_sep is None else vocal_sep
max_merge_duration = self.max_merge_duration if max_merge_duration is None else max_merge_duration
language = self.language if language is None else language
output_dir = Path(self.save_dir)
output_dir.mkdir(parents=True, exist_ok=True)
if vocal_sep:
# Perform vocal/accompaniment separation
sep = self.vocal_separator.process(audio_path)
vocal = sep.vocals_dereverbed.T
acc = sep.accompaniment.T
sample_rate = sep.sample_rate
vocal_path = output_dir / "vocal.wav"
acc_path = output_dir / "acc.wav"
sf.write(vocal_path, vocal, sample_rate)
sf.write(acc_path, acc, sample_rate)
else:
# Use the original audio as vocal source (no separation)
vocal, sample_rate = librosa.load(audio_path, sr=None, mono=True)
vocal_path = output_dir / "vocal.wav"
sf.write(vocal_path, vocal, sample_rate)
vocal_f0 = self.f0_extractor.process(str(vocal_path))
segments = self.vocal_detector.process(str(vocal_path), f0=vocal_f0)
metadata = []
for seg in segments:
self.f0_extractor.process(seg["wav_fn"], f0_path=seg["wav_fn"].replace(".wav", "_f0.npy"))
words, durs = self.lyric_transcriber.process(
seg["wav_fn"], language
)
seg["words"] = words
seg["word_durs"] = durs
seg["language"] = language
metadata.append(
self.note_transcriber.process(seg, segment_info=seg)
)
merged = merge_short_segments(
vocal,
sample_rate,
metadata,
output_dir / "long_cut_wavs",
max_duration_ms=max_merge_duration,
)
final_metadata = []
for item in merged:
self.f0_extractor.process(item.wav_fn, f0_path=item.wav_fn.replace(".wav", "_f0.npy"))
final_metadata.append(convert_metadata(item))
with open(output_dir / "metadata.json", "w", encoding="utf-8") as f:
json.dump(final_metadata, f, ensure_ascii=False, indent=2)
shutil.copy(output_dir / "metadata.json", audio_path.replace(".wav", ".json").replace(".mp3", ".json").replace(".flac", ".json"))
def main(args):
pipeline = PreprocessPipeline(
device=args.device,
language=args.language,
save_dir=args.save_dir,
vocal_sep=args.vocal_sep,
max_merge_duration=args.max_merge_duration,
)
pipeline.run(
audio_path=args.audio_path,
language=args.language
)
if __name__ == "__main__":
import argparse
parser = argparse.ArgumentParser()
parser.add_argument("--audio_path", type=str, required=True, help="Path to the input audio file")
parser.add_argument("--save_dir", type=str, required=True, help="Directory to save the output files")
parser.add_argument("--language", type=str, default="Mandarin", help="Language of the audio")
parser.add_argument("--device", type=str, default="cuda:0", help="Device to run the models on")
parser.add_argument("--vocal_sep", type=bool, default=True, help="Whether to perform vocal separation")
parser.add_argument("--max_merge_duration", type=int, default=60000, help="Maximum merged segment duration in milliseconds")
args = parser.parse_args()
main(args)
+53
View File
@@ -0,0 +1,53 @@
"""Preprocess tools.
This package provides a thin, stable import surface for common preprocess components.
Examples:
from preprocess.tools import (
F0Extractor,
PitchExtractor,
VocalDetectionModel,
VocalSeparationModel,
VocalExtractionModel,
NoteTranscriptionModel,
LyricTranscriptionModel,
)
Note:
Keep these imports lightweight. If a tool pulls heavy dependencies at import time,
consider switching to lazy imports.
"""
from __future__ import annotations
# Core tools
from .f0_extraction import F0Extractor
from .vocal_detection import VocalDetector
# Some tools may live outside this package in different layouts across branches.
# Keep the public surface stable while avoiding hard import failures.
try:
from .vocal_separation.model import VocalSeparator # type: ignore
except Exception: # pragma: no cover
VocalSeparator = None # type: ignore
try:
from .note_transcription.model import NoteTranscriber # type: ignore
except Exception: # pragma: no cover
NoteTranscriber = None # type: ignore
try:
from .lyric_transcription import LyricTranscriber
except Exception: # pragma: no cover
LyricTranscriber = None # type: ignore
__all__ = [
"F0Extractor",
"VocalDetector",
]
if VocalSeparator is not None:
__all__.append("VocalSeparator")
if LyricTranscriber is not None:
__all__.append("LyricTranscriber")
if NoteTranscriber is not None:
__all__.append("NoteTranscriber")
+527
View File
@@ -0,0 +1,527 @@
# https://github.com/Dream-High/RMVPE
import math
import time
import librosa
import numpy as np
from librosa.filters import mel
from scipy.interpolate import interp1d
from typing import Optional
import torch
import torch.nn as nn
import torch.nn.functional as F
class BiGRU(nn.Module):
def __init__(self, input_features, hidden_features, num_layers):
super(BiGRU, self).__init__()
self.gru = nn.GRU(
input_features,
hidden_features,
num_layers=num_layers,
batch_first=True,
bidirectional=True,
)
def forward(self, x):
return self.gru(x)[0]
class ConvBlockRes(nn.Module):
def __init__(self, in_channels, out_channels, momentum=0.01):
super(ConvBlockRes, self).__init__()
self.conv = nn.Sequential(
nn.Conv2d(
in_channels=in_channels,
out_channels=out_channels,
kernel_size=(3, 3),
stride=(1, 1),
padding=(1, 1),
bias=False,
),
nn.BatchNorm2d(out_channels, momentum=momentum),
nn.ReLU(),
nn.Conv2d(
in_channels=out_channels,
out_channels=out_channels,
kernel_size=(3, 3),
stride=(1, 1),
padding=(1, 1),
bias=False,
),
nn.BatchNorm2d(out_channels, momentum=momentum),
nn.ReLU(),
)
if in_channels != out_channels:
self.shortcut = nn.Conv2d(in_channels, out_channels, (1, 1))
def forward(self, x):
if not hasattr(self, "shortcut"):
return self.conv(x) + x
else:
return self.conv(x) + self.shortcut(x)
class ResEncoderBlock(nn.Module):
def __init__(self, in_channels, out_channels, kernel_size, n_blocks=1, momentum=0.01):
super(ResEncoderBlock, self).__init__()
self.n_blocks = n_blocks
self.conv = nn.ModuleList()
self.conv.append(ConvBlockRes(in_channels, out_channels, momentum))
for i in range(n_blocks - 1):
self.conv.append(ConvBlockRes(out_channels, out_channels, momentum))
self.kernel_size = kernel_size
if self.kernel_size is not None:
self.pool = nn.AvgPool2d(kernel_size=kernel_size)
def forward(self, x):
for conv in self.conv:
x = conv(x)
if self.kernel_size is not None:
return x, self.pool(x)
else:
return x
class Encoder(nn.Module):
def __init__(self, in_channels, in_size, n_encoders, kernel_size, n_blocks, out_channels=16, momentum=0.01):
super(Encoder, self).__init__()
self.n_encoders = n_encoders
self.bn = nn.BatchNorm2d(in_channels, momentum=momentum)
self.layers = nn.ModuleList()
self.latent_channels = []
for i in range(self.n_encoders):
self.layers.append(
ResEncoderBlock(in_channels, out_channels, kernel_size, n_blocks, momentum=momentum)
)
self.latent_channels.append([out_channels, in_size])
in_channels = out_channels
out_channels *= 2
in_size //= 2
self.out_size = in_size
self.out_channel = out_channels
def forward(self, x):
concat_tensors = []
x = self.bn(x)
for layer in self.layers:
t, x = layer(x)
concat_tensors.append(t)
return x, concat_tensors
class Intermediate(nn.Module):
def __init__(self, in_channels, out_channels, n_inters, n_blocks, momentum=0.01):
super(Intermediate, self).__init__()
self.n_inters = n_inters
self.layers = nn.ModuleList()
self.layers.append(ResEncoderBlock(in_channels, out_channels, None, n_blocks, momentum))
for i in range(self.n_inters - 1):
self.layers.append(ResEncoderBlock(out_channels, out_channels, None, n_blocks, momentum))
def forward(self, x):
for layer in self.layers:
x = layer(x)
return x
class ResDecoderBlock(nn.Module):
def __init__(self, in_channels, out_channels, stride, n_blocks=1, momentum=0.01):
super(ResDecoderBlock, self).__init__()
out_padding = (0, 1) if stride == (1, 2) else (1, 1)
self.n_blocks = n_blocks
self.conv1 = nn.Sequential(
nn.ConvTranspose2d(
in_channels=in_channels,
out_channels=out_channels,
kernel_size=(3, 3),
stride=stride,
padding=(1, 1),
output_padding=out_padding,
bias=False,
),
nn.BatchNorm2d(out_channels, momentum=momentum),
nn.ReLU(),
)
self.conv2 = nn.ModuleList()
self.conv2.append(ConvBlockRes(out_channels * 2, out_channels, momentum))
for i in range(n_blocks - 1):
self.conv2.append(ConvBlockRes(out_channels, out_channels, momentum))
def forward(self, x, concat_tensor):
x = self.conv1(x)
x = torch.cat((x, concat_tensor), dim=1)
for conv2 in self.conv2:
x = conv2(x)
return x
class Decoder(nn.Module):
def __init__(self, in_channels, n_decoders, stride, n_blocks, momentum=0.01):
super(Decoder, self).__init__()
self.layers = nn.ModuleList()
self.n_decoders = n_decoders
for i in range(self.n_decoders):
out_channels = in_channels // 2
self.layers.append(
ResDecoderBlock(in_channels, out_channels, stride, n_blocks, momentum)
)
in_channels = out_channels
def forward(self, x, concat_tensors):
for i, layer in enumerate(self.layers):
x = layer(x, concat_tensors[-1 - i])
return x
class DeepUnet(nn.Module):
def __init__(self, kernel_size, n_blocks, en_de_layers=5, inter_layers=4, in_channels=1, en_out_channels=16):
super(DeepUnet, self).__init__()
self.encoder = Encoder(in_channels, 128, en_de_layers, kernel_size, n_blocks, en_out_channels)
self.intermediate = Intermediate(
self.encoder.out_channel // 2,
self.encoder.out_channel,
inter_layers,
n_blocks,
)
self.decoder = Decoder(self.encoder.out_channel, en_de_layers, kernel_size, n_blocks)
def forward(self, x):
x, concat_tensors = self.encoder(x)
x = self.intermediate(x)
x = self.decoder(x, concat_tensors)
return x
class E2E(nn.Module):
def __init__(self, n_blocks, n_gru, kernel_size, en_de_layers=5, inter_layers=4, in_channels=1, en_out_channels=16):
super(E2E, self).__init__()
self.unet = DeepUnet(kernel_size, n_blocks, en_de_layers, inter_layers, in_channels, en_out_channels)
self.cnn = nn.Conv2d(en_out_channels, 3, (3, 3), padding=(1, 1))
if n_gru:
self.fc = nn.Sequential(
BiGRU(3 * 128, 256, n_gru),
nn.Linear(512, 360),
nn.Dropout(0.25),
nn.Sigmoid(),
)
else:
self.fc = nn.Sequential(
nn.Linear(3 * 128, 360),
nn.Dropout(0.25),
nn.Sigmoid()
)
def forward(self, mel):
mel = mel.transpose(-1, -2).unsqueeze(1)
x = self.cnn(self.unet(mel)).transpose(1, 2).flatten(-2)
x = self.fc(x)
return x
class MelSpectrogram(torch.nn.Module):
def __init__(self, is_half, n_mel_channels, sampling_rate, win_length, hop_length,
n_fft=None, mel_fmin=0, mel_fmax=None, clamp=1e-5):
super().__init__()
n_fft = win_length if n_fft is None else n_fft
self.hann_window = {}
mel_basis = mel(
sr=sampling_rate,
n_fft=n_fft,
n_mels=n_mel_channels,
fmin=mel_fmin,
fmax=mel_fmax,
htk=True,
)
mel_basis = torch.from_numpy(mel_basis).float()
self.register_buffer("mel_basis", mel_basis)
self.n_fft = win_length if n_fft is None else n_fft
self.hop_length = hop_length
self.win_length = win_length
self.sampling_rate = sampling_rate
self.n_mel_channels = n_mel_channels
self.clamp = clamp
self.is_half = is_half
def forward(self, audio, keyshift=0, speed=1, center=True):
factor = 2 ** (keyshift / 12)
n_fft_new = int(np.round(self.n_fft * factor))
win_length_new = int(np.round(self.win_length * factor))
hop_length_new = int(np.round(self.hop_length * speed))
keyshift_key = str(keyshift) + "_" + str(audio.device)
if keyshift_key not in self.hann_window:
self.hann_window[keyshift_key] = torch.hann_window(win_length_new).to(audio.device)
fft = torch.stft(
audio,
n_fft=n_fft_new,
hop_length=hop_length_new,
win_length=win_length_new,
window=self.hann_window[keyshift_key],
center=center,
return_complex=True,
)
magnitude = torch.sqrt(fft.real.pow(2) + fft.imag.pow(2))
if keyshift != 0:
size = self.n_fft // 2 + 1
resize = magnitude.size(1)
if resize < size:
magnitude = F.pad(magnitude, (0, 0, 0, size - resize))
magnitude = magnitude[:, :size, :] * self.win_length / win_length_new
mel_output = torch.matmul(self.mel_basis, magnitude)
if self.is_half:
mel_output = mel_output.half()
log_mel_spec = torch.log(torch.clamp(mel_output, min=self.clamp))
return log_mel_spec
class RMVPE:
def __init__(self, model_path: str, is_half, device=None):
self.is_half = is_half
if device is None:
device = "cuda:0" if torch.cuda.is_available() else "cpu"
self.device = torch.device(device) if isinstance(device, str) else device
self.mel_extractor = MelSpectrogram(
is_half=is_half,
n_mel_channels=128,
sampling_rate=16000,
win_length=1024,
hop_length=160,
n_fft=None,
mel_fmin=30,
mel_fmax=8000
).to(self.device)
model = E2E(n_blocks=4, n_gru=1, kernel_size=(2, 2))
ckpt = torch.load(model_path, map_location=self.device)
model.load_state_dict(ckpt)
model.eval()
if is_half:
model = model.half()
else:
model = model.float()
self.model = model.to(self.device)
cents_mapping = 20 * np.arange(360) + 1997.3794084376191
self.cents_mapping = np.pad(cents_mapping, (4, 4)) # 368
def mel2hidden(self, mel):
with torch.no_grad():
n_frames = mel.shape[-1]
n_pad = 32 * ((n_frames - 1) // 32 + 1) - n_frames
if n_pad > 0:
mel = F.pad(mel, (0, n_pad), mode="constant")
mel = mel.half() if self.is_half else mel.float()
hidden = self.model(mel)
return hidden[:, :n_frames]
def decode(self, hidden, thred=0.03):
cents_pred = self.to_local_average_cents(hidden, thred=thred)
f0 = 10 * (2 ** (cents_pred / 1200))
f0[f0 == 10] = 0
return f0
def infer_from_audio(self, audio, thred=0.03):
if not torch.is_tensor(audio):
audio = torch.from_numpy(audio)
mel = self.mel_extractor(audio.float().to(self.device).unsqueeze(0), center=True)
hidden = self.mel2hidden(mel)
hidden = hidden.squeeze(0).cpu().numpy()
if self.is_half:
hidden = hidden.astype("float32")
f0 = self.decode(hidden, thred=thred)
return f0
def to_local_average_cents(self, salience, thred=0.05):
center = np.argmax(salience, axis=1)
salience = np.pad(salience, ((0, 0), (4, 4)))
center += 4
todo_salience = []
todo_cents_mapping = []
starts = center - 4
ends = center + 5
for idx in range(salience.shape[0]):
todo_salience.append(salience[:, starts[idx]:ends[idx]][idx])
todo_cents_mapping.append(self.cents_mapping[starts[idx]:ends[idx]])
todo_salience = np.array(todo_salience)
todo_cents_mapping = np.array(todo_cents_mapping)
product_sum = np.sum(todo_salience * todo_cents_mapping, 1)
weight_sum = np.sum(todo_salience, 1)
devided = product_sum / weight_sum
maxx = np.max(salience, axis=1)
devided[maxx <= thred] = 0
return devided
class F0Extractor:
"""Extract frame-level f0 from singing voice.
Wrapper around an RMVPE network that:
1) loads the checkpoint once in ``__init__``
2) exposes a simple :py:meth:`process` API and optionally saves ``*_f0.npy``.
"""
def __init__(
self,
model_path: str,
device: str = "cpu",
*,
is_half: bool = False,
input_sr: int = 16000,
target_sr: int = 24000,
hop_size: int = 480,
max_duration: float = 300,
thred: float = 0.03,
verbose: bool = True,
):
"""Initialize the f0 extractor.
Args:
model_path: Path to RMVPE checkpoint.
device: Torch device string, e.g. ``"cuda:0"`` / ``"cpu"``.
is_half: Whether to run the model in fp16.
input_sr: Input resample rate used by RMVPE frontend.
target_sr: Target sample rate for the output f0 grid.
hop_size: Target hop size for the output f0 grid.
max_duration: Max duration (seconds) for interpolation grid.
thred: Voicing threshold used when decoding salience.
verbose: Whether to print verbose logs.
"""
self.model_path = model_path
self.input_sr = input_sr
self.target_sr = target_sr
self.hop_size = hop_size
self.max_duration = max_duration
self.thred = thred
self.verbose = verbose
self.model = RMVPE(model_path, is_half=is_half, device=device)
if self.verbose:
print(
"[f0 extraction] init success:",
f"device={device}",
f"model_path={model_path}",
f"is_half={is_half}",
f"input_sr={input_sr}",
f"target_sr={target_sr}",
f"hop_size={hop_size}",
f"thred={thred}",
)
@staticmethod
def interpolate_f0(
f0_16k: np.ndarray,
original_length: int,
original_sr: int,
*,
target_sr: int = 48000,
hop_size: int = 256,
max_duration: float = 20.0,
) -> np.ndarray:
"""Interpolate f0 from RMVPE's 16k hop grid to target mel hop grid."""
mel_target_sr = target_sr
mel_hop_size = hop_size
mel_max_duration = max_duration
batch_max_length = int(mel_max_duration * mel_target_sr / mel_hop_size)
duration_in_seconds = original_length / original_sr
effective_target_length = int(duration_in_seconds * mel_target_sr)
original_frames = math.ceil(effective_target_length / mel_hop_size)
target_frames = min(original_frames, batch_max_length)
rmvpe_hop = 160
t_16k = np.arange(len(f0_16k)) * (rmvpe_hop / 16000.0)
t_target = np.arange(target_frames) * (mel_hop_size / float(mel_target_sr))
if len(f0_16k) > 0:
f_interp = interp1d(
t_16k,
f0_16k,
kind="linear",
bounds_error=False,
fill_value=0.0,
assume_sorted=True,
)
f0 = f_interp(t_target)
else:
f0 = np.zeros(target_frames)
if len(f0) != target_frames:
f0 = (
f0[:target_frames]
if len(f0) > target_frames
else np.pad(f0, (0, target_frames - len(f0)), "constant")
)
return f0
def process(self, audio_path: str, *, f0_path: str | None = None, verbose: Optional[bool] = None) -> np.ndarray:
"""Run f0 extraction for a single wav.
Args:
audio_path: Path to the input wav file.
f0_path: if is not None, save the f0 data to this path.
verbose: Override instance-level verbose flag for this call.
Returns:
np.ndarray: shape ``[T]``, f0 in Hz (0 for unvoiced).
"""
verbose = self.verbose if verbose is None else verbose
if verbose:
print(f"[f0 extraction] process: start: {audio_path}")
t0 = time.time()
audio, _ = librosa.load(audio_path, sr=self.input_sr)
f0_16k = self.model.infer_from_audio(audio, thred=self.thred)
f0 = self.interpolate_f0(
f0_16k,
original_length=audio.shape[-1],
original_sr=self.input_sr,
target_sr=self.target_sr,
hop_size=self.hop_size,
max_duration=self.max_duration,
)
if verbose:
dt = time.time() - t0
voiced_ratio = float(np.mean(f0 > 0)) if len(f0) else 0.0
print(
"[f0 extraction] process: done:",
f"frames={len(f0)}",
f"voiced_ratio={voiced_ratio:.3f}",
f"time={dt:.3f}s",
)
if f0_path is not None:
np.save(f0_path, f0)
return f0
if __name__ == "__main__":
model_path = (
"pretrained_models/rmvpe/rmvpe.pt"
)
audio_path = "./outputs/transcription/test.wav"
pe = F0Extractor(
model_path,
device="cuda",
)
f0 = pe.process(audio_path)
+72
View File
@@ -0,0 +1,72 @@
import re
import ToJyutping
from g2pM import G2pM
from g2p_en import G2p as G2pE
_EN_WORD_RE = re.compile(r"^[A-Za-z]+(?:'[A-Za-z]+)*$")
_ZH_WORD_RE = re.compile(r"[\u4e00-\u9fff]")
EN_FLAG = "en_"
YUE_FLAG = "yue_"
ZH_FLAG = "zh_"
g2p_zh = G2pM()
g2p_en = G2pE()
def is_chinese_char(word: str) -> bool:
if len(word) != 1:
return False
return bool(_ZH_WORD_RE.fullmatch(word))
def is_english_word(word: str) -> bool:
if not word:
return False
return bool(_EN_WORD_RE.fullmatch(word))
def g2p_cantonese(sent):
return ToJyutping.get_jyutping_list(sent) # with tone
def g2p_mandarin(sent):
return g2p_zh(sent, tone=True, char_split=False)
def g2p_english(word):
return g2p_en(word)
def g2p_transform(words, lang):
zh_words = []
transformed_words = [0] * len(words)
for idx, w in enumerate(words):
if w == "<SP>":
transformed_words[idx] = w
continue
w = w.replace("?", "").replace(".", "").replace("!", "").replace(",", "")
if is_chinese_char(w):
zh_words.append([idx, w])
else:
if is_english_word(w):
w = EN_FLAG + "-".join(g2p_english(w.lower()))
else:
w = "<SP>"
transformed_words[idx] = w
sent = "".join([k[1] for k in zh_words])
# zh (zh and yue) transformer to g2p
if len(sent) > 0:
if lang == "Cantonese":
g2pm_rst = g2p_cantonese(sent) # with tone
g2pm_rst = [YUE_FLAG + k[1] for k in g2pm_rst]
else:
g2pm_rst = g2p_mandarin(sent)
g2pm_rst = [ZH_FLAG + k for k in g2pm_rst]
for p, w in zip([k[0] for k in zh_words], g2pm_rst):
transformed_words[p] = w
return transformed_words
+279
View File
@@ -0,0 +1,279 @@
# https://modelscope.cn/models/iic/speech_seaco_paraformer_large_asr_nat-zh-cn-16k-common-vocab8404-pytorch/summary
# https://huggingface.co/nvidia/parakeet-tdt-0.6b-v2
import os
import re
import time
from typing import Any, Dict, List, Tuple
import librosa
import numpy as np
from funasr import AutoModel
def _build_words_with_gaps(raw_words, raw_timestamps, wav_fn: str):
words, word_durs = [], []
prev = 0.0
for w, t in zip(raw_words, raw_timestamps):
s, e = float(t[0]), float(t[1])
if s > prev:
words.append("<SP>")
word_durs.append(s - prev)
words.append(w)
word_durs.append(e - s)
prev = e
wav_len = librosa.get_duration(filename=wav_fn)
if wav_len > prev:
if len(words) == 0:
words.append("<SP>")
word_durs.append(wav_len)
return words, word_durs
if words[-1] != "<SP>":
words.append("<SP>")
word_durs.append(wav_len - prev)
else:
word_durs[-1] += wav_len - prev
return words, word_durs
def _word_dur_post_process(words, word_durs, f0):
"""Post-process word durations using f0 to better place silences.
"""
# f0 time grid parameters
sr = 24000 # f0 sample rate
hop_length = 480 # f0 hop length
# Convert word durations (seconds) to frame boundaries on the f0 grid.
boundaries = np.cumsum([
0,
*[
int(dur * sr / hop_length)
for dur in word_durs
],
]).tolist()
sil_tolerance = 5 # tolerance frames for silence detection
ext_tolerance = 5 # tolerance frames for vocal extension
new_words: list[str] = []
new_word_durs: list[float] = []
if words:
new_words.append(words[0])
new_word_durs.append(word_durs[0])
for i in range(1, len(words)):
word = words[i]
if word == "<SP>":
start_frame = boundaries[i]
end_frame = boundaries[i + 1]
num_frames = end_frame - start_frame
frame_idx = start_frame
# Find first region with at least 5 consecutive "unvoiced" frames.
unvoiced_count = 0
while frame_idx < end_frame:
if f0[frame_idx] <= 1: # unvoiced
unvoiced_count += 1
if unvoiced_count >= sil_tolerance:
frame_idx -= sil_tolerance - 1 # back to the last voiced frame
break
else:
unvoiced_count = 0
frame_idx += 1
voice_frames = frame_idx - start_frame
if voice_frames >= int(num_frames * 0.9): # over 90% voiced
# Treat the whole "<SP>" as silence and merge into previous word.
new_word_durs[-1] += word_durs[i]
elif voice_frames >= ext_tolerance: # over 5 frames voiced
# Split the "<SP>" into two parts: leading silence and tail kept as "<SP>".
dur = voice_frames * hop_length / sr
new_word_durs[-1] += dur
new_words.append("<SP>")
new_word_durs.append(word_durs[i] - dur)
else:
# Too short to adjust, keep as-is.
new_words.append(word)
new_word_durs.append(word_durs[i])
else:
new_words.append(word)
new_word_durs.append(word_durs[i])
return new_words, new_word_durs
class _ASRZhModel:
"""Mandarin/Cantonese ASR wrapper."""
def __init__(self, model_path: str, device: str):
self.model = AutoModel(
model=model_path,
disable_update=True,
device=device,
)
def process(self, wav_fn):
out = self.model.generate(wav_fn, output_timestamp=True)[0]
raw_words = out["text"].replace("@", "").split(" ")
raw_timestamps = [[t[0] / 1000, t[1] / 1000] for t in out["timestamp"]]
words, word_durs = _build_words_with_gaps(raw_words, raw_timestamps, wav_fn)
if os.path.exists(wav_fn.replace(".wav", "_f0.npy")):
words, word_durs = _word_dur_post_process(
words, word_durs, np.load(wav_fn.replace(".wav", "_f0.npy"))
)
return words, word_durs
class _ASREnModel:
"""English ASR wrapper for NeMo Parakeet-TDT."""
def __init__(self, model_path: str, device: str):
try:
import nemo.collections.asr as nemo_asr # type: ignore
except Exception as e: # pragma: no cover
raise ImportError(
"NeMo (nemo_toolkit) is required for ASR English but is not available in this Python env. "
"Install it in the active environment, then retry."
) from e
self.model = nemo_asr.models.ASRModel.restore_from(
restore_path=model_path,
map_location=device,
)
self.model.eval()
@staticmethod
def _clean_word(word: str) -> str:
return re.sub(r"[\?\.,:]", "", word).strip()
@staticmethod
def _extract_word_segments(output: Any) -> List[Dict[str, Any]]:
ts = getattr(output, "timestamp", None)
if not ts or not isinstance(ts, dict):
return []
word_ts = ts.get("word")
return word_ts if isinstance(word_ts, list) else []
def process(self, wav_fn: str) -> Tuple[List[str], List[float]]:
outputs = self.model.transcribe(
[wav_fn],
timestamps=True,
batch_size=1,
num_workers=0,
)
output = outputs[0] if outputs else None
raw_words: List[str] = []
raw_timestamps: List[List[float]] = []
if output is not None:
for w in self._extract_word_segments(output):
s, e = float(w.get("start", 0.0)), float(w.get("end", 0.0))
word = self._clean_word(str(w.get("word", "")))
if word:
raw_words.append(word)
raw_timestamps.append([s, e])
words, durs = _build_words_with_gaps(raw_words, raw_timestamps, wav_fn)
if os.path.exists(wav_fn.replace(".wav", "_f0.npy")):
words, durs = _word_dur_post_process(
words, durs, np.load(wav_fn.replace(".wav", "_f0.npy"))
)
return words, durs
class LyricTranscriber:
"""Transcribe lyrics from singing voice segment
"""
def __init__(
self,
zh_model_path: str,
en_model_path: str,
device: str = "cuda",
*,
verbose: bool = True,
):
"""Initialize lyric transcriber.
Args:
zh_model_path (str): Path to the Chinese model file.
en_model_path (str): Path to the English model file.
device (str): Device to use for tensor operations.
verbose (bool): Whether to print verbose logs.
"""
self.verbose = verbose
self.device = device
self.zh_model_path = zh_model_path
self.en_model_path = en_model_path
if self.verbose:
print(
"[lyric transcription] init: start:",
f"device={device}",
f"model_path={zh_model_path}",
)
# Always initialize Chinese ASR.
self.zh_model = _ASRZhModel(device=device, model_path=zh_model_path)
# English ASR will be lazily initialized on first English request to avoid long waiting cost when importing NeMo
self.en_model = None
if self.verbose:
print("[lyric transcription] init: success")
def process(self, wav_fn, language: str | None = "Mandarin", *, verbose: bool | None = None):
""" Lyric transcriber process
Args:
wav_fn (str): Path to the audio file.
language (str | None): Language of the audio. Defaults to "Mandarin". Supports "Mandarin", "Cantonese" and "English".
verbose (bool | None): Whether to print verbose logs. Defaults to None.
"""
v = self.verbose if verbose is None else verbose
if language not in {"Mandarin", "Cantonese", "English"}:
raise ValueError(f"Unsupported language: {language}, should be one of ['Mandarin', 'Cantonese', 'English']")
if v:
print(f"[lyric transcription] process: start: wav_fn={wav_fn} language={language}")
t0 = time.time()
lang = (language or "auto").lower()
if lang in {"english"}:
if self.en_model is None:
# Lazy-load NeMo model only when English is actually used.
if v:
print("[lyric transcription] init English ASR, please make sure NeMo is installed")
self.en_model = _ASREnModel(model_path=self.en_model_path, device=self.device)
out = self.en_model.process(wav_fn)
else:
out = self.zh_model.process(wav_fn)
if v:
words, durs = out
n_words = len(words) if isinstance(words, list) else 0
dur_sum = float(sum(durs)) if isinstance(durs, list) else 0.0
dt = time.time() - t0
print(
"[lyric transcription] process: done:",
f"n_words={n_words}",
f"dur_sum={dur_sum:.3f}s",
f"time={dt:.3f}s",
)
return out
if __name__ == "__main__":
m = LyricTranscriber(
zh_model_path="pretrained_models/speech_seaco_paraformer_large_asr_nat-zh-cn-16k-common-vocab8404-pytorch",
en_model_path="pretrained_models/parakeet-tdt-0.6b-v2/parakeet-tdt-0.6b-v2.nemo",
device="cuda"
)
print(m.process("example/test/asr_zh.wav", language="Mandarin"))
print(m.process("example/test/asr_en.wav", language="English"))
+173
View File
@@ -0,0 +1,173 @@
# 🎹 MIDI Editor - Web-based Singing MIDI Editor
[English](README.md) | [简体中文](README_CN.md)
A full-featured web MIDI editor for singing voice production, similar to ACE-Studio and VOCALOID. It supports real-time drag editing of MIDI notes, lyric editing, audio waveform alignment, and importing/exporting MIDI files with lyrics.
![MIDI Editor](https://img.shields.io/badge/React-19.2-blue) ![TypeScript](https://img.shields.io/badge/TypeScript-5.9-blue) ![Vite](https://img.shields.io/badge/Vite-7.2-purple)
## ✨ Features
### 🎼 Piano Roll Editing
- **Visual note editing**: Full range from C1 to C8 with intuitive piano keys
- **Drag operations**:
- Move notes: drag note blocks to adjust position and pitch
- Resize start: drag the left edge to adjust start time
- Resize end: drag the right edge to adjust end time
- **Double-click to add**: Add new notes quickly in empty areas
- **Piano key preview**: Click a key to audition the pitch
### 🔍 Zoom & Navigation
- **Horizontal zoom**
- **Vertical zoom**
- **Dynamic snapping**: finer snap granularity at higher zoom (min 0.01s)
- **Auto scroll**: keep the playhead visible during playback
### 📝 Lyric Editing
- **Inline editing**: edit lyrics for each note in the side list
- **Batch fill**: enter a string and auto-fill notes in order
- **Fill from selection**: start batch fill from the selected note
- **Precise fields**: edit PITCH, START, and END directly
- **Confirm edits**: press Enter or click ✓ to confirm
### 🎵 Audio Alignment
- **Waveform display**: sync waveform with the MIDI timeline
- **Formats**: MP3, WAV, OGG, FLAC, M4A, AAC
- **Sync playback**: play audio and MIDI together with independent mute
- **Click to seek**: click waveform or timeline to seek
### ⚠️ Overlap Detection
- **Visual highlight**: overlapping notes blink in red
- **Smart tolerance**: adjacent notes (end equals next start) are not overlaps
- **One-click fix**: remove all overlaps automatically
- **Export warning**: warn if overlaps exist during export
### 📥 Import & Export
- **MIDI import**: parse standard MIDI and lyric metadata
- **MIDI export**: export MIDI with lyrics
- **Chinese support**: full UTF-8 lyrics support
### 🎨 UI & UX
- **Theme toggle**: light and dark modes
- **Responsive layout**: adapts to window size
- **SVG grid**: cross-browser grid rendering
- **Status feedback**: real-time state and error tips
## 🚀 Quick Start
### Requirements
- Node.js 18+
- npm or yarn
### Install
```bash
# Clone
git clone <repository-url>
cd MIDI_Editor
# Install dependencies
npm install
# Start dev server
npm run dev
# Expose to LAN
npm run dev -- --host 0.0.0.0
```
### Build
```bash
# Build for production
npm run build
# Preview build
npm run preview
```
## 📖 Usage
### Basic Workflow
1. **Import MIDI**: click Import MIDI and select a .mid file
2. **Edit notes**: drag notes in the piano roll to adjust time and pitch
3. **Add lyrics**: edit lyrics in the right-side list
4. **Align audio** (optional): import reference audio
5. **Export**: click Export MIDI with lyrics
### Shortcuts
| Action | Description |
|------|------|
| Double-click piano roll | Add a new note |
| Double-click note | Edit lyric |
| Drag note | Move note and pitch |
| Drag note edges | Resize note |
| Backspace / Delete | Delete selected note |
| Enter | Confirm value edits |
| Escape | Cancel value edits |
| Ctrl(Command) + Wheel | Horizontal zoom |
| Ctrl(Command) + Shift(Option) + Wheel | Vertical zoom |
### Playback Controls
| Button | Description |
|------|------|
| ⏮ | Go to start |
| ⏪ 2s | Back 2 seconds |
| ▶ / ⏸ | Play / Pause |
| 2s ⏩ | Forward 2 seconds |
| ⏭ | Go to end |
## 🛠 Tech Stack
- **Frontend**: React 19 + TypeScript
- **Build**: Vite 7
- **State**: Zustand
- **Audio**: Tone.js
- **Waveform**: WaveSurfer.js
- **MIDI**: @tonejs/midi
- **Styles**: CSS with custom variables
## 📁 Project Structure
```
.
├── eslint.config.js
├── index.html
├── package.json
├── postcss.config.js
├── README.md
├── README_CN.md
├── tailwind.config.js
├── tsconfig.app.json
├── tsconfig.json
├── tsconfig.node.json
├── vite.config.ts
├── public/
└── src/
├── App.css
├── App.tsx
├── constants.ts
├── index.css
├── main.tsx
├── types.ts
├── assets/
├── components/
│ ├── AudioTrack.tsx
│ ├── LyricTable.tsx
│ └── PianoRoll.tsx
├── lib/
│ └── midi.ts
└── store/
└── useMidiStore.ts
```
+173
View File
@@ -0,0 +1,173 @@
# 🎹 MIDI Editor - 网页端歌声 MIDI 编辑器
[English](README.md) | [简体中文](README_CN.md)
一个功能完整的网页端歌声 MIDI 文件编辑器,类似 ACE-Studio 和 VOCALOID。支持实时拖拽调整 MIDI 音符、歌词编辑、音频波形对齐,以及导入导出含歌词的 MIDI 文件。
![MIDI Editor](https://img.shields.io/badge/React-19.2-blue) ![TypeScript](https://img.shields.io/badge/TypeScript-5.9-blue) ![Vite](https://img.shields.io/badge/Vite-7.2-purple)
## ✨ 功能特性
### 🎼 钢琴卷帘编辑
- **可视化音符编辑**:支持 C1-C8 全音域显示,直观的钢琴键布局
- **拖拽操作**:
- 移动音符:拖拽音符块调整位置和音高
- 调整音头:拖拽音符左边缘调整开始时间
- 调整音尾:拖拽音符右边缘调整结束时间
- **双击添加**:在钢琴卷帘空白处双击快速添加新音符
- **钢琴键试听**:点击左侧钢琴键可试听对应音高
### 🔍 缩放与导航
- **水平缩放**
- **垂直缩放**
- **动态精度**:缩放越大,音符调整的 snap 粒度越精细(最小 0.01 秒)
- **自动滚动**:播放时播放头自动保持可见
### 📝 歌词编辑
- **实时编辑**:右侧列表直接编辑每个音符的歌词
- **批量填充**:输入一段歌词,按字顺序自动填充到音符
- **从选中开始**:批量填充可从当前选中的音符开始
- **精确调整**:可直接编辑 PITCH(音高)、START(开始时间)、END(结束时间)
- **确认机制**:修改数值后按 Enter 或点击 ✓ 确认,避免误操作
### 🎵 音频对齐
- **波形显示**:导入音频后显示波形,与 MIDI 同步滚动
- **格式支持**:MP3、WAV、OGG、FLAC、M4A、AAC
- **同步播放**:音频与 MIDI 同步播放,可分别静音
- **点击定位**:点击波形或时间尺可快速定位播放位置
### ⚠️ 重叠检测
- **可视化标注**:时间重叠的音符显示为红色并闪烁
- **智能容差**:紧邻的音符(上一个结束 = 下一个开始)不视为重叠
- **一键修复**:点击消除重叠按钮自动修复所有重叠
- **导出提醒**:导出时如有重叠会弹出警告
### 📥 导入导出
- **MIDI 导入**:支持标准 MIDI 文件,自动解析歌词元数据
- **MIDI 导出**:导出包含歌词信息的 MIDI 文件
- **中文支持**:完整支持中文歌词的导入导出(UTF-8 编码)
### 🎨 界面特性
- **主题切换**:支持浅色/深色主题
- **响应式布局**:自适应窗口大小
- **SVG 网格**:跨浏览器兼容的网格渲染
- **状态提示**:实时显示操作状态和错误信息
## 🚀 快速开始
### 环境要求
- Node.js 18+
- npm 或 yarn
### 安装
```bash
# 克隆项目
git clone <repository-url>
cd MIDI_Editor
# 安装依赖
npm install
# 启动开发服务器
npm run dev
# 在局域网启动
npm run dev -- --host 0.0.0.0
```
### 构建
```bash
# 构建生产版本
npm run build
# 预览构建结果
npm run preview
```
## 📖 使用指南
### 基本工作流
1. **导入 MIDI**:点击导入 MIDI 按钮选择 .mid 文件
2. **编辑音符**:在钢琴卷帘中拖拽调整音符位置和时长
3. **添加歌词**:在右侧列表中输入每个音符的歌词
4. **对齐音频**(可选):导入参考音频进行对照编辑
5. **导出文件**:点击导出含歌词 MIDI 保存文件
### 快捷操作
| 操作 | 说明 |
|------|------|
| 双击钢琴卷帘 | 添加新音符 |
| 双击音符 | 修改歌词 |
| 拖拽音符 | 移动音符位置/音高 |
| 拖拽音符边缘 | 调整音符时长 |
| Backspace / Delete | 删除选中音符 |
| Enter | 确认数值修改 |
| Escape | 取消数值修改 |
| Ctrl(Command) + 滚轮 | 水平缩放 |
| Ctrl(Command) + Shift(Option) + 滚轮 | 垂直缩放 |
### 播放控制
| 按钮 | 功能 |
|------|------|
| ⏮ | 回到开头 |
| ⏪ 2s | 后退 2 秒 |
| ▶ / ⏸ | 播放 / 暂停 |
| 2s ⏩ | 前进 2 秒 |
| ⏭ | 跳到结尾 |
## 🛠 技术栈
- **前端框架**:React 19 + TypeScript
- **构建工具**:Vite 7
- **状态管理**:Zustand
- **音频引擎**:Tone.js
- **波形显示**:WaveSurfer.js
- **MIDI 解析**:@tonejs/midi
- **样式**:CSS(自定义变量主题)
## 📁 项目结构
```
.
├── eslint.config.js
├── index.html
├── package.json
├── postcss.config.js
├── README.md
├── README_CN.md
├── tailwind.config.js
├── tsconfig.app.json
├── tsconfig.json
├── tsconfig.node.json
├── vite.config.ts
├── public/
└── src/
├── App.css
├── App.tsx
├── constants.ts
├── index.css
├── main.tsx
├── types.ts
├── assets/
├── components/
│ ├── AudioTrack.tsx
│ ├── LyricTable.tsx
│ └── PianoRoll.tsx
├── lib/
│ └── midi.ts
└── store/
└── useMidiStore.ts
```
@@ -0,0 +1,23 @@
import js from '@eslint/js'
import globals from 'globals'
import reactHooks from 'eslint-plugin-react-hooks'
import reactRefresh from 'eslint-plugin-react-refresh'
import tseslint from 'typescript-eslint'
import { defineConfig, globalIgnores } from 'eslint/config'
export default defineConfig([
globalIgnores(['dist']),
{
files: ['**/*.{ts,tsx}'],
extends: [
js.configs.recommended,
tseslint.configs.recommended,
reactHooks.configs.flat.recommended,
reactRefresh.configs.vite,
],
languageOptions: {
ecmaVersion: 2020,
globals: globals.browser,
},
},
])
+13
View File
@@ -0,0 +1,13 @@
<!doctype html>
<html lang="en">
<head>
<meta charset="UTF-8" />
<link rel="icon" type="image/svg+xml" href="/vite.svg" />
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
<title>midi-editor</title>
</head>
<body>
<div id="root"></div>
<script type="module" src="/src/main.tsx"></script>
</body>
</html>
File diff suppressed because it is too large Load Diff
+39
View File
@@ -0,0 +1,39 @@
{
"name": "midi-editor",
"private": true,
"version": "0.0.0",
"type": "module",
"scripts": {
"dev": "vite",
"build": "tsc -b && vite build",
"lint": "eslint .",
"preview": "vite preview"
},
"dependencies": {
"@tonejs/midi": "^2.0.28",
"class-variance-authority": "^0.7.1",
"nanoid": "^5.1.6",
"react": "^19.2.0",
"react-dom": "^19.2.0",
"tone": "^15.1.22",
"wavesurfer.js": "^7.12.1",
"zustand": "^5.0.10"
},
"devDependencies": {
"@eslint/js": "^9.39.1",
"@types/node": "^24.10.1",
"@types/react": "^19.2.5",
"@types/react-dom": "^19.2.3",
"@vitejs/plugin-react": "^5.1.1",
"autoprefixer": "^10.4.20",
"eslint": "^9.39.1",
"eslint-plugin-react-hooks": "^7.0.1",
"eslint-plugin-react-refresh": "^0.4.24",
"globals": "^16.5.0",
"postcss": "^8.4.47",
"tailwindcss": "^3.4.15",
"typescript": "~5.9.3",
"typescript-eslint": "^8.46.4",
"vite": "^7.2.4"
}
}
@@ -0,0 +1,6 @@
export default {
plugins: {
tailwindcss: {},
autoprefixer: {},
},
}
@@ -0,0 +1 @@
<svg xmlns="http://www.w3.org/2000/svg" xmlns:xlink="http://www.w3.org/1999/xlink" aria-hidden="true" role="img" class="iconify iconify--logos" width="31.88" height="32" preserveAspectRatio="xMidYMid meet" viewBox="0 0 256 257"><defs><linearGradient id="IconifyId1813088fe1fbc01fb466" x1="-.828%" x2="57.636%" y1="7.652%" y2="78.411%"><stop offset="0%" stop-color="#41D1FF"></stop><stop offset="100%" stop-color="#BD34FE"></stop></linearGradient><linearGradient id="IconifyId1813088fe1fbc01fb467" x1="43.376%" x2="50.316%" y1="2.242%" y2="89.03%"><stop offset="0%" stop-color="#FFEA83"></stop><stop offset="8.333%" stop-color="#FFDD35"></stop><stop offset="100%" stop-color="#FFA800"></stop></linearGradient></defs><path fill="url(#IconifyId1813088fe1fbc01fb466)" d="M255.153 37.938L134.897 252.976c-2.483 4.44-8.862 4.466-11.382.048L.875 37.958c-2.746-4.814 1.371-10.646 6.827-9.67l120.385 21.517a6.537 6.537 0 0 0 2.322-.004l117.867-21.483c5.438-.991 9.574 4.796 6.877 9.62Z"></path><path fill="url(#IconifyId1813088fe1fbc01fb467)" d="M185.432.063L96.44 17.501a3.268 3.268 0 0 0-2.634 3.014l-5.474 92.456a3.268 3.268 0 0 0 3.997 3.378l24.777-5.718c2.318-.535 4.413 1.507 3.936 3.838l-7.361 36.047c-.495 2.426 1.782 4.5 4.151 3.78l15.304-4.649c2.372-.72 4.652 1.36 4.15 3.788l-11.698 56.621c-.732 3.542 3.979 5.473 5.943 2.437l1.313-2.028l72.516-144.72c1.215-2.423-.88-5.186-3.54-4.672l-25.505 4.922c-2.396.462-4.435-1.77-3.759-4.114l16.646-57.705c.677-2.35-1.37-4.583-3.769-4.113Z"></path></svg>

After

Width:  |  Height:  |  Size: 1.5 KiB

+785
View File
@@ -0,0 +1,785 @@
.app-shell {
padding: 24px;
color: var(--text-primary);
width: 100%;
max-width: 100%;
margin: 0;
height: 100vh;
max-height: 100vh;
display: flex;
flex-direction: column;
overflow: hidden;
box-sizing: border-box;
}
.topbar {
display: flex;
align-items: center;
justify-content: space-between;
gap: 24px;
background: var(--panel-strong);
border: 1px solid var(--border-subtle);
border-radius: 16px;
padding: 20px 24px;
box-shadow: var(--shadow-panel);
}
.topbar h1 {
margin: 4px 0 0 0;
font-size: 26px;
letter-spacing: -0.5px;
}
.eyebrow {
margin: 0;
text-transform: uppercase;
font-size: 12px;
letter-spacing: 2px;
color: var(--text-muted);
}
.muted {
margin: 6px 0 0 0;
color: var(--text-muted);
}
.actions {
display: flex;
gap: 10px;
}
.icon-toggle {
width: 40px;
height: 40px;
border-radius: 999px;
border: 1px solid var(--border-soft);
background: var(--button-ghost-bg);
color: var(--button-ghost-text);
display: inline-flex;
align-items: center;
justify-content: center;
font-size: 18px;
cursor: pointer;
}
.icon-toggle:hover {
transform: translateY(-1px);
}
.audio-bar {
margin-top: 14px;
padding: 12px 16px;
border-radius: 14px;
background: var(--panel-strong);
border: 1px solid var(--border-subtle);
display: flex;
align-items: center;
justify-content: space-between;
gap: 16px;
}
.audio-left {
display: flex;
align-items: center;
gap: 12px;
}
.audio-hint {
color: var(--text-muted);
font-size: 12px;
}
.audio-right {
display: flex;
align-items: center;
gap: 20px;
}
.volume-control {
display: flex;
align-items: center;
gap: 8px;
}
.volume-label {
font-size: 12px;
color: var(--text-muted);
min-width: 32px;
}
.volume-slider {
width: 80px;
height: 4px;
cursor: pointer;
accent-color: var(--accent);
}
.volume-value {
font-size: 11px;
color: var(--text-muted);
min-width: 36px;
text-align: right;
}
.toggle {
display: inline-flex;
align-items: center;
gap: 8px;
font-size: 13px;
color: var(--text-primary);
}
.panel {
margin-top: 18px;
background: var(--panel);
border: 1px solid var(--border-subtle);
border-radius: 16px;
padding: 18px;
box-shadow: var(--shadow-panel);
display: flex;
flex-direction: column;
flex: 1;
min-height: 0;
overflow: hidden;
}
.panel-split {
display: grid;
grid-template-columns: minmax(0, 1fr) 360px;
gap: 16px;
align-items: stretch;
flex: 1;
min-height: 0;
max-height: 100%;
overflow: hidden;
}
.panel-main {
min-width: 0;
display: flex;
flex-direction: column;
min-height: 0;
max-height: 100%;
overflow: hidden;
}
.panel-side {
display: flex;
flex-direction: column;
gap: 16px;
width: 360px;
max-width: 360px;
/* Use absolute positioning to enforce height */
position: relative;
overflow: hidden;
}
.controls {
display: grid;
grid-template-columns: repeat(auto-fit, minmax(180px, 1fr));
gap: 14px;
align-items: center;
background: var(--panel-soft);
padding: 12px 14px;
border-radius: 12px;
border: 1px solid var(--border-soft);
flex-shrink: 0;
}
.controls label {
display: block;
font-size: 12px;
text-transform: uppercase;
letter-spacing: 1px;
color: var(--text-muted);
margin-bottom: 4px;
}
.controls input[type='number'] {
width: 100%;
padding: 10px 12px;
border-radius: 10px;
border: 1px solid var(--border-soft);
background: var(--input-bg);
color: var(--text-primary);
}
.timesig {
display: flex;
align-items: center;
gap: 6px;
}
.timesig span {
font-weight: 700;
color: var(--text-muted);
}
.transport {
display: flex;
gap: 8px;
align-items: center;
}
.selection-controls {
display: flex;
gap: 8px;
align-items: center;
padding: 6px 0;
}
.selection-btn {
font-size: 12px !important;
padding: 6px 10px !important;
}
.selection-btn.active {
background: var(--accent) !important;
color: white !important;
}
.selection-info {
font-size: 12px;
color: var(--accent);
font-weight: 500;
padding: 4px 8px;
background: rgba(var(--accent-rgb), 0.1);
border-radius: 6px;
}
.status {
color: var(--text-muted);
font-size: 13px;
}
.button,
.actions button,
.transport button,
.ghost,
.primary,
.soft {
cursor: pointer;
border-radius: 12px;
border: 1px solid transparent;
padding: 10px 14px;
font-weight: 600;
transition: transform 140ms ease, box-shadow 140ms ease, background 140ms ease, border 140ms ease;
color: #0f1528;
}
.ghost {
background: var(--button-ghost-bg);
color: var(--button-ghost-text);
border-color: var(--border-soft);
}
.primary {
background: linear-gradient(135deg, var(--accent), var(--accent-strong));
color: var(--button-primary-text);
box-shadow: 0 8px 26px rgba(72, 228, 194, 0.2);
}
.soft {
background: var(--button-soft-bg);
color: var(--button-soft-text);
border: 1px solid var(--border-soft);
}
.ghost:disabled,
.primary:disabled,
.soft:disabled {
opacity: 0.6;
cursor: not-allowed;
}
.ghost:hover,
.primary:hover,
.soft:hover {
transform: translateY(-1px);
}
.piano-shell {
border-radius: 12px;
background: var(--panel-strong);
border: 1px solid var(--border-subtle);
overflow: hidden;
flex: 1;
min-height: 0;
max-height: 100%;
display: flex;
flex-direction: column;
}
.ruler {
position: relative;
height: 32px;
background: var(--panel-soft);
border-bottom: 1px solid var(--border-soft);
min-width: 100%;
}
.ruler-shell {
display: flex;
}
.ruler-spacer {
background: var(--panel-soft);
border-bottom: 1px solid var(--border-soft);
height: 32px;
}
.ruler-scroll {
overflow: hidden;
flex: 1;
height: 32px;
cursor: pointer;
}
.measure-mark {
position: absolute;
top: 0;
height: 100%;
display: flex;
flex-direction: column;
align-items: flex-start;
font-size: 10px;
color: var(--text-muted);
padding-left: 4px;
border-left: 1px solid var(--border-soft);
}
.measure-mark span {
margin-top: 2px;
}
.ruler-playhead {
position: absolute;
top: 0;
width: 2px;
height: 100%;
background: #ff7043;
pointer-events: none;
z-index: 10;
}
.ruler-scroll.selecting {
cursor: crosshair;
}
.selection-range {
position: absolute;
top: 0;
height: 100%;
background: rgba(66, 165, 245, 0.35);
border-left: 2px solid #42a5f5;
border-right: 2px solid #42a5f5;
pointer-events: none;
z-index: 5;
}
.grid-selection-range {
position: absolute;
top: 0;
background: rgba(66, 165, 245, 0.15);
border-left: 2px dashed #42a5f5;
border-right: 2px dashed #42a5f5;
pointer-events: none;
z-index: 1;
}
.roll-body {
display: flex;
flex: 1;
min-height: 0;
overflow: hidden;
}
.pitch-rail {
background: var(--panel-strong);
border-right: 1px solid var(--border-subtle);
color: var(--text-primary);
font-size: 12px;
text-align: right;
overflow: hidden;
flex-shrink: 0;
height: 100%;
}
.pitch-cell {
border-bottom: 1px solid var(--border-soft);
display: flex;
align-items: center;
justify-content: flex-end;
padding: 0 4px;
font-variant-numeric: tabular-nums;
box-sizing: border-box;
}
.pitch-white {
background: rgba(255, 255, 255, 0.06);
color: var(--text-primary);
}
.pitch-black {
background: rgba(0, 0, 0, 0.35);
color: rgba(233, 238, 247, 0.9);
}
.pitch-c {
background: rgba(100, 150, 255, 0.15);
font-weight: 600;
}
.pitch-label {
font-size: 10px;
}
.roll-grid {
position: relative;
overflow: auto;
flex: 1;
min-height: 0;
background-color: var(--grid-bg);
}
.grid-content {
background-color: var(--grid-bg);
}
.grid-svg {
shape-rendering: crispEdges;
}
.grid-overlay {
position: relative;
}
.note-chip {
position: absolute;
background: linear-gradient(135deg, var(--accent), var(--accent-strong));
border-radius: 6px;
border: 1px solid rgba(255, 255, 255, 0.16);
box-shadow: 0 10px 22px rgba(0, 0, 0, 0.25);
display: flex;
align-items: center;
justify-content: center;
color: var(--note-text);
font-weight: 700;
user-select: none;
box-sizing: border-box;
}
.note-active {
outline: 2px solid #ff7043;
z-index: 2;
}
.note-overlap {
background: linear-gradient(135deg, #ef5350 0%, #ff7043 100%) !important;
animation: pulse-overlap 1s ease-in-out infinite;
}
/* Selected overlapping note - more visible outline */
.note-overlap.note-active {
outline: 3px solid #1e40af;
outline-offset: 1px;
box-shadow: 0 0 12px rgba(30, 64, 175, 0.8);
animation: none;
}
@keyframes pulse-overlap {
0%, 100% { opacity: 1; }
50% { opacity: 0.7; }
}
.playhead {
position: absolute;
top: 0;
width: 2px;
background: #ff7043;
box-shadow: 0 0 12px rgba(255, 112, 67, 0.6);
pointer-events: none;
z-index: 20;
}
.pitch-rail-inner {
will-change: transform;
}
.note-label {
width: 100%;
text-align: center;
font-size: 12px;
padding: 0 12px;
overflow: hidden;
text-overflow: ellipsis;
white-space: nowrap;
}
.note-handle {
position: absolute;
top: 0;
width: 8px;
height: 100%;
background: rgba(255, 255, 255, 0.25);
cursor: ew-resize;
}
.note-handle.start {
left: 0;
border-radius: 6px 0 0 6px;
}
.note-handle.end {
right: 0;
border-radius: 0 6px 6px 0;
}
.lyric-container {
flex: 1;
min-height: 0;
position: relative;
}
.lyric-card {
border: 1px solid rgba(255, 255, 255, 0.06);
border-radius: 12px;
background: var(--panel-soft);
overflow: hidden;
display: flex;
flex-direction: column;
/* Force fixed height with absolute positioning */
position: absolute;
top: 0;
left: 0;
right: 0;
bottom: 0;
}
.lyric-bulk {
display: flex;
gap: 8px;
padding: 10px 12px;
border-bottom: 1px solid rgba(255, 255, 255, 0.06);
align-items: center;
}
.lyric-bulk-input {
flex: 1;
padding: 8px 10px;
border-radius: 10px;
border: 1px solid var(--border-soft);
background: var(--input-bg);
color: var(--text-primary);
resize: vertical;
}
.lyric-header,
.lyric-row {
display: grid;
grid-template-columns: 1.4fr 0.5fr 0.5fr 0.5fr;
gap: 8px;
padding: 10px 12px;
align-items: center;
}
.lyric-header {
font-size: 12px;
text-transform: uppercase;
letter-spacing: 1px;
color: var(--text-muted);
border-bottom: 1px solid rgba(255, 255, 255, 0.06);
}
.lyric-list {
overflow-y: auto;
overflow-x: hidden;
flex: 1;
min-height: 0;
}
.lyric-row {
border-bottom: 1px solid rgba(255, 255, 255, 0.04);
}
.lyric-row:hover {
background: rgba(255, 255, 255, 0.03);
}
.lyric-row-active {
background: rgba(72, 228, 194, 0.08);
border-left: 3px solid #48e4c2;
}
.lyric-input {
width: 100%;
padding: 8px 10px;
border-radius: 10px;
border: 1px solid var(--border-soft);
background: var(--input-bg);
color: var(--text-primary);
}
.lyric-meta {
color: var(--text-muted);
font-variant-numeric: tabular-nums;
}
.editable-cell {
position: relative;
display: flex;
align-items: center;
gap: 2px;
}
.lyric-meta-input {
width: 100%;
padding: 2px 4px;
border: 1px solid transparent;
border-radius: 4px;
background: transparent;
color: var(--text-muted);
font-size: 12px;
font-variant-numeric: tabular-nums;
text-align: center;
outline: none;
transition: border-color 0.15s, background-color 0.15s;
}
.lyric-meta-input:hover {
background: var(--surface-elevated);
}
.lyric-meta-input:focus {
border-color: var(--accent);
background: var(--surface-elevated);
color: var(--text-primary);
}
.lyric-meta-dirty {
border-color: #f59e0b !important;
background: rgba(245, 158, 11, 0.1) !important;
}
.confirm-btn {
flex-shrink: 0;
width: 18px;
height: 18px;
padding: 0;
border: none;
border-radius: 4px;
background: #22c55e;
color: white;
font-size: 12px;
font-weight: bold;
cursor: pointer;
display: flex;
align-items: center;
justify-content: center;
transition: background 0.15s;
}
.confirm-btn:hover {
background: #16a34a;
}
/* Hide number input spinners */
.lyric-meta-input::-webkit-outer-spin-button,
.lyric-meta-input::-webkit-inner-spin-button {
-webkit-appearance: none;
margin: 0;
}
.lyric-meta-input[type=number] {
-moz-appearance: textfield;
}
.lyric-empty {
padding: 16px;
color: var(--text-muted);
text-align: center;
}
.audio-track {
display: grid;
grid-template-columns: 80px 1fr;
gap: 12px;
align-items: center;
padding: 12px 14px;
border-radius: 12px;
border: 1px solid var(--border-soft);
background: var(--panel-soft);
margin-bottom: 12px;
flex-shrink: 0;
}
.audio-track-label {
font-size: 12px;
text-transform: uppercase;
letter-spacing: 1px;
color: var(--text-muted);
}
.audio-wave {
width: 100%;
height: 80px;
min-height: 80px;
}
:root {
--text-primary: #e9eef7;
--text-muted: rgba(233, 238, 247, 0.7);
--panel: rgba(13, 16, 28, 0.8);
--panel-strong: rgba(16, 21, 35, 0.95);
--panel-soft: rgba(255, 255, 255, 0.03);
--border-subtle: rgba(255, 255, 255, 0.08);
--border-soft: rgba(255, 255, 255, 0.12);
--input-bg: rgba(255, 255, 255, 0.06);
--grid-bg: rgba(14, 18, 30, 0.9);
--grid-line-minor: rgba(233, 238, 247, 0.08);
--grid-line-major: rgba(233, 238, 247, 0.16);
--accent: #48e4c2;
--accent-strong: #4b64bc;
--note-text: #0b1122;
--button-ghost-bg: rgba(233, 238, 247, 0.18);
--button-ghost-text: #ffffff;
--button-soft-bg: rgba(255, 255, 255, 0.14);
--button-soft-text: #ffffff;
--button-primary-text: #0b1122;
--shadow-panel: 0 18px 40px rgba(0, 0, 0, 0.32);
}
:root[data-theme='light'] {
--text-primary: #1b2238;
--text-muted: rgba(27, 34, 56, 0.7);
--panel: rgba(255, 255, 255, 0.9);
--panel-strong: rgba(250, 252, 255, 0.98);
--panel-soft: rgba(15, 23, 42, 0.04);
--border-subtle: rgba(15, 23, 42, 0.12);
--border-soft: rgba(15, 23, 42, 0.16);
--input-bg: rgba(15, 23, 42, 0.06);
--grid-bg: rgba(248, 250, 255, 0.95);
--grid-line-minor: rgba(15, 23, 42, 0.12);
--grid-line-major: rgba(15, 23, 42, 0.24);
--accent: #3f8cff;
--accent-strong: #4b64bc;
--note-text: #ffffff;
--button-ghost-bg: rgba(15, 23, 42, 0.06);
--button-ghost-text: #1b2238;
--button-soft-bg: rgba(15, 23, 42, 0.06);
--button-soft-text: #1b2238;
--button-primary-text: #0b1122;
--shadow-panel: 0 18px 40px rgba(15, 23, 42, 0.15);
}
.sr-only {
position: absolute;
width: 1px;
height: 1px;
padding: 0;
margin: -1px;
overflow: hidden;
clip: rect(0, 0, 0, 0);
white-space: nowrap;
border: 0;
}
+654
View File
@@ -0,0 +1,654 @@
import { useCallback, useEffect, useMemo, useRef, useState } from 'react'
import * as Tone from 'tone'
import { PianoRoll } from './components/PianoRoll'
import { LyricTable } from './components/LyricTable'
import { AudioTrack } from './components/AudioTrack'
import { useMidiStore } from './store/useMidiStore'
import { exportMidi, importMidiFile } from './lib/midi'
import type { TimeSignature } from './types'
import { BASE_GRID_SECOND_WIDTH, BASE_ROW_HEIGHT, LOW_NOTE, HIGH_NOTE } from './constants'
import './App.css'
type PlayEvent = {
time: number
midi: number
duration: number
velocity: number
}
function App() {
const {
notes,
tempo,
timeSignature,
selectedId,
playhead,
ppq,
addNote,
updateNote,
removeNote,
setNotes,
setTempo,
setTimeSignature,
setPpq,
select,
setPlayhead,
} = useMidiStore()
const [status, setStatus] = useState('准备就绪')
const [isPlaying, setIsPlaying] = useState(false)
const [theme, setTheme] = useState<'dark' | 'light'>('light')
const [audioUrl, setAudioUrl] = useState<string | null>(null)
const [audioDuration, setAudioDuration] = useState(0)
const [midiVolume, setMidiVolume] = useState(80) // 0-100
const [audioVolume, setAudioVolume] = useState(80) // 0-100
const [horizontalZoom, setHorizontalZoom] = useState(1)
const [verticalZoom, setVerticalZoom] = useState(1)
const [focusLyricId, setFocusLyricId] = useState<string | null>(null)
// Selection range for loop playback (in seconds)
const [selectionStart, setSelectionStart] = useState<number | null>(null)
const [selectionEnd, setSelectionEnd] = useState<number | null>(null)
const [isSelectingRange, setIsSelectingRange] = useState(false)
const fileInputRef = useRef<HTMLInputElement | null>(null)
const audioInputRef = useRef<HTMLInputElement | null>(null)
const audioRef = useRef<HTMLAudioElement | null>(null)
const partRef = useRef<Tone.Part<PlayEvent> | null>(null)
const synthRef = useRef<Tone.PolySynth | null>(null)
const rafRef = useRef<number | null>(null)
const audioScrollRef = useRef<HTMLDivElement | null>(null)
useEffect(() => {
return () => {
stopPlayback()
synthRef.current?.dispose()
}
}, [])
useEffect(() => {
document.documentElement.dataset.theme = theme
}, [theme])
// Sync audio volume - also trigger when audioUrl changes (new audio loaded)
useEffect(() => {
if (audioRef.current) {
audioRef.current.volume = audioVolume / 100
}
}, [audioVolume, audioUrl])
// Sync MIDI synth volume
useEffect(() => {
if (synthRef.current) {
// Convert 0-100 to dB scale (-60 to 0)
const dbValue = midiVolume === 0 ? -Infinity : (midiVolume / 100) * 60 - 60
synthRef.current.volume.value = dbValue
}
}, [midiVolume])
useEffect(() => {
if (!audioUrl) return
return () => {
URL.revokeObjectURL(audioUrl)
}
}, [audioUrl])
const ensureSynth = async () => {
await Tone.start()
if (!synthRef.current) {
synthRef.current = new Tone.PolySynth(Tone.Synth).toDestination()
// Apply current volume
const dbValue = midiVolume === 0 ? -Infinity : (midiVolume / 100) * 60 - 60
synthRef.current.volume.value = dbValue
}
}
const playPreviewNote = useCallback(async (midi: number) => {
await ensureSynth()
const frequency = Tone.Frequency(midi, 'midi').toFrequency()
synthRef.current?.triggerAttackRelease(frequency, '8n', Tone.now(), 0.7)
}, [midiVolume])
useEffect(() => {
const onKeyDown = (event: KeyboardEvent) => {
if (!selectedId) return
const target = event.target as HTMLElement | null
if (target && ['INPUT', 'TEXTAREA'].includes(target.tagName)) return
// Delete note
if (event.key === 'Backspace' || event.key === 'Delete') {
event.preventDefault()
removeNote(selectedId)
select(null)
return
}
// Cmd/Ctrl + Up/Down to adjust pitch
const isCmdOrCtrl = event.metaKey || event.ctrlKey
if (isCmdOrCtrl && (event.key === 'ArrowUp' || event.key === 'ArrowDown')) {
event.preventDefault()
const selectedNote = notes.find(n => n.id === selectedId)
if (!selectedNote) return
const delta = event.key === 'ArrowUp' ? 1 : -1
const newMidi = Math.max(LOW_NOTE, Math.min(HIGH_NOTE, selectedNote.midi + delta))
if (newMidi !== selectedNote.midi) {
updateNote(selectedId, { midi: newMidi })
playPreviewNote(newMidi)
}
}
}
window.addEventListener('keydown', onKeyDown)
return () => window.removeEventListener('keydown', onKeyDown)
}, [selectedId, notes, removeNote, select, updateNote, playPreviewNote])
const noteEvents = useMemo<PlayEvent[]>(
() =>
notes.map((note) => ({
time: (60 / tempo) * note.start,
duration: (60 / tempo) * note.duration,
midi: note.midi,
velocity: note.velocity,
})),
[notes, tempo],
)
const beatToSeconds = (beat: number) => beat * (60 / tempo)
const secondsToBeat = (seconds: number) => seconds / (60 / tempo)
const seekBySeconds = (deltaSeconds: number) => {
const maxNoteEnd = notes.reduce((acc, n) => Math.max(acc, n.start + n.duration), 0)
const maxBeat = Math.max(secondsToBeat(audioDuration), maxNoteEnd)
const nextSeconds = Math.max(0, Math.min(beatToSeconds(maxBeat), beatToSeconds(playhead) + deltaSeconds))
seekToBeat(secondsToBeat(nextSeconds))
}
const gridSecondWidth = BASE_GRID_SECOND_WIDTH * horizontalZoom
const rowHeight = BASE_ROW_HEIGHT * verticalZoom
// Calculate MIDI content width to sync with audio track
const midiContentWidth = useMemo(() => {
const noteEndSeconds = notes.reduce((acc, n) => {
const endBeat = n.start + n.duration
return Math.max(acc, beatToSeconds(endBeat))
}, 8)
const maxSeconds = Math.max(noteEndSeconds + 10, audioDuration + 10, 30)
return maxSeconds * gridSecondWidth
}, [notes, audioDuration, gridSecondWidth, beatToSeconds])
const seekToBeat = (beat: number) => {
setPlayhead(beat)
Tone.Transport.seconds = beatToSeconds(beat)
if (audioRef.current) {
audioRef.current.currentTime = beatToSeconds(beat)
}
}
const schedulePlayback = async () => {
if (!notes.length && !audioUrl) return
await ensureSynth()
partRef.current?.dispose()
Tone.Transport.cancel()
Tone.Transport.stop()
Tone.Transport.bpm.value = tempo
// Determine playback range
const hasSelection = selectionStart !== null && selectionEnd !== null && selectionEnd > selectionStart
const startSeconds = hasSelection ? selectionStart : beatToSeconds(playhead)
const endSeconds = hasSelection ? selectionEnd : null
Tone.Transport.seconds = startSeconds
// Filter notes within selection range if applicable
const filteredEvents = hasSelection
? noteEvents.filter(e => e.time >= startSeconds && e.time < endSeconds!)
: noteEvents
if (filteredEvents.length) {
partRef.current = new Tone.Part((time, event) => {
if (midiVolume === 0) return
const frequency = Tone.Frequency(event.midi, 'midi').toFrequency()
synthRef.current?.triggerAttackRelease(frequency, event.duration, time, event.velocity)
}, filteredEvents)
partRef.current.start(0)
}
Tone.Transport.start()
if (audioRef.current && audioUrl) {
audioRef.current.currentTime = startSeconds
if (audioVolume > 0) {
audioRef.current.play().catch(() => null)
}
}
setIsPlaying(true)
setStatus(hasSelection ? '选区回放中...' : '正在回放...')
const tick = () => {
const seconds =
audioRef.current && audioUrl && !audioRef.current.paused
? audioRef.current.currentTime
: Tone.Transport.seconds
// Stop at selection end
if (endSeconds !== null && seconds >= endSeconds) {
pausePlayback()
seekToBeat(secondsToBeat(selectionStart!))
setStatus('选区播放完毕')
return
}
const beat = seconds / (60 / tempo)
setPlayhead(beat)
rafRef.current = requestAnimationFrame(tick)
}
rafRef.current = requestAnimationFrame(tick)
}
const stopPlayback = () => {
Tone.Transport.stop()
Tone.Transport.cancel()
partRef.current?.dispose()
partRef.current = null
setIsPlaying(false)
setPlayhead(0)
if (audioRef.current) {
audioRef.current.pause()
audioRef.current.currentTime = 0
}
if (rafRef.current) {
cancelAnimationFrame(rafRef.current)
rafRef.current = null
}
}
const pausePlayback = () => {
Tone.Transport.stop()
partRef.current?.dispose()
partRef.current = null
setIsPlaying(false)
if (audioRef.current) {
audioRef.current.pause()
}
if (rafRef.current) {
cancelAnimationFrame(rafRef.current)
rafRef.current = null
}
}
const handlePlayToggle = async () => {
if (isPlaying) {
pausePlayback()
setStatus('已暂停')
} else {
await schedulePlayback()
}
}
const handleImportClick = () => fileInputRef.current?.click()
const handleAudioImportClick = () => audioInputRef.current?.click()
const handleFileChange = async (event: React.ChangeEvent<HTMLInputElement>) => {
const file = event.target.files?.[0]
if (!file) return
try {
const snapshot = await importMidiFile(file)
setNotes(snapshot.notes)
setTempo(snapshot.tempo)
setTimeSignature(snapshot.timeSignature as TimeSignature)
setPpq(snapshot.ppq) // Preserve original ppq for accurate export
setStatus(`已载入 ${file.name}`)
} catch (error) {
console.error(error)
setStatus('导入失败,请确认文件合法')
} finally {
event.target.value = ''
}
}
const handleAudioChange = (event: React.ChangeEvent<HTMLInputElement>) => {
const file = event.target.files?.[0]
if (!file) return
// Validate audio file type
const validAudioTypes = ['audio/mpeg', 'audio/wav', 'audio/ogg', 'audio/flac', 'audio/mp4', 'audio/aac', 'audio/x-m4a']
const validExtensions = ['.mp3', '.wav', '.ogg', '.flac', '.m4a', '.aac']
const fileName = file.name.toLowerCase()
const isValidType = validAudioTypes.includes(file.type) || file.type.startsWith('audio/')
const isValidExtension = validExtensions.some(ext => fileName.endsWith(ext))
if (!isValidType && !isValidExtension) {
setStatus(`不支持的文件格式,请选择音频文件(${validExtensions.join(', ')})`)
event.target.value = ''
return
}
const url = URL.createObjectURL(file)
setAudioUrl(url)
setStatus(`已载入音频 ${file.name}`)
event.target.value = ''
}
// Check for overlapping notes (any pitch)
const getOverlappingNotes = () => {
const overlapping: string[] = []
const sortedNotes = [...notes].sort((a, b) => a.start - b.start)
const EPSILON = 0.05 // Tolerance for floating point comparison
for (let i = 0; i < sortedNotes.length; i++) {
for (let j = i + 1; j < sortedNotes.length; j++) {
const noteA = sortedNotes[i]
const noteB = sortedNotes[j]
const noteAEnd = noteA.start + noteA.duration
// If noteB starts at or after noteA ends (with tolerance), no overlap
if (noteB.start >= noteAEnd - EPSILON) break
// True overlap: noteB starts before noteA ends
if (!overlapping.includes(noteA.id)) overlapping.push(noteA.id)
if (!overlapping.includes(noteB.id)) overlapping.push(noteB.id)
}
}
return overlapping
}
// Auto-fix overlapping notes by trimming the first note to end where the second begins
const handleFixOverlaps = () => {
const sortedNotes = [...notes].sort((a, b) => a.start - b.start)
let fixCount = 0
for (let i = 0; i < sortedNotes.length - 1; i++) {
const noteA = sortedNotes[i]
const noteB = sortedNotes[i + 1]
const noteAEnd = noteA.start + noteA.duration
// If noteA overlaps with noteB
if (noteAEnd > noteB.start) {
// Trim noteA to end at noteB's start
const newDuration = Math.max(0.01, noteB.start - noteA.start)
updateNote(noteA.id, { duration: newDuration })
fixCount++
}
}
if (fixCount > 0) {
setStatus(`已修复 ${fixCount} 个重叠音符`)
} else {
setStatus('没有检测到重叠音符')
}
}
const handleExport = () => {
const overlapping = getOverlappingNotes()
if (overlapping.length > 0) {
const confirm = window.confirm(
`检测到 ${overlapping.length} 个音符存在时间重叠(标红色的音符),这可能导致播放异常。\n\n是否仍要导出?`
)
if (!confirm) return
}
const blob = exportMidi({ notes, tempo, timeSignature, ppq })
const url = URL.createObjectURL(blob)
const anchor = document.createElement('a')
anchor.href = url
anchor.download = 'vocal-midi.mid'
anchor.click()
URL.revokeObjectURL(url)
setStatus('已导出包含歌词的 MIDI 文件')
}
return (
<div className="app-shell">
<header className="topbar">
<div>
<p className="eyebrow">歌声 MIDI 编辑器</p>
<h1>Lyric-ready Piano Roll</h1>
<p className="muted">导入、拖拽、实时修改歌词并导出标准 MIDI。</p>
</div>
<div className="actions">
<button className="icon-toggle" onClick={() => setTheme(theme === 'dark' ? 'light' : 'dark')}>
{theme === 'dark' ? (
<span className="icon" aria-label="切换到亮色">
☀️
</span>
) : (
<span className="icon" aria-label="切换到暗色">
🌙
</span>
)}
</button>
<button className="primary" onClick={handleImportClick}>
导入 MIDI
</button>
<button className="primary" onClick={handleExport}>
导出含歌词 MIDI
</button>
<button className="soft" onClick={handleFixOverlaps} title="自动消除重叠:将重叠音符的音尾提前到下一个音的音头">
消除重叠
</button>
<input ref={fileInputRef} type="file" accept=".mid,.midi" className="sr-only" onChange={handleFileChange} />
</div>
</header>
<section className="audio-bar">
<div className="audio-left">
<button className="ghost" onClick={handleAudioImportClick}>
对齐音频导入
</button>
<input
ref={audioInputRef}
type="file"
accept=".mp3,.wav,.ogg,.flac,.m4a,.aac"
className="sr-only"
onChange={handleAudioChange}
/>
<span className="audio-hint">导入后显示音频波形并与 MIDI 同步走带</span>
</div>
<div className="audio-right">
<div className="volume-control">
<span className="volume-label">MIDI</span>
<input
type="range"
min={0}
max={100}
value={midiVolume}
onChange={(e) => setMidiVolume(Number(e.target.value))}
className="volume-slider"
/>
<span className="volume-value">{midiVolume}%</span>
</div>
<div className="volume-control">
<span className="volume-label">音频</span>
<input
type="range"
min={0}
max={100}
value={audioVolume}
onChange={(e) => setAudioVolume(Number(e.target.value))}
className="volume-slider"
/>
<span className="volume-value">{audioVolume}%</span>
</div>
</div>
</section>
<section className="panel panel-split">
<div className="panel-main">
{audioUrl && (
<AudioTrack
key={audioUrl}
ref={audioScrollRef}
audioUrl={audioUrl}
muted={audioVolume === 0}
onSeek={(seconds) => seekToBeat(secondsToBeat(seconds))}
playheadSeconds={beatToSeconds(playhead)}
gridSecondWidth={gridSecondWidth}
minContentWidth={midiContentWidth}
/>
)}
<PianoRoll
notes={notes}
selectedId={selectedId}
timeSignature={timeSignature}
tempo={tempo}
playhead={playhead}
selectionStart={selectionStart}
selectionEnd={selectionEnd}
onAddNote={addNote}
onSelect={select}
onUpdateNote={updateNote}
onSeek={seekToBeat}
onScroll={(left) => {
if (audioScrollRef.current) {
audioScrollRef.current.scrollLeft = left
}
}}
onZoom={(deltaH, deltaV) => {
if (deltaH !== 0) {
setHorizontalZoom(prev => Math.max(0.5, prev + deltaH))
}
if (deltaV !== 0) {
setVerticalZoom(prev => Math.max(0.6, Math.min(2.5, prev + deltaV)))
}
}}
onPlayNote={playPreviewNote}
onFocusLyric={(noteId) => {
select(noteId)
setFocusLyricId(noteId)
}}
onSelectionChange={(start, end) => {
setSelectionStart(start)
setSelectionEnd(end)
}}
isSelectingRange={isSelectingRange}
audioDuration={audioDuration}
gridSecondWidth={gridSecondWidth}
rowHeight={rowHeight}
/>
</div>
<aside className="panel-side">
<div className="controls">
<div className="toggle" style={{ justifyContent: 'space-between' }}>
<span>水平缩放</span>
<input
type="range"
min={0.5}
max={10}
step={0.1}
value={Math.min(horizontalZoom, 10)}
onChange={(e) => setHorizontalZoom(Number(e.target.value))}
style={{ width: '140px' }}
/>
<span style={{ width: 42, textAlign: 'right' }}>{horizontalZoom.toFixed(1)}x</span>
</div>
<div className="toggle" style={{ justifyContent: 'space-between' }}>
<span>垂直缩放</span>
<input
type="range"
min={0.6}
max={2.5}
step={0.1}
value={verticalZoom}
onChange={(e) => setVerticalZoom(Number(e.target.value))}
style={{ width: '140px' }}
/>
<span style={{ width: 42, textAlign: 'right' }}>{verticalZoom.toFixed(1)}x</span>
</div>
<div className="transport">
<button
className="soft"
onClick={() => {
setPlayhead(0)
seekToBeat(0)
}}
title="回到开头"
>
⏮
</button>
<button
className="soft"
onClick={() => seekBySeconds(-2)}
title="后退 2 秒"
>
⏪ 2s
</button>
<button
className="primary"
onClick={handlePlayToggle}
disabled={!notes.length && !audioUrl}
title={isPlaying ? "暂停" : (selectionStart !== null && selectionEnd !== null ? "播放选区" : "播放")}
>
{isPlaying ? '⏸' : '▶'}
</button>
<button
className="soft"
onClick={() => seekBySeconds(2)}
title="前进 2 秒"
>
2s ⏩
</button>
<button
className="soft"
onClick={() => {
// Logic to find end of song? Max note end or audio duration
const maxNoteEnd = notes.reduce((acc, n) => Math.max(acc, n.start + n.duration), 0)
seekToBeat(Math.max(secondsToBeat(audioDuration), maxNoteEnd))
}}
title="回到结尾"
>
⏭
</button>
</div>
<div className="selection-controls">
<button
className={`soft selection-btn ${isSelectingRange ? 'active' : ''}`}
onClick={() => setIsSelectingRange(!isSelectingRange)}
title={isSelectingRange ? "退出选区模式" : "设置选区:在时间轴上拖拽选择播放范围"}
>
{isSelectingRange ? '📍 选区中' : '📍 设选区'}
</button>
{selectionStart !== null && selectionEnd !== null && (
<>
<span className="selection-info">
{selectionStart.toFixed(1)}s - {selectionEnd.toFixed(1)}s
</span>
<button
className="soft"
onClick={() => {
setSelectionStart(null)
setSelectionEnd(null)
}}
title="清除选区"
>
✕
</button>
</>
)}
</div>
<div className="status">{status}</div>
</div>
<div className="lyric-container">
<LyricTable
notes={notes}
selectedId={selectedId}
tempo={tempo}
focusLyricId={focusLyricId}
onSelect={select}
onUpdate={updateNote}
onFocusHandled={() => setFocusLyricId(null)}
/>
</div>
</aside>
</section>
<audio
ref={audioRef}
src={audioUrl ?? undefined}
preload="auto"
className="sr-only"
onLoadedMetadata={(e) => {
setAudioDuration(e.currentTarget.duration)
// Ensure volume is set when audio loads
e.currentTarget.volume = audioVolume / 100
}}
/>
</div>
)
}
export default App
@@ -0,0 +1,182 @@
import { useEffect, useRef, forwardRef, useState } from 'react'
import WaveSurfer from 'wavesurfer.js'
import { PITCH_WIDTH } from '../constants'
export type AudioTrackProps = {
audioUrl: string | null
muted: boolean
onSeek: (seconds: number) => void
mediaElement?: HTMLAudioElement | null
playheadSeconds: number
gridSecondWidth: number
minContentWidth?: number // Minimum width to match MIDI editor area
}
export const AudioTrack = forwardRef<HTMLDivElement, AudioTrackProps>(
({ audioUrl, muted, onSeek, playheadSeconds, gridSecondWidth, minContentWidth = 0 }, ref) => {
const containerRef = useRef<HTMLDivElement | null>(null)
const waveRef = useRef<WaveSurfer | null>(null)
const [waveWidth, setWaveWidth] = useState(0)
useEffect(() => {
if (!containerRef.current) return
if (!audioUrl) {
try {
waveRef.current?.destroy()
} catch {
// ignore teardown errors
}
waveRef.current = null
setWaveWidth(0)
return
}
let cancelled = false
// Clean up existing instance
if (waveRef.current) {
try {
waveRef.current.destroy()
} catch {
// ignore teardown errors
}
}
waveRef.current = WaveSurfer.create({
container: containerRef.current,
waveColor: '#4b64bc',
progressColor: '#4b64bc',
cursorColor: 'transparent',
barWidth: 2,
barGap: 2,
height: 60,
normalize: true,
minPxPerSec: gridSecondWidth,
interact: false,
hideScrollbar: true,
autoScroll: false,
})
waveRef.current.load(audioUrl).catch(() => null)
waveRef.current.on('error', () => null)
waveRef.current.on('ready', () => {
if (cancelled || !waveRef.current) return
const duration = waveRef.current.getDuration()
const requiredWidth = duration * gridSecondWidth
setWaveWidth(requiredWidth)
})
return () => {
cancelled = true
try {
waveRef.current?.destroy()
} catch {
// ignore teardown errors
}
waveRef.current = null
}
}, [audioUrl, gridSecondWidth])
useEffect(() => {
if (!waveRef.current) return
waveRef.current.setOptions({
waveColor: muted ? '#9aa6b2' : '#4b64bc',
progressColor: muted ? '#c0c9d4' : '#4b64bc',
})
}, [muted])
if (!audioUrl) return null
// Content width should be at least as wide as MIDI editor
const contentWidth = Math.max(waveWidth, minContentWidth)
return (
<div
className="audio-track-row"
style={{
display: 'flex',
borderBottom: '1px solid var(--border-soft)',
height: '70px',
flexShrink: 0
}}
>
<div
className="audio-gutter"
style={{
width: PITCH_WIDTH,
flexShrink: 0,
background: 'var(--panel-strong)',
borderRight: '1px solid var(--border-subtle)',
display: 'flex',
alignItems: 'center',
justifyContent: 'center',
fontSize: '11px',
color: 'var(--text-muted)',
fontWeight: 600,
}}
>
AUDIO
</div>
{/* Scroll Mask - Controlled by parent via ref */}
<div
ref={ref}
className="audio-scroll-mask"
style={{
flex: 1,
overflow: 'hidden',
position: 'relative',
background: 'var(--panel-soft)',
}}
onClick={(e) => {
const rect = e.currentTarget.getBoundingClientRect()
const scrollMask = e.currentTarget as HTMLDivElement
const x = e.clientX - rect.left + scrollMask.scrollLeft
const seconds = x / gridSecondWidth
onSeek(seconds)
}}
>
{/* Container that matches MIDI editor width */}
<div
className="audio-content"
style={{
width: contentWidth > 0 ? contentWidth : '100%',
height: '100%',
position: 'relative'
}}
>
{/* WaveSurfer container - only as wide as audio */}
<div
ref={containerRef}
className="wave-container"
style={{
width: waveWidth > 0 ? waveWidth : '100%',
height: '100%',
position: 'absolute',
left: 0,
top: 0
}}
/>
{/* Custom Playhead */}
<div
className="audio-playhead"
style={{
position: 'absolute',
top: 0,
bottom: 0,
width: '2px',
background: '#ff7043',
boxShadow: '0 0 12px rgba(255, 112, 67, 0.6)',
left: playheadSeconds * gridSecondWidth,
zIndex: 10,
pointerEvents: 'none',
}}
/>
</div>
</div>
</div>
)
}
)
@@ -0,0 +1,288 @@
import { useEffect, useMemo, useRef, useState } from 'react'
import type { NoteEvent } from '../types'
export type LyricTableProps = {
notes: NoteEvent[]
selectedId: string | null
tempo: number
focusLyricId: string | null
onSelect: (id: string | null) => void
onUpdate: (id: string, patch: Partial<NoteEvent>) => void
onScrollToNote?: (noteId: string) => void
onFocusHandled?: () => void
}
const formatSeconds = (beats: number, tempo: number) => {
const seconds = beats * (60 / tempo)
return Number.parseFloat(seconds.toFixed(2))
}
const secondsToBeats = (seconds: number, tempo: number) => {
return seconds * (tempo / 60)
}
// Editable cell with confirmation
function EditableCell({
value,
noteId,
field,
tempo,
onConfirm,
type = 'number',
min,
step
}: {
value: number
noteId: string
field: 'midi' | 'start' | 'end'
tempo: number
onConfirm: (noteId: string, field: string, value: number) => void
type?: string
min?: number
step?: number
}) {
const displayValue = field === 'midi' ? value : formatSeconds(value, tempo)
const [localValue, setLocalValue] = useState(String(displayValue))
const [isDirty, setIsDirty] = useState(false)
const inputRef = useRef<HTMLInputElement>(null)
// Sync with external value when it changes (and not dirty)
useEffect(() => {
if (!isDirty) {
setLocalValue(String(displayValue))
}
}, [displayValue, isDirty])
const handleChange = (e: React.ChangeEvent<HTMLInputElement>) => {
setLocalValue(e.target.value)
setIsDirty(true)
}
const handleConfirm = () => {
const parsed = parseFloat(localValue)
if (!isNaN(parsed)) {
if (field === 'midi') {
if (parsed >= 0 && parsed <= 127) {
onConfirm(noteId, field, Math.round(parsed))
}
} else {
if (parsed >= 0) {
onConfirm(noteId, field, secondsToBeats(parsed, tempo))
}
}
}
setIsDirty(false)
}
const handleKeyDown = (e: React.KeyboardEvent) => {
if (e.key === 'Enter') {
e.preventDefault()
handleConfirm()
inputRef.current?.blur()
} else if (e.key === 'Escape') {
setLocalValue(String(displayValue))
setIsDirty(false)
inputRef.current?.blur()
}
}
const handleBlur = () => {
if (isDirty) {
// Reset to original on blur without confirm
setLocalValue(String(displayValue))
setIsDirty(false)
}
}
return (
<div className="editable-cell">
<input
ref={inputRef}
className={`lyric-meta-input ${isDirty ? 'lyric-meta-dirty' : ''}`}
type={type}
min={min}
step={step}
value={localValue}
onChange={handleChange}
onKeyDown={handleKeyDown}
onBlur={handleBlur}
onClick={(e) => e.stopPropagation()}
/>
{isDirty && (
<button
className="confirm-btn"
onMouseDown={(e) => {
e.preventDefault() // Prevent input blur
e.stopPropagation()
}}
onClick={(e) => {
e.stopPropagation()
handleConfirm()
}}
title="确认修改 (Enter)"
>
✓
</button>
)}
</div>
)
}
export function LyricTable({ notes, selectedId, tempo, focusLyricId, onSelect, onUpdate, onScrollToNote, onFocusHandled }: LyricTableProps) {
const listRef = useRef<HTMLDivElement | null>(null)
const inputRefs = useRef<Map<string, HTMLInputElement>>(new Map())
const sorted = useMemo(() => [...notes].sort((a, b) => a.start - b.start), [notes])
// Scroll to selected note (no auto-focus on single click)
useEffect(() => {
if (!selectedId || !listRef.current) return
const target = listRef.current.querySelector<HTMLDivElement>(`[data-note-id="${selectedId}"]`)
if (target) {
target.scrollIntoView({ block: 'nearest', behavior: 'smooth' })
}
}, [selectedId])
// Focus lyric input when requested (double-click on note or click on list row)
useEffect(() => {
if (!focusLyricId) return
const input = inputRefs.current.get(focusLyricId)
if (input) {
setTimeout(() => {
input.focus()
input.select()
}, 50)
}
onFocusHandled?.()
}, [focusLyricId, onFocusHandled])
// Fill lyrics from selected note onwards
const handleBulkFill = (bulkText: string) => {
if (!sorted.length) return
const chars = Array.from(bulkText.replace(/\s+/g, ''))
let startIndex = 0
if (selectedId) {
const selectedIndex = sorted.findIndex(n => n.id === selectedId)
if (selectedIndex >= 0) {
startIndex = selectedIndex
}
}
let charIndex = 0
for (let i = startIndex; i < sorted.length && charIndex < chars.length; i++) {
onUpdate(sorted[i].id, { lyric: chars[charIndex] })
charIndex++
}
}
const handleRowClick = (noteId: string) => {
onSelect(noteId)
onScrollToNote?.(noteId)
}
const handleFieldConfirm = (noteId: string, field: string, value: number) => {
const note = notes.find(n => n.id === noteId)
if (!note) return
if (field === 'midi') {
onUpdate(noteId, { midi: value })
} else if (field === 'start') {
// Keep END the same, adjust duration accordingly
const currentEnd = note.start + note.duration
const newDuration = Math.max(0.01, currentEnd - value)
onUpdate(noteId, { start: value, duration: newDuration })
} else if (field === 'end') {
// End changed, update duration
const newDuration = Math.max(0.01, value - note.start)
onUpdate(noteId, { duration: newDuration })
}
}
return (
<div className="lyric-card">
<div className="lyric-bulk">
<textarea
className="lyric-bulk-input"
rows={2}
placeholder={selectedId ? "从选中音符开始按字填充" : "输入歌词,点击按字填充"}
onKeyDown={(e) => {
if (e.key === 'Enter' && !e.shiftKey) {
e.preventDefault()
handleBulkFill(e.currentTarget.value)
}
}}
/>
<button
className="soft"
type="button"
onClick={(e) => {
const textarea = e.currentTarget.previousElementSibling as HTMLTextAreaElement
handleBulkFill(textarea.value)
}}
>
按字<br/>填充
</button>
</div>
<div className="lyric-header" style={{ flexShrink: 0 }}>
<div>LYRIC</div>
<div>PITCH</div>
<div>START</div>
<div>END</div>
</div>
<div className="lyric-list" ref={listRef}>
{sorted.map((note) => (
<div
key={note.id}
className={`lyric-row ${selectedId === note.id ? 'lyric-row-active' : ''}`}
data-note-id={note.id}
onClick={() => handleRowClick(note.id)}
>
<input
ref={(el) => {
if (el) {
inputRefs.current.set(note.id, el)
} else {
inputRefs.current.delete(note.id)
}
}}
className="lyric-input"
value={note.lyric}
placeholder="Type lyric"
onChange={(event) => onUpdate(note.id, { lyric: event.target.value })}
onClick={(e) => e.stopPropagation()}
/>
<EditableCell
value={note.midi}
noteId={note.id}
field="midi"
tempo={tempo}
onConfirm={handleFieldConfirm}
min={0}
/>
<EditableCell
value={note.start}
noteId={note.id}
field="start"
tempo={tempo}
onConfirm={handleFieldConfirm}
min={0}
step={0.01}
/>
<EditableCell
value={note.start + note.duration}
noteId={note.id}
field="end"
tempo={tempo}
onConfirm={handleFieldConfirm}
min={0}
step={0.01}
/>
</div>
))}
{sorted.length === 0 && <div className="lyric-empty">Import或双击钢琴卷帘以添加音符</div>}
</div>
</div>
)
}
@@ -0,0 +1,704 @@
import { useEffect, useMemo, useRef, useState, useCallback, memo } from 'react'
import type React from 'react'
import type { NoteEvent, TimeSignature } from '../types'
import { PITCH_WIDTH, LOW_NOTE, HIGH_NOTE } from '../constants'
const midiToName = (midi: number) => {
const names = ['C', 'C#', 'D', 'D#', 'E', 'F', 'F#', 'G', 'G#', 'A', 'A#', 'B']
const octave = Math.floor(midi / 12) - 1
return `${names[midi % 12]}${octave}`
}
// Memoized note component to prevent unnecessary re-renders
const NoteChip = memo(function NoteChip({
note,
left,
top,
width,
height,
fontSize,
isSelected,
isOverlapping,
onPointerDown,
onDoubleClick,
}: {
note: NoteEvent
left: number
top: number
width: number
height: number
fontSize: number
isSelected: boolean
isOverlapping: boolean
onPointerDown: (event: React.PointerEvent<HTMLDivElement>, mode: 'move' | 'resize-start' | 'resize-end') => void
onDoubleClick: (event: React.MouseEvent<HTMLDivElement>) => void
}) {
return (
<div
className={`note-chip ${isSelected ? 'note-active' : ''} ${isOverlapping ? 'note-overlap' : ''}`}
style={{
left,
top: top + 1,
width,
height,
willChange: 'transform', // GPU acceleration hint
}}
onPointerDown={(e) => onPointerDown(e, 'move')}
onDoubleClick={onDoubleClick}
>
<div className="note-label" style={{ fontSize }}>
<span>{note.lyric || '\u00a0'}</span>
</div>
<div className="note-handle start" onPointerDown={(e) => { e.stopPropagation(); onPointerDown(e, 'resize-start') }} />
<div className="note-handle end" onPointerDown={(e) => { e.stopPropagation(); onPointerDown(e, 'resize-end') }} />
</div>
)
})
// Dynamic snap based on zoom level - higher zoom = finer snap
const getSnapSeconds = (gridSecondWidth: number) => {
// At base width (80px/s), snap is 0.1s
// At 2x zoom (160px/s), snap is 0.05s
// At 4x zoom (320px/s), snap is 0.025s
// At 8x zoom (640px/s), snap is 0.01s
const baseSnap = 0.1
const zoomFactor = gridSecondWidth / 80
return Math.max(0.01, baseSnap / zoomFactor)
}
const snapSeconds = (value: number, gridSecondWidth: number) => {
const snap = getSnapSeconds(gridSecondWidth)
return Math.max(0, Math.round(value / snap) * snap)
}
export type PianoRollProps = {
notes: NoteEvent[]
selectedId: string | null
timeSignature: TimeSignature
tempo: number
playhead: number // in beats
selectionStart: number | null // in seconds
selectionEnd: number | null // in seconds
onAddNote: (note: Partial<NoteEvent>) => NoteEvent
onUpdateNote: (id: string, patch: Partial<NoteEvent>) => void
onSelect: (id: string | null) => void
onSeek: (beat: number) => void
onScroll?: (left: number) => void
onZoom?: (deltaH: number, deltaV: number) => void
onPlayNote?: (midi: number) => void
onFocusLyric?: (noteId: string) => void
onSelectionChange?: (start: number | null, end: number | null) => void
isSelectingRange?: boolean
audioDuration?: number
gridSecondWidth: number
rowHeight: number
}
export function PianoRoll({
notes,
selectedId,
timeSignature: _timeSignature,
tempo,
playhead,
selectionStart,
selectionEnd,
onAddNote,
onSelect,
onUpdateNote,
onSeek,
onScroll,
onZoom,
onPlayNote,
onFocusLyric,
onSelectionChange,
isSelectingRange = false,
audioDuration = 0,
gridSecondWidth,
rowHeight
}: PianoRollProps) {
const scrollContainerRef = useRef<HTMLDivElement | null>(null)
const rulerScrollRef = useRef<HTMLDivElement | null>(null)
const [scrollTop, setScrollTop] = useState(0)
const [scrollLeft, setScrollLeft] = useState(0)
const [viewportWidth, setViewportWidth] = useState(800)
const [viewportHeight, setViewportHeight] = useState(400)
const dragRef = useRef<{
id: string
mode: 'move' | 'resize-start' | 'resize-end'
originX: number
originY: number
startSeconds: number
durationSeconds: number
midi: number
lastMidi?: number // Track last midi for pitch change sound
} | null>(null)
// Selection drag state
const selectionDragRef = useRef<{
startX: number
startSeconds: number
} | null>(null)
// Store callbacks in refs to avoid stale closures in event handlers
const onPlayNoteRef = useRef(onPlayNote)
const onUpdateNoteRef = useRef(onUpdateNote)
useEffect(() => {
onPlayNoteRef.current = onPlayNote
onUpdateNoteRef.current = onUpdateNote
}, [onPlayNote, onUpdateNote])
// Conversion helpers
const beatToSeconds = useCallback((beat: number) => beat * (60 / tempo), [tempo])
const secondsToBeat = useCallback((seconds: number) => seconds / (60 / tempo), [tempo])
// Calculate dimensions
const totalRows = HIGH_NOTE - LOW_NOTE + 1
const contentHeight = totalRows * rowHeight
const [containerWidth, setContainerWidth] = useState(1200)
// Track container size
useEffect(() => {
const container = scrollContainerRef.current
if (!container) return
const observer = new ResizeObserver((entries) => {
for (const entry of entries) {
setContainerWidth(entry.contentRect.width)
setViewportWidth(entry.contentRect.width)
setViewportHeight(entry.contentRect.height)
}
})
observer.observe(container)
return () => observer.disconnect()
}, [])
const maxSeconds = useMemo(() => {
const noteEndSeconds = notes.reduce((acc, n) => {
const endBeat = n.start + n.duration
return Math.max(acc, beatToSeconds(endBeat))
}, 8)
// Ensure grid extends at least 2x the visible area for smoother scrolling
const minSecondsForView = (containerWidth / gridSecondWidth) * 2
return Math.max(noteEndSeconds + 10, audioDuration + 10, minSecondsForView, 30)
}, [notes, audioDuration, beatToSeconds, containerWidth, gridSecondWidth])
const contentWidth = maxSeconds * gridSecondWidth
// Drag handlers - use refs to avoid stale closure issues
const handlePointerMove = useCallback((event: PointerEvent) => {
const drag = dragRef.current
if (!drag) return
const dxSeconds = (event.clientX - drag.originX) / gridSecondWidth
const dy = (event.clientY - drag.originY) / rowHeight
if (drag.mode === 'move') {
const nextSeconds = snapSeconds(drag.startSeconds + dxSeconds, gridSecondWidth)
const nextMidi = Math.min(HIGH_NOTE, Math.max(LOW_NOTE, Math.round(drag.midi - dy)))
// Play sound when pitch changes
if (nextMidi !== drag.lastMidi && onPlayNoteRef.current) {
onPlayNoteRef.current(nextMidi)
drag.lastMidi = nextMidi
}
onUpdateNoteRef.current(drag.id, {
start: secondsToBeat(nextSeconds),
midi: nextMidi
})
}
if (drag.mode === 'resize-start') {
const nextSeconds = snapSeconds(drag.startSeconds + dxSeconds, gridSecondWidth)
const delta = drag.startSeconds - nextSeconds
const nextDurationSeconds = Math.max(0.05, drag.durationSeconds + delta)
onUpdateNoteRef.current(drag.id, {
start: secondsToBeat(nextSeconds),
duration: secondsToBeat(nextDurationSeconds)
})
}
if (drag.mode === 'resize-end') {
const nextDurationSeconds = Math.max(0.05, snapSeconds(drag.durationSeconds + dxSeconds, gridSecondWidth))
onUpdateNoteRef.current(drag.id, { duration: secondsToBeat(nextDurationSeconds) })
}
}, [gridSecondWidth, rowHeight, secondsToBeat])
const handlePointerUp = useCallback(() => {
dragRef.current = null
window.removeEventListener('pointermove', handlePointerMove)
window.removeEventListener('pointerup', handlePointerUp)
}, [handlePointerMove])
useEffect(() => {
return () => {
window.removeEventListener('pointermove', handlePointerMove)
window.removeEventListener('pointerup', handlePointerUp)
}
}, [handlePointerMove, handlePointerUp])
// Scroll sync
useEffect(() => {
const container = scrollContainerRef.current
const ruler = rulerScrollRef.current
if (!container || !ruler) return
const handleScroll = () => {
ruler.scrollLeft = container.scrollLeft
setScrollTop(container.scrollTop)
setScrollLeft(container.scrollLeft)
if (onScroll) onScroll(container.scrollLeft)
}
container.addEventListener('scroll', handleScroll)
return () => container.removeEventListener('scroll', handleScroll)
}, [onScroll])
// Zoom support via wheel/trackpad
// Mac: Cmd+滚轮 (水平缩放), Cmd+Shift+滚轮 (垂直缩放), 或双指捏合
// Windows/Linux: Ctrl+滚轮 (水平缩放), Ctrl+Shift+滚轮 (垂直缩放)
useEffect(() => {
const container = scrollContainerRef.current
if (!container || !onZoom) return
const handleWheel = (e: WheelEvent) => {
// Ctrl (Windows/Linux/捏合) or Cmd (Mac) triggers zoom
const isZoomTrigger = e.ctrlKey || e.metaKey
if (isZoomTrigger) {
e.preventDefault()
e.stopPropagation()
// Use deltaY for zoom amount, normalize for different input methods
// Pinch gestures typically have smaller delta values
let delta = -e.deltaY
if (Math.abs(delta) > 10) {
// Likely a mouse wheel, scale down
delta = delta * 0.01
} else {
// Likely a trackpad pinch, scale appropriately
delta = delta * 0.05
}
// Shift or Alt/Option for vertical zoom, otherwise horizontal
if (e.shiftKey || e.altKey) {
onZoom(0, delta)
} else {
onZoom(delta, 0)
}
}
}
container.addEventListener('wheel', handleWheel, { passive: false })
return () => container.removeEventListener('wheel', handleWheel)
}, [onZoom])
// Playhead auto-scroll
useEffect(() => {
if (!scrollContainerRef.current) return
const container = scrollContainerRef.current
const playheadX = beatToSeconds(playhead) * gridSecondWidth
const viewStart = container.scrollLeft
const viewEnd = viewStart + container.clientWidth
if (playheadX > viewEnd) {
container.scrollLeft = playheadX
} else if (playheadX < viewStart) {
container.scrollLeft = playheadX
}
}, [playhead, gridSecondWidth, beatToSeconds])
// Selection auto-scroll
useEffect(() => {
if (!scrollContainerRef.current || !selectedId) return
const note = notes.find((n) => n.id === selectedId)
if (!note) return
const container = scrollContainerRef.current
const noteX = beatToSeconds(note.start) * gridSecondWidth
const noteY = (HIGH_NOTE - note.midi) * rowHeight
const viewStart = container.scrollLeft
const viewEnd = viewStart + container.clientWidth
if (noteX < viewStart + 50 || noteX > viewEnd - 50) {
container.scrollLeft = Math.max(0, noteX - container.clientWidth * 0.35)
}
const viewTop = container.scrollTop
const viewBottom = viewTop + container.clientHeight
if (noteY < viewTop || noteY > viewBottom - rowHeight) {
container.scrollTop = Math.max(0, noteY - container.clientHeight * 0.4)
}
}, [selectedId, notes, gridSecondWidth, rowHeight, beatToSeconds])
const handleGridDoubleClick = (event: React.MouseEvent<HTMLDivElement>) => {
// Only add note if clicking on empty space (not on a note)
const target = event.target as HTMLElement
if (target.closest('.note-chip')) return
if (!scrollContainerRef.current) return
const container = scrollContainerRef.current
const rect = container.getBoundingClientRect()
const x = event.clientX - rect.left + container.scrollLeft
const y = event.clientY - rect.top + container.scrollTop
const seconds = snapSeconds(x / gridSecondWidth, gridSecondWidth)
const pitch = Math.min(HIGH_NOTE, Math.max(LOW_NOTE, HIGH_NOTE - Math.floor(y / rowHeight)))
const created = onAddNote({
start: secondsToBeat(seconds),
midi: pitch,
duration: secondsToBeat(0.5),
lyric: ''
})
onSelect(created.id)
}
const startDrag = (
event: React.PointerEvent<HTMLDivElement>,
note: NoteEvent,
mode: 'move' | 'resize-start' | 'resize-end',
) => {
event.preventDefault()
event.stopPropagation()
dragRef.current = {
id: note.id,
mode,
originX: event.clientX,
originY: event.clientY,
startSeconds: beatToSeconds(note.start),
durationSeconds: beatToSeconds(note.duration),
midi: note.midi,
lastMidi: note.midi, // Initialize last midi
}
window.addEventListener('pointermove', handlePointerMove)
window.addEventListener('pointerup', handlePointerUp)
onSelect(note.id)
// Play sound when clicking/selecting note
if (onPlayNote) {
onPlayNote(note.midi)
}
}
// Second-based ruler labels
const secondLabels = useMemo(() => {
const labels = [] as Array<{ left: number; label: string }>
const totalSeconds = Math.ceil(maxSeconds)
for (let s = 0; s <= totalSeconds; s += 1) {
labels.push({ left: s * gridSecondWidth, label: `${s}s` })
}
return labels
}, [maxSeconds, gridSecondWidth])
// Piano keys
const pitchRows = useMemo(() => {
const rows = [] as Array<{ midi: number; isBlack: boolean; label: string; isC: boolean }>
const black = new Set([1, 3, 6, 8, 10])
for (let p = HIGH_NOTE; p >= LOW_NOTE; p -= 1) {
const name = midiToName(p)
const isC = p % 12 === 0
rows.push({ midi: p, isBlack: black.has(p % 12), label: name, isC })
}
return rows
}, [])
// Detect overlapping notes using optimized sweep line algorithm
const overlappingNoteIds = useMemo(() => {
if (notes.length < 2) return new Set<string>()
const overlapping = new Set<string>()
const sortedNotes = [...notes].sort((a, b) => a.start - b.start)
const EPSILON = 0.05 // Tolerance for floating point comparison
// Use a sliding window approach - more efficient for typical music data
// Active notes: notes that haven't ended yet
const activeNotes: typeof sortedNotes = []
for (const note of sortedNotes) {
// Remove notes that have ended before current note starts
while (activeNotes.length > 0) {
const firstActive = activeNotes[0]
const firstActiveEnd = firstActive.start + firstActive.duration
if (firstActiveEnd <= note.start + EPSILON) {
activeNotes.shift()
} else {
break
}
}
// Check overlap with remaining active notes
for (const activeNote of activeNotes) {
const activeEnd = activeNote.start + activeNote.duration
if (note.start < activeEnd - EPSILON) {
overlapping.add(activeNote.id)
overlapping.add(note.id)
}
}
// Add current note to active set (maintain sorted order by end time)
const noteEnd = note.start + note.duration
let insertIndex = activeNotes.length
for (let i = 0; i < activeNotes.length; i++) {
const aEnd = activeNotes[i].start + activeNotes[i].duration
if (noteEnd < aEnd) {
insertIndex = i
break
}
}
activeNotes.splice(insertIndex, 0, note)
}
return overlapping
}, [notes])
// Calculate visible area with buffer for smooth scrolling
const BUFFER_PX = 200 // Render notes slightly outside viewport for smooth scrolling
const visibleArea = useMemo(() => {
return {
left: Math.max(0, scrollLeft - BUFFER_PX),
right: scrollLeft + viewportWidth + BUFFER_PX,
top: Math.max(0, scrollTop - BUFFER_PX),
bottom: scrollTop + viewportHeight + BUFFER_PX,
}
}, [scrollLeft, scrollTop, viewportWidth, viewportHeight])
// Filter notes to only render visible ones (virtualization)
const visibleNotes = useMemo(() => {
return notes.filter(note => {
const noteSeconds = beatToSeconds(note.start)
const noteDurationSeconds = beatToSeconds(note.duration)
const noteLeft = noteSeconds * gridSecondWidth
const noteRight = noteLeft + noteDurationSeconds * gridSecondWidth
const noteTop = (HIGH_NOTE - note.midi) * rowHeight
const noteBottom = noteTop + rowHeight
// Check if note intersects with visible area
const horizontallyVisible = noteRight >= visibleArea.left && noteLeft <= visibleArea.right
const verticallyVisible = noteBottom >= visibleArea.top && noteTop <= visibleArea.bottom
return horizontallyVisible && verticallyVisible
})
}, [notes, visibleArea, gridSecondWidth, rowHeight, beatToSeconds])
// Calculate visible grid lines (virtualization)
const visibleGridLines = useMemo(() => {
const startSecond = Math.max(0, Math.floor(visibleArea.left / gridSecondWidth) - 1)
const endSecond = Math.ceil(visibleArea.right / gridSecondWidth) + 1
const startRow = Math.max(0, Math.floor(visibleArea.top / rowHeight) - 1)
const endRow = Math.min(totalRows, Math.ceil(visibleArea.bottom / rowHeight) + 1)
return {
horizontalLines: Array.from({ length: endRow - startRow + 1 }, (_, i) => startRow + i),
verticalLines: Array.from({ length: endSecond - startSecond + 1 }, (_, i) => startSecond + i),
}
}, [visibleArea, gridSecondWidth, rowHeight, totalRows])
const playheadSeconds = beatToSeconds(playhead)
// Selection drag handlers
const handleRulerPointerDown = (event: React.PointerEvent<HTMLDivElement>) => {
if (!isSelectingRange) {
// Normal click to seek
const rect = event.currentTarget.getBoundingClientRect()
const x = event.clientX - rect.left + (rulerScrollRef.current?.scrollLeft ?? 0)
const seconds = x / gridSecondWidth
onSeek(secondsToBeat(seconds))
return
}
// Start selection drag
const rect = event.currentTarget.getBoundingClientRect()
const x = event.clientX - rect.left + (rulerScrollRef.current?.scrollLeft ?? 0)
const seconds = Math.max(0, x / gridSecondWidth)
selectionDragRef.current = {
startX: event.clientX,
startSeconds: seconds,
}
onSelectionChange?.(seconds, seconds)
const handleSelectionMove = (e: PointerEvent) => {
if (!selectionDragRef.current) return
const currentX = e.clientX - rect.left + (rulerScrollRef.current?.scrollLeft ?? 0)
const currentSeconds = Math.max(0, currentX / gridSecondWidth)
const start = Math.min(selectionDragRef.current.startSeconds, currentSeconds)
const end = Math.max(selectionDragRef.current.startSeconds, currentSeconds)
onSelectionChange?.(start, end)
}
const handleSelectionUp = () => {
selectionDragRef.current = null
window.removeEventListener('pointermove', handleSelectionMove)
window.removeEventListener('pointerup', handleSelectionUp)
}
window.addEventListener('pointermove', handleSelectionMove)
window.addEventListener('pointerup', handleSelectionUp)
}
return (
<div className="piano-shell">
{/* Ruler */}
<div className="ruler-shell">
<div className="ruler-spacer" style={{ width: PITCH_WIDTH, flexShrink: 0 }} />
<div
ref={rulerScrollRef}
className={`ruler-scroll ${isSelectingRange ? 'selecting' : ''}`}
onPointerDown={handleRulerPointerDown}
>
<div className="ruler" style={{ width: contentWidth }}>
{secondLabels.map((mark) => (
<div key={mark.left} className="measure-mark" style={{ left: mark.left }}>
<span>{mark.label}</span>
</div>
))}
{/* Selection range indicator */}
{selectionStart !== null && selectionEnd !== null && selectionEnd > selectionStart && (
<div
className="selection-range"
style={{
left: selectionStart * gridSecondWidth,
width: (selectionEnd - selectionStart) * gridSecondWidth
}}
/>
)}
{/* Ruler playhead indicator */}
<div
className="ruler-playhead"
style={{ left: playheadSeconds * gridSecondWidth }}
/>
</div>
</div>
</div>
{/* Main content area */}
<div className="roll-body">
{/* Piano keys - synced with vertical scroll */}
<div className="pitch-rail" style={{ width: PITCH_WIDTH }}>
<div
className="pitch-rail-inner"
style={{
transform: `translateY(${-scrollTop}px)`,
height: contentHeight
}}
>
{pitchRows.map((pitch) => (
<div
key={pitch.midi}
className={`pitch-cell ${pitch.isBlack ? 'pitch-black' : 'pitch-white'} ${pitch.isC ? 'pitch-c' : ''}`}
style={{ height: rowHeight, cursor: 'pointer' }}
onClick={() => onPlayNote?.(pitch.midi)}
onMouseDown={(e) => e.preventDefault()}
>
<span className="pitch-label">{pitch.label}</span>
</div>
))}
</div>
</div>
{/* Scrollable grid area */}
<div
ref={scrollContainerRef}
className="roll-grid"
onDoubleClick={handleGridDoubleClick}
>
<div
className="grid-content"
style={{
width: contentWidth,
height: contentHeight,
position: 'relative'
}}
>
{/* SVG Grid - virtualized for performance */}
<svg
className="grid-svg"
width={contentWidth}
height={contentHeight}
style={{ position: 'absolute', top: 0, left: 0, pointerEvents: 'none' }}
>
{/* Horizontal lines (pitch rows) - only visible ones */}
{visibleGridLines.horizontalLines.map(i => (
<line
key={`h-${i}`}
x1={visibleArea.left}
y1={i * rowHeight}
x2={visibleArea.right}
y2={i * rowHeight}
stroke="var(--grid-line-minor)"
strokeWidth={1}
/>
))}
{/* Vertical lines (seconds) - only visible ones */}
{visibleGridLines.verticalLines.map(i => (
<line
key={`v-${i}`}
x1={i * gridSecondWidth}
y1={visibleArea.top}
x2={i * gridSecondWidth}
y2={visibleArea.bottom}
stroke="var(--grid-line-minor)"
strokeWidth={1}
/>
))}
</svg>
{/* Selection range in grid */}
{selectionStart !== null && selectionEnd !== null && selectionEnd > selectionStart && (
<div
className="grid-selection-range"
style={{
left: selectionStart * gridSecondWidth,
width: (selectionEnd - selectionStart) * gridSecondWidth,
height: contentHeight
}}
/>
)}
{/* Playhead */}
<div
className="playhead"
style={{
left: playheadSeconds * gridSecondWidth,
height: contentHeight
}}
/>
{/* Notes - virtualized: only render visible notes */}
{visibleNotes.map((note) => {
const noteSeconds = beatToSeconds(note.start)
const noteDurationSeconds = beatToSeconds(note.duration)
const left = noteSeconds * gridSecondWidth
const top = (HIGH_NOTE - note.midi) * rowHeight
const noteWidthPx = Math.max(noteDurationSeconds * gridSecondWidth, 4)
const noteHeight = rowHeight - 2
const isOverlapping = overlappingNoteIds.has(note.id)
// Dynamic font size based on row height (base: 12px at 20px row height)
const fontSize = Math.max(10, Math.min(24, rowHeight * 0.6))
return (
<NoteChip
key={note.id}
note={note}
left={left}
top={top}
width={noteWidthPx}
height={noteHeight}
fontSize={fontSize}
isSelected={selectedId === note.id}
isOverlapping={isOverlapping}
onPointerDown={(event, mode) => startDrag(event, note, mode)}
onDoubleClick={(event) => {
event.stopPropagation()
onFocusLyric?.(note.id)
}}
/>
)
})}
</div>
</div>
</div>
</div>
)
}
@@ -0,0 +1,8 @@
// Base values used for scaling; actual runtime values are derived in components
export const BASE_GRID_SECOND_WIDTH = 80
export const BASE_ROW_HEIGHT = 20
export const PITCH_WIDTH = 60
// C-1 to C8 range (MIDI note numbers)
// LOW_NOTE = 0 to support SP markers (pitch=0) in some MIDI files
export const LOW_NOTE = 0 // C-1 (also supports pitch=0 for SP markers)
export const HIGH_NOTE = 108 // C8
@@ -0,0 +1,37 @@
@tailwind base;
@tailwind components;
@tailwind utilities;
:root {
font-family: 'Space Grotesk', 'IBM Plex Sans', system-ui, sans-serif;
color: var(--text-primary);
background: radial-gradient(circle at 20% 20%, rgba(72, 228, 194, 0.08), transparent 35%),
radial-gradient(circle at 80% 0%, rgba(75, 100, 188, 0.24), transparent 40%),
#0f1528;
text-rendering: optimizeLegibility;
-webkit-font-smoothing: antialiased;
}
:root[data-theme='light'] {
background: radial-gradient(circle at 20% 20%, rgba(63, 140, 255, 0.08), transparent 35%),
radial-gradient(circle at 80% 0%, rgba(75, 100, 188, 0.14), transparent 40%),
#f5f7fb;
}
* {
box-sizing: border-box;
}
body {
margin: 0;
min-height: 100vh;
background: transparent;
}
#root {
min-height: 100vh;
}
a {
color: inherit;
}
@@ -0,0 +1,224 @@
import { Midi } from '@tonejs/midi'
import { writeMidi } from 'midi-file'
import type { MidiData, MidiEvent } from 'midi-file'
import type { NoteEvent, ProjectSnapshot, TimeSignature } from '../types'
const DEFAULT_SIGNATURE: TimeSignature = [4, 4]
// Decode UTF-8 byte string (latin1 encoded) to proper Unicode string
// This matches: text.encode("latin1").decode("utf-8") in Python
function decodeUtf8ByteString(byteString: string): string {
try {
const bytes = new Uint8Array(byteString.length)
for (let i = 0; i < byteString.length; i++) {
bytes[i] = byteString.charCodeAt(i)
}
return new TextDecoder('utf-8').decode(bytes)
} catch {
return byteString
}
}
// Encode Unicode string to UTF-8 byte string (latin1 encoding)
// This matches: text.encode("utf-8").decode("latin1") in Python
function encodeUtf8ByteString(text: string): string {
const bytes = new TextEncoder().encode(text)
let output = ''
bytes.forEach((b) => {
output += String.fromCharCode(b)
})
return output
}
export async function importMidiFile(file: File): Promise<ProjectSnapshot> {
const buffer = await file.arrayBuffer()
return parseMidiBuffer(buffer)
}
export async function parseMidiBuffer(buffer: ArrayBuffer): Promise<ProjectSnapshot> {
const midi = new Midi(buffer)
const tempo = midi.header.tempos[0]?.bpm ?? 120
const timeSignature = (midi.header.timeSignatures[0]?.timeSignature as TimeSignature | undefined) ?? DEFAULT_SIGNATURE
// Merge notes from all tracks and sort by ticks then by midi (for stable ordering)
const allNotes = midi.tracks
.flatMap(t => t.notes)
.sort((a, b) => a.ticks - b.ticks || a.midi - b.midi)
// Get lyrics from header.meta and sort by ticks
const lyricEvents = midi.header.meta
.filter((event) => event.type === 'lyrics')
.sort((a, b) => a.ticks - b.ticks)
// Match lyrics to notes by tick position
// Each lyric should be consumed by exactly one note at the same tick
const lyricsByTick = new Map<number, string[]>()
for (const event of lyricEvents) {
const existing = lyricsByTick.get(event.ticks) || []
existing.push(decodeUtf8ByteString(event.text))
lyricsByTick.set(event.ticks, existing)
}
// Track which lyrics have been used at each tick position
const usedLyricIndices = new Map<number, number>()
const notes: NoteEvent[] = allNotes.map((note, index) => {
const beat = note.ticks / midi.header.ppq
const durationBeats = note.durationTicks / midi.header.ppq
let lyric = ''
// First try exact tick match
const lyricsAtTick = lyricsByTick.get(note.ticks)
if (lyricsAtTick && lyricsAtTick.length > 0) {
const usedIndex = usedLyricIndices.get(note.ticks) || 0
if (usedIndex < lyricsAtTick.length) {
lyric = lyricsAtTick[usedIndex]
usedLyricIndices.set(note.ticks, usedIndex + 1)
}
}
// If no exact match, try nearby ticks (within small tolerance)
if (!lyric) {
const tolerance = midi.header.ppq / 100 // Very small tolerance
for (const [tick, lyrics] of lyricsByTick.entries()) {
if (Math.abs(tick - note.ticks) <= tolerance) {
const usedIndex = usedLyricIndices.get(tick) || 0
if (usedIndex < lyrics.length) {
lyric = lyrics[usedIndex]
usedLyricIndices.set(tick, usedIndex + 1)
break
}
}
}
}
return {
id: `${index}-${note.midi}-${Math.round(note.ticks)}`,
midi: note.midi,
start: beat,
duration: Math.max(durationBeats, 0.0625),
velocity: note.velocity,
lyric,
}
})
return { tempo, timeSignature, notes, ppq: midi.header.ppq }
}
// Used to add absoluteTime property for sorting
type WithAbsoluteTime<T> = T & { absoluteTime: number }
export function exportMidi(snapshot: ProjectSnapshot): Blob {
const ppq = snapshot.ppq ?? 480 // Use original ppq if available, otherwise default to 480
const microsecondsPerBeat = Math.round(60000000 / snapshot.tempo) // Convert BPM to microseconds per beat
// Sort notes by start time, then by midi for stable ordering
const sortedNotes = [...snapshot.notes].sort((a, b) => a.start - b.start || a.midi - b.midi)
// Build events for a single track containing both lyrics and notes
// Event order at same tick: note_off (0) < lyrics (1) < note_on (2)
// This matches meta.py's tg2midi implementation
const events: Array<WithAbsoluteTime<MidiEvent>> = []
// Add all note events and their corresponding lyrics
sortedNotes.forEach((note) => {
const startTicks = Math.round(note.start * ppq)
const endTicks = Math.round((note.start + note.duration) * ppq)
const velocity = Math.round(note.velocity * 127)
// Add lyric event at the same tick as note_on (but will be sorted before it)
const lyricText = note.lyric ?? ''
const encodedLyric = encodeUtf8ByteString(lyricText)
// Lyric event - sort key 1 (after note_off, before note_on)
events.push({
absoluteTime: startTicks,
deltaTime: 0,
meta: true,
type: 'lyrics',
text: encodedLyric,
_sortKey: 1,
} as WithAbsoluteTime<MidiEvent> & { _sortKey: number })
// Note on event - sort key 2 (after lyrics)
events.push({
absoluteTime: startTicks,
deltaTime: 0,
type: 'noteOn',
channel: 0,
noteNumber: note.midi,
velocity: velocity,
_sortKey: 2,
} as WithAbsoluteTime<MidiEvent> & { _sortKey: number })
// Note off event - sort key 0 (before everything at same tick)
events.push({
absoluteTime: endTicks,
deltaTime: 0,
type: 'noteOff',
channel: 0,
noteNumber: note.midi,
velocity: 0,
_sortKey: 0,
} as WithAbsoluteTime<MidiEvent> & { _sortKey: number })
})
// Sort events by absoluteTime, then by _sortKey
events.sort((a, b) => {
const aKey = (a as { _sortKey?: number })._sortKey ?? 1
const bKey = (b as { _sortKey?: number })._sortKey ?? 1
return a.absoluteTime - b.absoluteTime || aKey - bKey
})
// Convert absolute time to delta time
let lastTick = 0
events.forEach(event => {
event.deltaTime = event.absoluteTime - lastTick
lastTick = event.absoluteTime
delete (event as { absoluteTime?: number }).absoluteTime
delete (event as { _sortKey?: number })._sortKey
})
// Build the MIDI track with header events
const track: MidiEvent[] = [
// Set tempo
{
deltaTime: 0,
meta: true,
type: 'setTempo',
microsecondsPerBeat: microsecondsPerBeat,
},
// Time signature
{
deltaTime: 0,
meta: true,
type: 'timeSignature',
numerator: snapshot.timeSignature[0],
denominator: snapshot.timeSignature[1],
metronome: 24,
thirtyseconds: 8,
},
// All note and lyric events
...events,
// End of track
{
deltaTime: 0,
meta: true,
type: 'endOfTrack',
},
]
// Build MIDI data structure
const midiData: MidiData = {
header: {
format: 0, // Single track format (type 0)
numTracks: 1,
ticksPerBeat: ppq,
},
tracks: [track],
}
const bytes = writeMidi(midiData)
return new Blob([new Uint8Array(bytes)], { type: 'audio/midi' })
}
+10
View File
@@ -0,0 +1,10 @@
import { StrictMode } from 'react'
import { createRoot } from 'react-dom/client'
import './index.css'
import App from './App.tsx'
createRoot(document.getElementById('root')!).render(
<StrictMode>
<App />
</StrictMode>,
)
@@ -0,0 +1,78 @@
import { nanoid } from 'nanoid'
import { create } from 'zustand'
import type { NoteEvent, TimeSignature } from '../types'
const clamp = (value: number, min: number, max: number) =>
Math.min(Math.max(value, min), max)
export type MidiStore = {
tempo: number
timeSignature: TimeSignature
notes: NoteEvent[]
selectedId: string | null
playhead: number
ppq: number | undefined // Ticks per quarter note (for preserving original MIDI timing)
addNote: (partial?: Partial<NoteEvent>) => NoteEvent
updateNote: (id: string, partial: Partial<NoteEvent>) => void
removeNote: (id: string) => void
setNotes: (notes: NoteEvent[]) => void
setTempo: (tempo: number) => void
setTimeSignature: (sig: TimeSignature) => void
setPpq: (ppq: number | undefined) => void
select: (id: string | null) => void
setLyric: (id: string, lyric: string) => void
setPlayhead: (beat: number) => void
clear: () => void
}
const defaultNotes: NoteEvent[] = [
{ id: nanoid(), midi: 64, start: 0, duration: 1.5, velocity: 0.9, lyric: 'la' },
{ id: nanoid(), midi: 67, start: 1.5, duration: 1.5, velocity: 0.85, lyric: 'na' },
{ id: nanoid(), midi: 69, start: 3, duration: 2, velocity: 0.8, lyric: 'ah' },
]
export const useMidiStore = create<MidiStore>((set) => ({
tempo: 110,
timeSignature: [4, 4],
notes: defaultNotes,
selectedId: null,
playhead: 0,
ppq: undefined,
addNote: (partial = {}) => {
const note: NoteEvent = {
id: nanoid(),
midi: partial.midi ?? 64,
start: partial.start ?? 0,
duration: partial.duration ?? 1,
velocity: clamp(partial.velocity ?? 0.85, 0, 1),
lyric: partial.lyric ?? '',
}
set((state) => ({ notes: [...state.notes, note] }))
return note
},
updateNote: (id, partial) => {
set((state) => ({
notes: state.notes.map((note) =>
note.id === id
? {
...note,
...partial,
duration: Math.max(partial.duration ?? note.duration, 0.0625),
}
: note,
),
}))
},
removeNote: (id) => set((state) => ({ notes: state.notes.filter((n) => n.id !== id) })),
setNotes: (notes) => set(() => ({ notes })),
setTempo: (tempo) => set(() => ({ tempo: clamp(tempo, 30, 240) })),
setTimeSignature: (sig) => set(() => ({ timeSignature: sig })),
setPpq: (ppq) => set(() => ({ ppq })),
select: (id) => set(() => ({ selectedId: id })),
setLyric: (id, lyric) =>
set((state) => ({
notes: state.notes.map((note) => (note.id === id ? { ...note, lyric } : note)),
})),
setPlayhead: (beat) => set(() => ({ playhead: Math.max(beat, 0) })),
clear: () => set(() => ({ notes: [], selectedId: null })),
}))
+17
View File
@@ -0,0 +1,17 @@
export type NoteEvent = {
id: string
midi: number
start: number // in beats
duration: number // in beats
velocity: number
lyric: string
}
export type TimeSignature = [number, number]
export type ProjectSnapshot = {
tempo: number
timeSignature: TimeSignature
notes: NoteEvent[]
ppq?: number // Ticks per quarter note (for preserving original MIDI timing)
}
@@ -0,0 +1,33 @@
/** @type {import('tailwindcss').Config} */
export default {
content: ['./index.html', './src/**/*.{ts,tsx,js,jsx}'],
theme: {
extend: {
fontFamily: {
display: ['"Space Grotesk"', '"IBM Plex Sans"', 'system-ui', 'sans-serif'],
mono: ['"JetBrains Mono"', 'ui-monospace', 'SFMono-Regular', 'monospace'],
},
colors: {
ink: {
50: '#f4f7fb',
100: '#dfe7f5',
200: '#beceec',
300: '#95addf',
400: '#6a87ce',
500: '#4b64bc',
600: '#3b4ea7',
700: '#32418a',
800: '#2c376f',
900: '#262f5c',
},
ember: '#ff7043',
mint: '#48e4c2',
},
boxShadow: {
panel: '0 14px 35px rgba(0, 0, 0, 0.25)',
},
},
},
plugins: [],
}
@@ -0,0 +1,28 @@
{
"compilerOptions": {
"tsBuildInfoFile": "./node_modules/.tmp/tsconfig.app.tsbuildinfo",
"target": "ES2022",
"useDefineForClassFields": true,
"lib": ["ES2022", "DOM", "DOM.Iterable"],
"module": "ESNext",
"types": ["vite/client"],
"skipLibCheck": true,
/* Bundler mode */
"moduleResolution": "bundler",
"allowImportingTsExtensions": true,
"verbatimModuleSyntax": true,
"moduleDetection": "force",
"noEmit": true,
"jsx": "react-jsx",
/* Linting */
"strict": true,
"noUnusedLocals": true,
"noUnusedParameters": true,
"erasableSyntaxOnly": true,
"noFallthroughCasesInSwitch": true,
"noUncheckedSideEffectImports": true
},
"include": ["src"]
}
@@ -0,0 +1,7 @@
{
"files": [],
"references": [
{ "path": "./tsconfig.app.json" },
{ "path": "./tsconfig.node.json" }
]
}
@@ -0,0 +1,26 @@
{
"compilerOptions": {
"tsBuildInfoFile": "./node_modules/.tmp/tsconfig.node.tsbuildinfo",
"target": "ES2023",
"lib": ["ES2023"],
"module": "ESNext",
"types": ["node"],
"skipLibCheck": true,
/* Bundler mode */
"moduleResolution": "bundler",
"allowImportingTsExtensions": true,
"verbatimModuleSyntax": true,
"moduleDetection": "force",
"noEmit": true,
/* Linting */
"strict": true,
"noUnusedLocals": true,
"noUnusedParameters": true,
"erasableSyntaxOnly": true,
"noFallthroughCasesInSwitch": true,
"noUncheckedSideEffectImports": true
},
"include": ["vite.config.ts"]
}
@@ -0,0 +1,7 @@
import { defineConfig } from 'vite'
import react from '@vitejs/plugin-react'
// https://vite.dev/config/
export default defineConfig({
plugins: [react()],
})
+669
View File
@@ -0,0 +1,669 @@
"""
SoulX-Singer MIDI <-> metadata converter.
Converts between SoulX-Singer-style metadata JSON (with note_text, note_dur,
note_pitch, note_type per segment) and standard MIDI files. Uses an internal
Note dataclass (start_s, note_dur, note_text, note_pitch, note_type) as the
intermediate representation.
"""
import os
import json
import shutil
from dataclasses import dataclass
from typing import Any, List, Tuple, Union
import librosa
import mido
from soundfile import write
from .f0_extraction import F0Extractor
from .g2p import g2p_transform
# Audio and segmenting constants (used by _edit_data_to_meta)
SAMPLE_RATE = 44100
DEFAULT_LANGUAGE = "Mandarin"
MAX_GAP_SEC = 5.0 # gap (sec) above which we start a new segment
MAX_SEGMENT_DUR_SUM_SEC = 60.0 # max cumulative note duration per segment (sec)
MIN_GAP_THRESHOLD_SEC = 0.001 # ignore gaps smaller than this
LONG_SILENCE_THRESHOLD_SEC = 0.05 # treat as separate <SP> if gap larger
MAX_LEADING_SP_DUR_SEC = 2.0 # cap leading silence in a segment to this (sec)
DEFAULT_RMVPE_MODEL_PATH = "pretrained_models/SoulX-Singer-Preprocess/rmvpe/rmvpe.pt"
@dataclass
class Note:
"""Single note: text, duration (seconds), pitch (MIDI), type. start_s is absolute start time in seconds (for ordering / MIDI)."""
start_s: float
note_dur: float
note_text: str
note_pitch: int
note_type: int
@property
def end_s(self) -> float:
return self.start_s + self.note_dur
def remove_duplicate_segments(meta_data: List[dict]) -> None:
"""Merge consecutive identical notes (same text, pitch, type) within each segment. Mutates meta_data in place."""
for idx, segment in enumerate(meta_data):
texts = segment["note_text"]
durs = segment["note_dur"]
pitches = segment["note_pitch"]
types = segment["note_type"]
new_texts = []
new_durs = []
new_pitches = []
new_types = []
for i in range(len(texts)):
if i == 0:
new_texts.append(texts[i])
new_durs.append(durs[i])
new_pitches.append(pitches[i])
new_types.append(types[i])
continue
t, d, p, ty = texts[i], durs[i], pitches[i], types[i]
if t == "<SP>" and texts[i - 1] == "<SP>":
new_durs[-1] += d
continue
if t == texts[i - 1] and p == pitches[i - 1] and ty == types[i - 1]:
new_durs[-1] += d
else:
new_texts.append(t)
new_durs.append(d)
new_pitches.append(p)
new_types.append(ty)
meta_data[idx]["note_text"] = new_texts
meta_data[idx]["note_dur"] = new_durs
meta_data[idx]["note_pitch"] = new_pitches
meta_data[idx]["note_type"] = new_types
def meta2notes(meta_path: str) -> List[Note]:
"""Parse SoulX-Singer metadata JSON into a flat list of Note (absolute start_s)."""
with open(meta_path, "r", encoding="utf-8") as f:
segments = json.load(f)
if not isinstance(segments, list):
raise ValueError(f"Metadata must be a list of segments, got {type(segments).__name__}")
if not segments:
raise ValueError("Metadata has no segments.")
notes: List[Note] = []
for seg in segments:
offset_s = seg["time"][0] / 1000
words = [str(x).replace("<AP>", "<SP>") for i, x in enumerate(seg["text"].split())]
word_durs = [float(x) for x in seg["duration"].split()]
pitches = [int(x) for x in seg["note_pitch"].split()]
types = [int(x) if words[i] != "<SP>" else 1 for i, x in enumerate(seg["note_type"].split())]
if len(words) != len(word_durs) or len(word_durs) != len(pitches) or len(pitches) != len(types):
raise ValueError(
f"Length mismatch in segment {seg.get('item_name', '?')}: "
"note_text, note_dur, note_pitch, note_type must have same length"
)
current_s = offset_s
for text, dur, pitch, type_ in zip(words, word_durs, pitches, types):
notes.append(
Note(
start_s=current_s,
note_dur=float(dur),
note_text=str(text),
note_pitch=int(pitch),
note_type=int(type_),
)
)
current_s += float(dur)
return notes
def _append_segment_to_meta(
meta_path_str: str,
cut_wavs_output_dir: str,
vocal_file: str,
audio_data: Any,
meta_data: List[dict],
note_start: List[float],
note_end: List[float],
note_text: List[Any],
note_pitch: List[Any],
note_type: List[Any],
note_dur: List[float],
end_time_ms_override: float | None = None,
) -> None:
"""Write one segment wav and append one segment dict to meta_data. Caller clears note_* lists after."""
base_name = os.path.splitext(os.path.basename(meta_path_str))[0]
item_name = f"{base_name}_{len(meta_data)}"
wav_fn = os.path.join(cut_wavs_output_dir, f"{item_name}.wav")
start_ms = int(note_start[0] * 1000)
end_ms = (
int(end_time_ms_override)
if end_time_ms_override is not None
else int(note_end[-1] * 1000)
)
start_sample = int(note_start[0] * SAMPLE_RATE)
end_sample = int(note_end[-1] * SAMPLE_RATE)
write(wav_fn, audio_data[start_sample:end_sample], SAMPLE_RATE)
meta_data.append({
"item_name": item_name,
"wav_fn": wav_fn,
"origin_wav_fn": vocal_file,
"start_time_ms": start_ms,
"end_time_ms": end_ms,
"language": DEFAULT_LANGUAGE,
"note_text": list(note_text),
"note_pitch": list(note_pitch),
"note_type": list(note_type),
"note_dur": list(note_dur),
})
def convert_meta(meta_data: List[dict], rmvpe_model_path, device="cuda"):
pitch_extractor = F0Extractor(rmvpe_model_path, device=device, verbose=False)
converted_data = []
for item in meta_data:
wav_fn = item.get("wav_fn")
if not wav_fn or not os.path.isfile(wav_fn):
raise FileNotFoundError(f"Segment wav file not found: {wav_fn}")
f0 = pitch_extractor.process(wav_fn)
converted_item = {
"index": item.get("item_name"),
"language": item.get("language"),
"time": [item.get("start_time_ms", 0), item.get("end_time_ms", sum(item["note_dur"]) * 1000)],
"duration": " ".join(str(round(x, 2)) for x in item.get("note_dur", [])),
"text": " ".join(item.get("note_text", [])),
"phoneme": " ".join(g2p_transform(item.get("note_text", []), DEFAULT_LANGUAGE)),
"note_pitch": " ".join(str(x) for x in item.get("note_pitch", [])),
"note_type": " ".join(str(x) for x in item.get("note_type", [])),
"f0": " ".join(str(round(float(x), 1)) for x in f0),
}
converted_data.append(converted_item)
return converted_data
def _edit_data_to_meta(
meta_path_str: str,
edit_data: List[dict],
vocal_file: str,
rmvpe_model_path: str | None = None,
device: str = "cuda",
) -> None:
"""Write SoulX-Singer metadata JSON from edit_data (list of {start, end, note_text, note_pitch, note_type})."""
# Use a fixed temporary directory for cut wavs
cut_wavs_output_dir = os.path.join(os.path.dirname(vocal_file), "cut_wavs_tmp")
os.makedirs(cut_wavs_output_dir, exist_ok=True)
note_text: List[Any] = []
note_pitch: List[Any] = []
note_type: List[Any] = []
note_dur: List[float] = []
note_start: List[float] = []
note_end: List[float] = []
prev_end = 0.0
meta_data: List[dict] = []
audio_data, _ = librosa.load(vocal_file, sr=SAMPLE_RATE, mono=True)
dur_sum = 0.0
for entry in edit_data:
start = float(entry["start"])
end = float(entry["end"])
text = entry["note_text"]
pitch = entry["note_pitch"]
type_ = entry["note_type"]
if text == "" or pitch == "" or type_ == "":
note_text.append("<SP>")
note_pitch.append(0)
note_type.append(1)
note_dur.append(end - start)
note_start.append(start)
note_end.append(end)
prev_end = end
dur_sum += end - start
continue
if (
len(note_text) > 0
and note_text[-1] == "<SP>"
and note_dur[-1] > MAX_LEADING_SP_DUR_SEC
):
cut_time = note_dur[-1] - MAX_LEADING_SP_DUR_SEC
note_dur[-1] = MAX_LEADING_SP_DUR_SEC
end_ms_override = note_end[-1] * 1000 - cut_time * 1000
_append_segment_to_meta(
meta_path_str,
cut_wavs_output_dir,
vocal_file,
audio_data,
meta_data,
note_start,
note_end,
note_text,
note_pitch,
note_type,
note_dur,
end_time_ms_override=end_ms_override,
)
note_text = []
note_pitch = []
note_type = []
note_dur = []
note_start = []
note_end = []
prev_end = start
dur_sum = 0.0
gap_from_prev = start - prev_end
gap_from_last_note = (start - note_end[-1]) if note_end else 0.0
if (
gap_from_prev >= MAX_GAP_SEC
or gap_from_last_note >= MAX_GAP_SEC
or dur_sum >= MAX_SEGMENT_DUR_SUM_SEC
):
if len(note_text) > 0:
_append_segment_to_meta(
meta_path_str,
cut_wavs_output_dir,
vocal_file,
audio_data,
meta_data,
note_start,
note_end,
note_text,
note_pitch,
note_type,
note_dur,
)
note_text = []
note_pitch = []
note_type = []
note_dur = []
note_start = []
note_end = []
prev_end = start
dur_sum = 0.0
if start - prev_end > MIN_GAP_THRESHOLD_SEC:
if start - prev_end > LONG_SILENCE_THRESHOLD_SEC or len(note_text) == 0:
note_text.append("<SP>")
note_pitch.append(0)
note_type.append(1)
note_dur.append(start - prev_end)
note_start.append(prev_end)
note_end.append(start)
else:
if len(note_dur) > 0:
note_dur[-1] += start - prev_end
note_end[-1] = start
prev_end = end
note_text.append(text)
note_pitch.append(int(pitch))
note_type.append(int(type_))
note_dur.append(end - start)
note_start.append(start)
note_end.append(end)
dur_sum += end - start
if len(note_text) > 0:
_append_segment_to_meta(
meta_path_str,
cut_wavs_output_dir,
vocal_file,
audio_data,
meta_data,
note_start,
note_end,
note_text,
note_pitch,
note_type,
note_dur,
)
remove_duplicate_segments(meta_data)
_rmvpe_path = rmvpe_model_path or DEFAULT_RMVPE_MODEL_PATH
converted_data = convert_meta(meta_data, _rmvpe_path, device)
with open(meta_path_str, "w", encoding="utf-8") as f:
json.dump(converted_data, f, ensure_ascii=False, indent=2)
# Clean up temporary cut wavs directory
try:
shutil.rmtree(cut_wavs_output_dir, ignore_errors=True)
except Exception:
pass
def notes2meta(
notes: List[Note],
meta_path: str,
vocal_file: str,
rmvpe_model_path: str | None = None,
device: str = "cuda",
) -> None:
"""Write SoulX-Singer metadata JSON from a list of Note (segmenting + wav cuts)."""
edit_data = [
{
"start": n.start_s,
"end": n.end_s,
"note_text": n.note_text,
"note_pitch": str(n.note_pitch),
"note_type": str(n.note_type),
}
for n in notes
]
_edit_data_to_meta(
str(meta_path),
edit_data,
vocal_file,
rmvpe_model_path=rmvpe_model_path,
device=device,
)
@dataclass(frozen=True)
class MidiDefaults:
ticks_per_beat: int = 500
tempo: int = 500000 # microseconds per beat (120 BPM)
time_signature: Tuple[int, int] = (4, 4)
velocity: int = 64
def _seconds_to_ticks(seconds: float, ticks_per_beat: int, tempo: int) -> int:
return int(round(seconds * ticks_per_beat * 1_000_000 / tempo))
def notes2midi(
notes: List[Note],
midi_path: str,
defaults: MidiDefaults | None = None,
) -> None:
"""Write MIDI file from a list of Note."""
defaults = defaults or MidiDefaults()
if not notes:
raise ValueError("Empty note list.")
events: List[Tuple[int, int, Union[mido.Message, mido.MetaMessage]]] = []
for n in notes:
start_s = n.start_s
end_s = n.end_s
if end_s <= start_s:
continue
start_ticks = _seconds_to_ticks(
start_s, defaults.ticks_per_beat, defaults.tempo
)
end_ticks = _seconds_to_ticks(
end_s, defaults.ticks_per_beat, defaults.tempo
)
if end_ticks <= start_ticks:
end_ticks = start_ticks + 1
lyric = n.note_text
try:
lyric = lyric.encode("utf-8").decode("latin1")
except (UnicodeEncodeError, UnicodeDecodeError):
pass
if n.note_type == 3:
lyric = "-"
events.append(
(start_ticks, 1, mido.MetaMessage("lyrics", text=lyric, time=0))
)
events.append(
(
start_ticks,
2,
mido.Message(
"note_on",
note=n.note_pitch,
velocity=defaults.velocity,
time=0,
),
)
)
events.append(
(
end_ticks,
0,
mido.Message("note_off", note=n.note_pitch, velocity=0, time=0),
)
)
events.sort(key=lambda x: (x[0], x[1]))
mid = mido.MidiFile(ticks_per_beat=defaults.ticks_per_beat)
track = mido.MidiTrack()
mid.tracks.append(track)
track.append(mido.MetaMessage("set_tempo", tempo=defaults.tempo, time=0))
track.append(
mido.MetaMessage(
"time_signature",
numerator=defaults.time_signature[0],
denominator=defaults.time_signature[1],
time=0,
)
)
last_tick = 0
for tick, _, msg in events:
msg.time = max(0, tick - last_tick)
track.append(msg)
last_tick = tick
track.append(mido.MetaMessage("end_of_track", time=0))
mid.save(midi_path)
def midi2notes(midi_path: str) -> List[Note]:
"""Parse MIDI file into a list of Note. Merges all tracks; tempo from last set_tempo event."""
mid = mido.MidiFile(midi_path)
ticks_per_beat = mid.ticks_per_beat
tempo = 500000
raw_notes: List[dict] = []
lyrics: List[Tuple[int, str]] = []
for track in mid.tracks:
abs_ticks = 0
active = {}
for msg in track:
abs_ticks += msg.time
if msg.type == "set_tempo":
tempo = msg.tempo
elif msg.type == "lyrics":
text = msg.text
try:
text = text.encode("latin1").decode("utf-8")
except Exception:
pass
lyrics.append((abs_ticks, text))
elif msg.type == "note_on":
key = (msg.channel, msg.note)
if msg.velocity > 0:
active[key] = (abs_ticks, msg.velocity)
else:
if key in active:
start_ticks, vel = active.pop(key)
raw_notes.append(
{
"midi": msg.note,
"start_ticks": start_ticks,
"duration_ticks": abs_ticks - start_ticks,
"velocity": vel,
"lyric": "",
}
)
elif msg.type == "note_off":
key = (msg.channel, msg.note)
if key in active:
start_ticks, vel = active.pop(key)
raw_notes.append(
{
"midi": msg.note,
"start_ticks": start_ticks,
"duration_ticks": abs_ticks - start_ticks,
"velocity": vel,
"lyric": "",
}
)
if not raw_notes:
raise ValueError("No notes found in MIDI file")
for n in raw_notes:
n["end_ticks"] = n["start_ticks"] + n["duration_ticks"]
raw_notes.sort(key=lambda n: n["start_ticks"])
lyrics.sort(key=lambda x: x[0])
trimmed = []
for note in raw_notes:
while trimmed:
prev = trimmed[-1]
if note["start_ticks"] < prev["end_ticks"]:
prev["end_ticks"] = note["start_ticks"]
prev["duration_ticks"] = prev["end_ticks"] - prev["start_ticks"]
if prev["duration_ticks"] <= 0:
trimmed.pop()
continue
break
trimmed.append(note)
raw_notes = trimmed
tolerance = ticks_per_beat // 100
lyric_idx = 0
for note in raw_notes:
while lyric_idx < len(lyrics) and lyrics[lyric_idx][0] < note["start_ticks"] - tolerance:
lyric_idx += 1
if lyric_idx < len(lyrics):
lyric_ticks, lyric_text = lyrics[lyric_idx]
if abs(lyric_ticks - note["start_ticks"]) <= tolerance:
note["lyric"] = lyric_text
lyric_idx += 1
def ticks_to_seconds(ticks: int) -> float:
return (ticks / ticks_per_beat) * (tempo / 1_000_000)
result: List[Note] = []
prev_end_s = 0.0
for idx, n in enumerate(raw_notes):
start_s = ticks_to_seconds(n["start_ticks"])
end_s = ticks_to_seconds(n["end_ticks"])
if prev_end_s > start_s:
start_s = prev_end_s
dur_s = end_s - start_s
if dur_s <= 0:
continue
lyric = n.get("lyric", "")
if not lyric:
tp = 2
text = "啦"
elif lyric == "<SP>":
tp = 1
text = "<SP>"
elif lyric == "-":
tp = 3
text = raw_notes[idx - 1].get("lyric", "-") if idx > 0 else "-"
else:
tp = 2
text = lyric
result.append(
Note(
start_s=start_s,
note_dur=dur_s,
note_text=text,
note_pitch=n["midi"],
note_type=tp,
)
)
prev_end_s = end_s
return result
def meta2midi(meta_path: str, midi_path: str, defaults: MidiDefaults | None = None) -> None:
"""Convert SoulX-Singer metadata JSON to MIDI file (meta -> List[Note] -> midi)."""
notes = meta2notes(meta_path)
notes2midi(notes, midi_path, defaults)
print(f"Saved MIDI to {midi_path}")
def midi2meta(
midi_path: str,
meta_path: str,
vocal_file: str,
rmvpe_model_path: str | None = None,
device: str = "cuda",
) -> None:
"""Convert MIDI file to SoulX-Singer metadata JSON (midi -> List[Note] -> meta)."""
meta_dir = os.path.dirname(meta_path)
if meta_dir:
os.makedirs(meta_dir, exist_ok=True)
# cut_wavs will be written to a fixed temporary directory inside _edit_data_to_meta
notes = midi2notes(midi_path)
notes2meta(
notes,
meta_path,
vocal_file,
rmvpe_model_path=rmvpe_model_path,
device=device,
)
print(f"Saved Meta to {meta_path}")
if __name__ == "__main__":
import argparse
parser = argparse.ArgumentParser(
description="Convert SoulX-Singer metadata JSON <-> MIDI."
)
parser.add_argument("--meta", type=str, help="Path to metadata JSON")
parser.add_argument("--midi", type=str, help="Path to MIDI file")
parser.add_argument("--vocal", type=str, help="Path to vocal wav (for midi2meta)")
parser.add_argument(
"--meta2midi",
action="store_true",
help="Convert meta -> midi (requires --meta and --midi)",
)
parser.add_argument(
"--midi2meta",
action="store_true",
help="Convert midi -> meta (requires --midi, --meta, --vocal, --cut_wavs_dir)",
)
parser.add_argument(
"--rmvpe_model_path",
type=str,
help="Path to RMVPE model",
default="pretrained_models/SoulX-Singer-Preprocess/rmvpe/rmvpe.pt",
)
parser.add_argument(
"--device",
type=str,
help="Device to use for RMVPE",
default="cuda",
)
args = parser.parse_args()
if args.meta2midi:
if not args.meta or not args.midi:
parser.error("--meta2midi requires --meta and --midi")
meta2midi(args.meta, args.midi)
elif args.midi2meta:
if not args.midi or not args.meta or not args.vocal:
parser.error(
"--midi2meta requires --midi, --meta, --vocal"
)
midi2meta(
args.midi,
args.meta,
args.vocal,
rmvpe_model_path=args.rmvpe_model_path,
device=args.device,
)
else:
parser.print_help()
@@ -0,0 +1,522 @@
# https://github.com/RickyL-2000/ROSVOT
import math
import sys
import traceback
import json
import time
from pathlib import Path
from typing import Any, Dict, Optional
import librosa
import numpy as np
import torch
import matplotlib.pyplot as plt
from .utils.os_utils import safe_path
from .utils.commons.hparams import set_hparams
from .utils.commons.ckpt_utils import load_ckpt
from .utils.commons.dataset_utils import pad_or_cut_xd
from .utils.audio.mel import MelNet
from .utils.audio.pitch_utils import (
norm_interp_f0,
denorm_f0,
f0_to_coarse,
boundary2Interval,
save_midi,
midi_to_hz,
)
from .utils.rosvot_utils import (
get_mel_len,
align_word,
regulate_real_note_itv,
regulate_ill_slur,
bd_to_durs,
)
from .modules.pe.rmvpe import RMVPE
from .modules.rosvot.rosvot import MidiExtractor, WordbdExtractor
@torch.no_grad()
def infer_sample(
item: Dict[str, Any],
hparams: Dict[str, Any],
models: Dict[str, Any],
device: torch.device,
*,
save_dir: Optional[str] = None,
apply_rwbd: Optional[bool] = None,
# outputs
save_plot: bool = False,
no_save_midi: bool = True,
no_save_npy: bool = True,
verbose: bool = False,
) -> Dict[str, Any]:
if "item_name" not in item or "wav_fn" not in item:
raise ValueError('item must contain keys: "item_name" and "wav_fn"')
item_name = item["item_name"]
wav_src = item["wav_fn"]
# Decide RWBD usage
if apply_rwbd is None:
apply_rwbd_ = ("word_durs" not in item)
else:
apply_rwbd_ = bool(apply_rwbd)
# Models
model = models["model"]
mel_net = models["mel_net"]
pe = models.get("pe")
wbd_predictor = models.get("wbd_predictor")
if wbd_predictor is None and apply_rwbd_:
raise ValueError("apply_rwbd is True but wbd_predictor model is not provided in models")
# ---- Prepare Data ----
if isinstance(wav_src, str):
wav, _ = librosa.core.load(wav_src, sr=hparams["audio_sample_rate"])
else:
wav = wav_src
if not isinstance(wav, np.ndarray):
wav = np.asarray(wav)
wav = wav.astype(np.float32)
# Calculate timestamps and alignment lengths
wav_len_samples = wav.shape[-1]
mel_len = get_mel_len(wav_len_samples, hparams["hop_size"])
# Word boundary preparation
mel2word = None
word_durs_filtered = None
if not apply_rwbd_:
if "word_durs" not in item:
raise ValueError('apply_rwbd=False but item has no "word_durs"')
wd_raw = list(item["word_durs"])
min_word_dur = hparams.get("min_word_dur", 20) / 1000
word_durs_filtered = []
for i, wd in enumerate(wd_raw):
if wd < min_word_dur:
if i == 0 and len(wd_raw) > 1:
wd_raw[i + 1] += wd
elif len(word_durs_filtered) > 0:
word_durs_filtered[-1] += wd
else:
word_durs_filtered.append(wd)
mel2word, _ = align_word(word_durs_filtered, mel_len, hparams["hop_size"], hparams["audio_sample_rate"])
mel2word = np.asarray(mel2word)
if mel2word.size > 0 and mel2word[0] == 0:
mel2word = mel2word + 1
mel2word_len = int(np.sum(mel2word > 0))
real_len = min(mel_len, mel2word_len)
else:
real_len = min(mel_len, hparams["max_frames"])
T = math.ceil(min(real_len, hparams["max_frames"]) / hparams["frames_multiple"]) * hparams["frames_multiple"]
# ---- Input Tensors & Padding ----
target_samples = T * hparams["hop_size"]
wav_t = torch.from_numpy(wav).float().to(device).unsqueeze(0) # [1, L]
if wav_t.shape[-1] < target_samples:
wav_t = pad_or_cut_xd(wav_t, target_samples, 1)
# ---- Pitch Extraction ----
if pe is not None:
f0s, uvs = pe.get_pitch_batch(
wav_t,
sample_rate=hparams["audio_sample_rate"],
hop_size=hparams["hop_size"],
lengths=[real_len],
fmax=hparams["f0_max"],
fmin=hparams["f0_min"],
)
f0_1d, uv_1d = norm_interp_f0(f0s[0][:T])
f0_t = pad_or_cut_xd(torch.FloatTensor(f0_1d).to(device), T, 0).unsqueeze(0)
uv_t = pad_or_cut_xd(torch.FloatTensor(uv_1d).to(device), T, 0).long().unsqueeze(0)
pitch_coarse = f0_to_coarse(denorm_f0(f0_t, uv_t)).to(device)
f0_np = denorm_f0(f0_t, uv_t)[0].detach().cpu().numpy()[:real_len]
else:
f0_t = uv_t = pitch_coarse = None
f0_np = None
# ---- Mel Extraction ----
mel = mel_net(wav_t) # [1, T_padded, C]
mel = pad_or_cut_xd(mel, T, 1)
# Construct non-padding mask
mel_nonpadding_mask = torch.zeros(1, T, device=device)
mel_nonpadding_mask[:, :real_len] = 1.0
# Apply mask to mel (zero out padding)
mel = (mel.transpose(1, 2) * mel_nonpadding_mask.unsqueeze(1)).transpose(1, 2)
# Re-calculate non_padding bool mask
mel_nonpadding = mel.abs().sum(-1) > 0
# ---- Word Boundary ----
word_durs_used = None
if apply_rwbd_:
mel_input = mel[:, :, : hparams.get("wbd_use_mel_bins", 80)]
wbd_outputs = wbd_predictor(
mel=mel_input,
pitch=pitch_coarse,
uv=uv_t,
non_padding=mel_nonpadding,
train=False,
)
word_bd = wbd_outputs["word_bd_pred"] # [1, T]
else:
# Construct word_bd from provided durs
mel2word_t = pad_or_cut_xd(torch.LongTensor(mel2word).to(device), T, 0)
word_bd = torch.zeros_like(mel2word_t)
# Vectorized check
word_bd[1:] = (mel2word_t[1:] != mel2word_t[:-1]).long()
word_bd[real_len:] = 0
word_bd = word_bd.unsqueeze(0) # [1, T]
word_durs_used = np.array(word_durs_filtered)
# ---- Main Inference ----
mel_input = mel[:, :, : hparams.get("use_mel_bins", 80)]
outputs = model(
mel=mel_input,
word_bd=word_bd,
pitch=pitch_coarse,
uv=uv_t,
non_padding=mel_nonpadding,
train=False,
)
note_lengths = outputs["note_lengths"].detach().cpu().numpy()
note_bd_pred = outputs["note_bd_pred"][0].detach().cpu().numpy()[:real_len]
note_pred = outputs["note_pred"][0].detach().cpu().numpy()[: note_lengths[0]]
note_bd_logits = torch.sigmoid(outputs["note_bd_logits"])[0].detach().cpu().numpy()[:real_len]
if note_pred.shape == (0,):
if verbose:
print(f"skip {item_name}: no notes detected")
return {
"item_name": item_name,
"pitches": [],
"note_durs": [],
"note2words": None,
}
# ---- Post-Processing & Regulation ----
note_itv_pred = boundary2Interval(note_bd_pred)
note2words = None
if apply_rwbd_:
word_bd_np = outputs['word_bd_pred'][0].detach().cpu().numpy()[:real_len]
word_durs_derived = np.array(bd_to_durs(word_bd_np)) * hparams['hop_size'] / hparams['audio_sample_rate']
word_durs_for_reg = word_durs_derived
word_bd_for_reg = word_bd_np
else:
word_bd_for_reg = word_bd[0].detach().cpu().numpy()[:real_len]
word_durs_for_reg = word_durs_used
should_regulate = hparams.get("infer_regulate_real_note_itv", True) and (not apply_rwbd_)
if should_regulate and (word_durs_for_reg is not None):
try:
note_itv_pred_secs, note2words = regulate_real_note_itv(
note_itv_pred,
note_bd_pred,
word_bd_for_reg,
word_durs_for_reg,
hparams["hop_size"],
hparams["audio_sample_rate"],
)
note_pred, note_itv_pred_secs, note2words = regulate_ill_slur(note_pred, note_itv_pred_secs, note2words)
except Exception as err:
if verbose:
_, exc_value, exc_tb = sys.exc_info()
tb = traceback.extract_tb(exc_tb)[-1]
print(f"postprocess failed: {err}: {exc_value} in {tb[0]}:{tb[1]} '{tb[2]}' in {tb[3]}")
# Fallback
note_itv_pred_secs = note_itv_pred * hparams["hop_size"] / hparams["audio_sample_rate"]
note2words = None
else:
note_itv_pred_secs = note_itv_pred * hparams["hop_size"] / hparams["audio_sample_rate"]
# ---- Output ----
note_durs = [float((itv[1] - itv[0])) for itv in note_itv_pred_secs]
out = {
"item_name": item_name,
"pitches": note_pred.tolist(),
"note_durs": note_durs,
"note2words": note2words.tolist() if note2words is not None else None,
}
# ---- Saving ----
if save_dir is not None:
save_dir_path = Path(save_dir)
save_dir_path.mkdir(parents=True, exist_ok=True)
fn = str(item_name)
if not no_save_midi:
save_midi(note_pred, note_itv_pred_secs, safe_path(save_dir_path / "midi" / f"{fn}.mid"))
if not no_save_npy:
np.save(safe_path(save_dir_path / "npy" / f"[note]{fn}.npy"), out, allow_pickle=True)
if save_plot:
fig = plt.figure()
if f0_np is not None:
plt.plot(f0_np, color="red", label="f0")
midi_pred = np.zeros(note_bd_pred.shape[0], dtype=np.float32)
itvs = np.round(note_itv_pred_secs * hparams["audio_sample_rate"] / hparams["hop_size"]).astype(int)
for i, itv in enumerate(itvs):
midi_pred[itv[0] : itv[1]] = note_pred[i]
plt.plot(midi_to_hz(midi_pred), color="blue", label="pred midi")
plt.plot(note_bd_logits * 100, color="green", label="note bd logits x100")
plt.legend()
plt.tight_layout()
plt.savefig(safe_path(save_dir_path / "plot" / f"[MIDI]{fn}.png"), format="png")
plt.close(fig)
return out
def load_rosvot_models(ckpt, config="", wbd_ckpt="", wbd_config="", device="cuda:0", verbose=False, thr=0.85):
"""
Load models once to reuse across multiple items.
"""
dev = torch.device(device)
# 1. Hparams
config_path = Path(ckpt).with_name("config.yaml") if config == "" else config
pe_ckpt = Path(ckpt).parent.parent / "rmvpe/model.pt"
hparams = set_hparams(
config=config_path,
print_hparams=verbose,
hparams_str=f"note_bd_threshold={thr}",
)
# 2. Main Model
model = MidiExtractor(hparams)
load_ckpt(model, ckpt, verbose=verbose)
model.eval().to(dev)
# 3. MelNet
mel_net = MelNet(hparams)
mel_net.to(dev)
# 4. Pitch Extractor
pe = None
if hparams.get("use_pitch_embed", False):
pe = RMVPE(pe_ckpt, device=dev)
# 5. Word Boundary Predictor (optional but we load if ckpt provided or needed)
wbd_predictor = None
if wbd_ckpt:
wbd_config_path = Path(wbd_ckpt).with_name("config.yaml") if wbd_config == "" else wbd_config
wbd_hparams = set_hparams(
config=wbd_config_path,
print_hparams=False,
hparams_str="",
)
hparams.update({
"wbd_use_mel_bins": wbd_hparams["use_mel_bins"],
"min_word_dur": wbd_hparams["min_word_dur"],
})
wbd_predictor = WordbdExtractor(wbd_hparams)
load_ckpt(wbd_predictor, wbd_ckpt, verbose=verbose)
wbd_predictor.eval().to(dev)
models = {
"model": model,
"mel_net": mel_net,
"pe": pe,
"wbd_predictor": wbd_predictor
}
return hparams, models
class NoteTranscriber:
"""Note transcription wrapper based on ROSVOT.
Loads ROSVOT and optional RWBD models once in ``__init__`` and
exposes a :py:meth:`process` API that turns an item dict into
aligned note metadata for downstream SVS.
"""
def __init__(
self,
rosvot_model_path: str,
rwbd_model_path: str,
*,
rosvot_config_path: str = "",
rwbd_config_path: str = "",
device: str = "cuda:0",
thr: float = 0.85,
verbose: bool = True,
):
"""Initialize the note transcriber.
Args:
ckpt: Path to the main ROSVOT checkpoint.
config: Optional config YAML path for ROSVOT.
wbd_ckpt: Optional word-boundary checkpoint path.
wbd_config: Optional config YAML path for RWBD.
device: Torch device string, e.g. ``"cuda:0"`` / ``"cpu"``.
thr: Note boundary threshold.
verbose: Whether to print verbose logs.
"""
self.verbose = verbose
self.device = torch.device(device)
self.hparams, self.models = load_rosvot_models(
ckpt=rosvot_model_path,
config=rosvot_config_path,
wbd_ckpt=rwbd_model_path,
wbd_config=rwbd_config_path,
device=device,
verbose=verbose,
thr=thr,
)
if self.verbose:
print(
"[note transcription] init success:",
f"device={self.device}",
f"rosvot_model_path={rosvot_model_path}",
f"rwbd_model_path={rwbd_model_path if rwbd_model_path else 'None'}",
f"thr={thr}",
)
def process(
self,
item: Dict[str, Any],
*,
segment_info: Optional[Dict[str, Any]] = None,
save_dir: Optional[str] = None,
apply_rwbd: Optional[bool] = None,
save_plot: bool = False,
no_save_midi: bool = True,
no_save_npy: bool = True,
verbose: Optional[bool] = None,
) -> Dict[str, Any]:
"""Run ROSVOT on a single item and post-process outputs.
Args:
item: Input metadata dict with at least ``item_name`` and ``wav_fn``.
segment_info: Optional segment metadata for sliced audio.
save_dir: Optional directory for debug artifacts (plots, midis).
apply_rwbd: Whether to run RWBD-based word boundary refinement.
save_plot: Whether to save diagnostic plots.
no_save_midi: If True, skip saving midi.
no_save_npy: If True, skip saving numpy intermediates.
verbose: Override instance-level verbose flag for this call.
Returns:
Dict with aligned note information for downstream SVS.
"""
v = self.verbose if verbose is None else verbose
if v:
item_name = item.get("item_name", "")
wav_fn = item.get("wav_fn", "")
print(f"[note transcription] process: start: item_name={item_name} wav_fn={wav_fn}")
t0 = time.time()
rosvot_out = infer_sample(
item,
self.hparams,
self.models,
device=self.device,
save_dir=save_dir,
apply_rwbd=apply_rwbd,
save_plot=save_plot,
no_save_midi=no_save_midi,
no_save_npy=no_save_npy,
verbose=v,
)
out = self.post_process(
metadata=item,
segment_info=segment_info,
rosvot_out=rosvot_out,
)
if v:
dt = time.time() - t0
print(
"[note transcription] process: done:",
f"item_name={out.get('item_name','')}",
f"n_notes={len(out.get('note_pitch', []) or [])}",
f"time={dt:.3f}s",
)
return out
@staticmethod
def _normalize_note2words(note2words: list[int]) -> list[int]:
if not note2words:
return []
normalized = [note2words[0]]
for idx in range(1, len(note2words)):
if note2words[idx] < normalized[-1]:
normalized.append(normalized[-1])
else:
normalized.append(note2words[idx])
return normalized
@staticmethod
def _build_ep_types(note2words: list[int], align_words: list[str]) -> list[int]:
ep_types: list[int] = []
prev = -1
for i, w in zip(note2words, align_words):
if w == "<SP>":
ep_types.append(1)
else:
ep_types.append(2 if i != prev else 3)
prev = i
return ep_types
def post_process(
self,
*,
metadata: Dict[str, Any],
segment_info: Dict[str, Any],
rosvot_out: Dict[str, Any],
) -> Dict[str, Any]:
"""Build aligned note metadata using ROSVOT outputs."""
note2words_raw = rosvot_out.get("note2words") or []
note2words = self._normalize_note2words(note2words_raw)
align_words = [
metadata["words"][idx - 1]
for idx in note2words_raw
if 0 < idx <= len(metadata["words"])
]
ep_types = self._build_ep_types(note2words, align_words) if align_words else []
return {
"item_name": rosvot_out.get("item_name", "") if not segment_info else segment_info["item_name"],
"wav_fn": metadata.get("wav_fn", "") if not segment_info else segment_info["wav_fn"],
"origin_wav_fn": metadata.get("origin_wav_fn", "") if not segment_info else segment_info["origin_wav_fn"],
"start_time_ms": "" if not segment_info else segment_info["start_time_ms"],
"end_time_ms": "" if not segment_info else segment_info["end_time_ms"],
"language": metadata.get("language", ""),
"note_text": align_words,
"note_dur": rosvot_out.get("note_durs", []),
"note_type": ep_types,
"note_pitch": rosvot_out.get("pitches", []),
}
if __name__ == "__main__":
items = json.load(open("example/test/rosvot_input.json", "r"))
item = items[0]
m = NoteTranscriber(
rosvot_model_path="pretrained_models/rosvot/rosvot/model.pt",
rwbd_model_path="pretrained_models/rosvot/rwbd/model.pt",
device="cuda"
)
out = m.process(item)
print(out)
@@ -0,0 +1 @@
"""ROSVOT model submodules."""
@@ -0,0 +1 @@
"""Common ROSVOT layers and utilities."""
@@ -0,0 +1 @@
"""Conformer layers for ROSVOT."""
@@ -0,0 +1,96 @@
from torch import nn
from .espnet_positional_embedding import RelPositionalEncoding, ScaledPositionalEncoding, PositionalEncoding
from .espnet_transformer_attn import RelPositionMultiHeadedAttention, MultiHeadedAttention
from .layers import Swish, ConvolutionModule, EncoderLayer, MultiLayeredConv1d
from ..layers import Embedding
class ConformerLayers(nn.Module):
def __init__(self, hidden_size, num_layers, kernel_size=9, dropout=0.0, num_heads=4,
use_last_norm=True, save_hidden=False):
super().__init__()
self.use_last_norm = use_last_norm
self.layers = nn.ModuleList()
positionwise_layer = MultiLayeredConv1d
positionwise_layer_args = (hidden_size, hidden_size * 4, 1, dropout)
self.pos_embed = RelPositionalEncoding(hidden_size, dropout)
self.encoder_layers = nn.ModuleList([EncoderLayer(
hidden_size,
RelPositionMultiHeadedAttention(num_heads, hidden_size, 0.0),
positionwise_layer(*positionwise_layer_args),
positionwise_layer(*positionwise_layer_args),
ConvolutionModule(hidden_size, kernel_size, Swish()),
dropout,
) for _ in range(num_layers)])
if self.use_last_norm:
self.layer_norm = nn.LayerNorm(hidden_size)
else:
self.layer_norm = nn.Linear(hidden_size, hidden_size)
self.save_hidden = save_hidden
if save_hidden:
self.hiddens = []
def forward(self, x, padding_mask=None):
"""
:param x: [B, T, H]
:param padding_mask: [B, T]
:return: [B, T, H]
"""
self.hiddens = []
nonpadding_mask = x.abs().sum(-1) > 0
x = self.pos_embed(x)
for l in self.encoder_layers:
x, mask = l(x, nonpadding_mask[:, None, :])
if self.save_hidden:
self.hiddens.append(x[0])
x = x[0]
x = self.layer_norm(x) * nonpadding_mask.float()[:, :, None]
return x
class FastConformerLayers(ConformerLayers):
def __init__(self, hidden_size, num_layers, kernel_size=9, dropout=0.0, num_heads=4,
use_last_norm=True, save_hidden=False):
super(ConformerLayers, self).__init__()
self.use_last_norm = use_last_norm
self.layers = nn.ModuleList()
positionwise_layer = MultiLayeredConv1d
positionwise_layer_args = (hidden_size, hidden_size * 4, 1, dropout)
self.pos_embed = PositionalEncoding(hidden_size, dropout)
self.encoder_layers = nn.ModuleList([EncoderLayer(
hidden_size,
MultiHeadedAttention(num_heads, hidden_size, 0.0, flash=True),
positionwise_layer(*positionwise_layer_args),
positionwise_layer(*positionwise_layer_args),
ConvolutionModule(hidden_size, kernel_size, Swish()),
dropout,
) for _ in range(num_layers)])
if self.use_last_norm:
self.layer_norm = nn.LayerNorm(hidden_size)
else:
self.layer_norm = nn.Linear(hidden_size, hidden_size)
self.save_hidden = save_hidden
if save_hidden:
self.hiddens = []
class ConformerEncoder(ConformerLayers):
def __init__(self, hidden_size, dict_size, num_layers=None):
conformer_enc_kernel_size = 9
super().__init__(hidden_size, num_layers, conformer_enc_kernel_size)
self.embed = Embedding(dict_size, hidden_size, padding_idx=0)
def forward(self, x):
"""
:param src_tokens: [B, T]
:return: [B x T x C]
"""
x = self.embed(x) # [B, T, H]
x = super(ConformerEncoder, self).forward(x)
return x
class ConformerDecoder(ConformerLayers):
def __init__(self, hidden_size, num_layers):
conformer_dec_kernel_size = 9
super().__init__(hidden_size, num_layers, conformer_dec_kernel_size)
@@ -0,0 +1,113 @@
import math
import torch
class PositionalEncoding(torch.nn.Module):
"""Positional encoding.
Args:
d_model (int): Embedding dimension.
dropout_rate (float): Dropout rate.
max_len (int): Maximum input length.
reverse (bool): Whether to reverse the input position.
"""
def __init__(self, d_model, dropout_rate, max_len=5000, reverse=False):
"""Construct an PositionalEncoding object."""
super(PositionalEncoding, self).__init__()
self.d_model = d_model
self.reverse = reverse
self.xscale = math.sqrt(self.d_model)
self.dropout = torch.nn.Dropout(p=dropout_rate)
self.pe = None
self.extend_pe(torch.tensor(0.0).expand(1, max_len))
def extend_pe(self, x):
"""Reset the positional encodings."""
if self.pe is not None:
if self.pe.size(1) >= x.size(1):
if self.pe.dtype != x.dtype or self.pe.device != x.device:
self.pe = self.pe.to(dtype=x.dtype, device=x.device)
return
pe = torch.zeros(x.size(1), self.d_model)
if self.reverse:
position = torch.arange(
x.size(1) - 1, -1, -1.0, dtype=torch.float32
).unsqueeze(1)
else:
position = torch.arange(0, x.size(1), dtype=torch.float32).unsqueeze(1)
div_term = torch.exp(
torch.arange(0, self.d_model, 2, dtype=torch.float32)
* -(math.log(10000.0) / self.d_model)
)
pe[:, 0::2] = torch.sin(position * div_term)
pe[:, 1::2] = torch.cos(position * div_term)
pe = pe.unsqueeze(0)
self.pe = pe.to(device=x.device, dtype=x.dtype)
def forward(self, x: torch.Tensor):
"""Add positional encoding.
Args:
x (torch.Tensor): Input tensor (batch, time, `*`).
Returns:
torch.Tensor: Encoded tensor (batch, time, `*`).
"""
self.extend_pe(x)
x = x * self.xscale + self.pe[:, : x.size(1)]
return self.dropout(x)
class ScaledPositionalEncoding(PositionalEncoding):
"""Scaled positional encoding module.
See Sec. 3.2 https://arxiv.org/abs/1809.08895
Args:
d_model (int): Embedding dimension.
dropout_rate (float): Dropout rate.
max_len (int): Maximum input length.
"""
def __init__(self, d_model, dropout_rate, max_len=5000):
"""Initialize class."""
super().__init__(d_model=d_model, dropout_rate=dropout_rate, max_len=max_len)
self.alpha = torch.nn.Parameter(torch.tensor(1.0))
def reset_parameters(self):
"""Reset parameters."""
self.alpha.data = torch.tensor(1.0)
def forward(self, x):
"""Add positional encoding.
Args:
x (torch.Tensor): Input tensor (batch, time, `*`).
Returns:
torch.Tensor: Encoded tensor (batch, time, `*`).
"""
self.extend_pe(x)
x = x + self.alpha * self.pe[:, : x.size(1)]
return self.dropout(x)
class RelPositionalEncoding(PositionalEncoding):
"""Relative positional encoding module.
See : Appendix B in https://arxiv.org/abs/1901.02860
Args:
d_model (int): Embedding dimension.
dropout_rate (float): Dropout rate.
max_len (int): Maximum input length.
"""
def __init__(self, d_model, dropout_rate, max_len=5000):
"""Initialize class."""
super().__init__(d_model, dropout_rate, max_len, reverse=True)
def forward(self, x):
"""Compute positional encoding.
Args:
x (torch.Tensor): Input tensor (batch, time, `*`).
Returns:
torch.Tensor: Encoded tensor (batch, time, `*`).
torch.Tensor: Positional embedding tensor (1, time, `*`).
"""
self.extend_pe(x)
x = x * self.xscale
pos_emb = self.pe[:, : x.size(1)]
return self.dropout(x), self.dropout(pos_emb)
@@ -0,0 +1,198 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
# Copyright 2019 Shigeki Karita
# Apache 2.0 (http://www.apache.org/licenses/LICENSE-2.0)
"""Multi-Head Attention layer definition."""
from packaging import version
import math
import numpy
import torch
from torch import nn
class MultiHeadedAttention(nn.Module):
"""Multi-Head Attention layer.
Args:
n_head (int): The number of heads.
n_feat (int): The number of features.
dropout_rate (float): Dropout rate.
"""
def __init__(self, n_head, n_feat, dropout_rate, flash=False):
"""Construct an MultiHeadedAttention object."""
super(MultiHeadedAttention, self).__init__()
assert n_feat % n_head == 0
# We assume d_v always equals d_k
self.d_k = n_feat // n_head
self.h = n_head
self.linear_q = nn.Linear(n_feat, n_feat)
self.linear_k = nn.Linear(n_feat, n_feat)
self.linear_v = nn.Linear(n_feat, n_feat)
self.linear_out = nn.Linear(n_feat, n_feat)
self.attn = None
self.dropout = nn.Dropout(p=dropout_rate)
self.dropout_rate = dropout_rate
self.flash = flash
def forward_qkv(self, query, key, value):
"""Transform query, key and value.
Args:
query (torch.Tensor): Query tensor (#batch, time1, size).
key (torch.Tensor): Key tensor (#batch, time2, size).
value (torch.Tensor): Value tensor (#batch, time2, size).
Returns:
torch.Tensor: Transformed query tensor (#batch, n_head, time1, d_k).
torch.Tensor: Transformed key tensor (#batch, n_head, time2, d_k).
torch.Tensor: Transformed value tensor (#batch, n_head, time2, d_k).
"""
n_batch = query.size(0)
q = self.linear_q(query).view(n_batch, -1, self.h, self.d_k)
k = self.linear_k(key).view(n_batch, -1, self.h, self.d_k)
v = self.linear_v(value).view(n_batch, -1, self.h, self.d_k)
q = q.transpose(1, 2) # (batch, head, time1, d_k)
k = k.transpose(1, 2) # (batch, head, time2, d_k)
v = v.transpose(1, 2) # (batch, head, time2, d_k)
return q, k, v
def forward_attention(self, value, scores, mask):
"""Compute attention context vector.
Args:
value (torch.Tensor): Transformed value (#batch, n_head, time2, d_k).
scores (torch.Tensor): Attention score (#batch, n_head, time1, time2).
mask (torch.Tensor): Mask (#batch, 1, time2) or (#batch, time1, time2).
Returns:
torch.Tensor: Transformed value (#batch, time1, d_model)
weighted by the attention score (#batch, time1, time2).
"""
n_batch = value.size(0)
if mask is not None:
mask = mask.unsqueeze(1).eq(0) # (batch, 1, *, time2)
min_value = float(
numpy.finfo(torch.tensor(0, dtype=scores.dtype).numpy().dtype).min
)
scores = scores.masked_fill(mask, min_value)
self.attn = torch.softmax(scores, dim=-1).masked_fill(
mask, 0.0
) # (batch, head, time1, time2)
else:
self.attn = torch.softmax(scores, dim=-1) # (batch, head, time1, time2)
p_attn = self.dropout(self.attn)
x = torch.matmul(p_attn, value) # (batch, head, time1, d_k)
x = (
x.transpose(1, 2).contiguous().view(n_batch, -1, self.h * self.d_k)
) # (batch, time1, d_model)
return self.linear_out(x) # (batch, time1, d_model)
def forward(self, query, key, value, mask):
"""Compute scaled dot product attention.
Args:
query (torch.Tensor): Query tensor (#batch, time1, size).
key (torch.Tensor): Key tensor (#batch, time2, size).
value (torch.Tensor): Value tensor (#batch, time2, size).
mask (torch.Tensor): Mask tensor (#batch, 1, time2) or
(#batch, time1, time2).
Returns:
torch.Tensor: Output tensor (#batch, time1, d_model).
"""
q, k, v = self.forward_qkv(query, key, value)
if version.parse(torch.__version__) >= version.parse("2.0") and self.flash:
n_batch = value.size(0)
x = torch.nn.functional.scaled_dot_product_attention(
q, k, v, attn_mask=mask.unsqueeze(1) if mask is not None else None, dropout_p=self.dropout_rate)
x = (
x.transpose(1, 2).contiguous().view(n_batch, -1, self.h * self.d_k)
) # (batch, time1, d_model)
return self.linear_out(x)
else:
scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.d_k)
return self.forward_attention(v, scores, mask)
class RelPositionMultiHeadedAttention(MultiHeadedAttention):
"""Multi-Head Attention layer with relative position encoding.
Paper: https://arxiv.org/abs/1901.02860
Args:
n_head (int): The number of heads.
n_feat (int): The number of features.
dropout_rate (float): Dropout rate.
"""
def __init__(self, n_head, n_feat, dropout_rate):
"""Construct an RelPositionMultiHeadedAttention object."""
super().__init__(n_head, n_feat, dropout_rate)
# linear transformation for positional ecoding
self.linear_pos = nn.Linear(n_feat, n_feat, bias=False)
# these two learnable bias are used in matrix c and matrix d
# as described in https://arxiv.org/abs/1901.02860 Section 3.3
self.pos_bias_u = nn.Parameter(torch.Tensor(self.h, self.d_k))
self.pos_bias_v = nn.Parameter(torch.Tensor(self.h, self.d_k))
torch.nn.init.xavier_uniform_(self.pos_bias_u)
torch.nn.init.xavier_uniform_(self.pos_bias_v)
def rel_shift(self, x, zero_triu=False):
"""Compute relative positinal encoding.
Args:
x (torch.Tensor): Input tensor (batch, time, size).
zero_triu (bool): If true, return the lower triangular part of the matrix.
Returns:
torch.Tensor: Output tensor.
"""
zero_pad = torch.zeros((*x.size()[:3], 1), device=x.device, dtype=x.dtype)
x_padded = torch.cat([zero_pad, x], dim=-1)
x_padded = x_padded.view(*x.size()[:2], x.size(3) + 1, x.size(2))
x = x_padded[:, :, 1:].view_as(x)
if zero_triu:
ones = torch.ones((x.size(2), x.size(3)))
x = x * torch.tril(ones, x.size(3) - x.size(2))[None, None, :, :]
return x
def forward(self, query, key, value, pos_emb, mask):
"""Compute 'Scaled Dot Product Attention' with rel. positional encoding.
Args:
query (torch.Tensor): Query tensor (#batch, time1, size).
key (torch.Tensor): Key tensor (#batch, time2, size).
value (torch.Tensor): Value tensor (#batch, time2, size).
pos_emb (torch.Tensor): Positional embedding tensor (#batch, time2, size).
mask (torch.Tensor): Mask tensor (#batch, 1, time2) or
(#batch, time1, time2).
Returns:
torch.Tensor: Output tensor (#batch, time1, d_model).
"""
q, k, v = self.forward_qkv(query, key, value)
q = q.transpose(1, 2) # (batch, time1, head, d_k)
n_batch_pos = pos_emb.size(0)
p = self.linear_pos(pos_emb).view(n_batch_pos, -1, self.h, self.d_k)
p = p.transpose(1, 2) # (batch, head, time1, d_k)
# (batch, head, time1, d_k)
q_with_bias_u = (q + self.pos_bias_u).transpose(1, 2)
# (batch, head, time1, d_k)
q_with_bias_v = (q + self.pos_bias_v).transpose(1, 2)
# compute attention score
# first compute matrix a and matrix c
# as described in https://arxiv.org/abs/1901.02860 Section 3.3
# (batch, head, time1, time2)
matrix_ac = torch.matmul(q_with_bias_u, k.transpose(-2, -1))
# compute matrix b and matrix d
# (batch, head, time1, time2)
matrix_bd = torch.matmul(q_with_bias_v, p.transpose(-2, -1))
matrix_bd = self.rel_shift(matrix_bd)
scores = (matrix_ac + matrix_bd) / math.sqrt(
self.d_k
) # (batch, head, time1, time2)
return self.forward_attention(v, scores, mask)
@@ -0,0 +1,260 @@
from torch import nn
import torch
from ..layers import LayerNorm
class ConvolutionModule(nn.Module):
"""ConvolutionModule in Conformer model.
Args:
channels (int): The number of channels of conv layers.
kernel_size (int): Kernerl size of conv layers.
"""
def __init__(self, channels, kernel_size, activation=nn.ReLU(), bias=True):
"""Construct an ConvolutionModule object."""
super(ConvolutionModule, self).__init__()
# kernerl_size should be a odd number for 'SAME' padding
assert (kernel_size - 1) % 2 == 0
self.pointwise_conv1 = nn.Conv1d(
channels,
2 * channels,
kernel_size=1,
stride=1,
padding=0,
bias=bias,
)
self.depthwise_conv = nn.Conv1d(
channels,
channels,
kernel_size,
stride=1,
padding=(kernel_size - 1) // 2,
groups=channels,
bias=bias,
)
self.norm = nn.BatchNorm1d(channels)
self.pointwise_conv2 = nn.Conv1d(
channels,
channels,
kernel_size=1,
stride=1,
padding=0,
bias=bias,
)
self.activation = activation
def forward(self, x):
"""Compute convolution module.
Args:
x (torch.Tensor): Input tensor (#batch, time, channels).
Returns:
torch.Tensor: Output tensor (#batch, time, channels).
"""
# exchange the temporal dimension and the feature dimension
x = x.transpose(1, 2)
# GLU mechanism
x = self.pointwise_conv1(x) # (batch, 2*channel, dim)
x = nn.functional.glu(x, dim=1) # (batch, channel, dim)
# 1D Depthwise Conv
x = self.depthwise_conv(x)
x = self.activation(self.norm(x))
x = self.pointwise_conv2(x)
return x.transpose(1, 2)
class MultiLayeredConv1d(torch.nn.Module):
"""Multi-layered conv1d for Transformer block.
This is a module of multi-leyered conv1d designed
to replace positionwise feed-forward network
in Transforner block, which is introduced in
`FastSpeech: Fast, Robust and Controllable Text to Speech`_.
.. _`FastSpeech: Fast, Robust and Controllable Text to Speech`:
https://arxiv.org/pdf/1905.09263.pdf
"""
def __init__(self, in_chans, hidden_chans, kernel_size, dropout_rate):
"""Initialize MultiLayeredConv1d module.
Args:
in_chans (int): Number of input channels.
hidden_chans (int): Number of hidden channels.
kernel_size (int): Kernel size of conv1d.
dropout_rate (float): Dropout rate.
"""
super(MultiLayeredConv1d, self).__init__()
self.w_1 = torch.nn.Conv1d(
in_chans,
hidden_chans,
kernel_size,
stride=1,
padding=(kernel_size - 1) // 2,
)
self.w_2 = torch.nn.Conv1d(
hidden_chans,
in_chans,
kernel_size,
stride=1,
padding=(kernel_size - 1) // 2,
)
self.dropout = torch.nn.Dropout(dropout_rate)
def forward(self, x):
"""Calculate forward propagation.
Args:
x (torch.Tensor): Batch of input tensors (B, T, in_chans).
Returns:
torch.Tensor: Batch of output tensors (B, T, hidden_chans).
"""
x = torch.relu(self.w_1(x.transpose(-1, 1))).transpose(-1, 1)
return self.w_2(self.dropout(x).transpose(-1, 1)).transpose(-1, 1)
class Swish(torch.nn.Module):
"""Construct an Swish object."""
def forward(self, x):
"""Return Swich activation function."""
return x * torch.sigmoid(x)
class EncoderLayer(nn.Module):
"""Encoder layer module.
Args:
size (int): Input dimension.
self_attn (torch.nn.Module): Self-attention module instance.
`MultiHeadedAttention` or `RelPositionMultiHeadedAttention` instance
can be used as the argument.
feed_forward (torch.nn.Module): Feed-forward module instance.
`PositionwiseFeedForward`, `MultiLayeredConv1d`, or `Conv1dLinear` instance
can be used as the argument.
feed_forward_macaron (torch.nn.Module): Additional feed-forward module instance.
`PositionwiseFeedForward`, `MultiLayeredConv1d`, or `Conv1dLinear` instance
can be used as the argument.
conv_module (torch.nn.Module): Convolution module instance.
`ConvlutionModule` instance can be used as the argument.
dropout_rate (float): Dropout rate.
normalize_before (bool): Whether to use layer_norm before the first block.
concat_after (bool): Whether to concat attention layer's input and output.
if True, additional linear will be applied.
i.e. x -> x + linear(concat(x, att(x)))
if False, no additional linear will be applied. i.e. x -> x + att(x)
"""
def __init__(
self,
size,
self_attn,
feed_forward,
feed_forward_macaron,
conv_module,
dropout_rate,
normalize_before=True,
concat_after=False,
):
"""Construct an EncoderLayer object."""
super(EncoderLayer, self).__init__()
self.self_attn = self_attn
self.feed_forward = feed_forward
self.feed_forward_macaron = feed_forward_macaron
self.conv_module = conv_module
self.norm_ff = LayerNorm(size) # for the FNN module
self.norm_mha = LayerNorm(size) # for the MHA module
if feed_forward_macaron is not None:
self.norm_ff_macaron = LayerNorm(size)
self.ff_scale = 0.5
else:
self.ff_scale = 1.0
if self.conv_module is not None:
self.norm_conv = LayerNorm(size) # for the CNN module
self.norm_final = LayerNorm(size) # for the final output of the block
self.dropout = nn.Dropout(dropout_rate)
self.size = size
self.normalize_before = normalize_before
self.concat_after = concat_after
if self.concat_after:
self.concat_linear = nn.Linear(size + size, size)
def forward(self, x_input, mask, cache=None):
"""Compute encoded features.
Args:
x_input (Union[Tuple, torch.Tensor]): Input tensor w/ or w/o pos emb.
- w/ pos emb: Tuple of tensors [(#batch, time, size), (1, time, size)].
- w/o pos emb: Tensor (#batch, time, size).
mask (torch.Tensor): Mask tensor for the input (#batch, time).
cache (torch.Tensor): Cache tensor of the input (#batch, time - 1, size).
Returns:
torch.Tensor: Output tensor (#batch, time, size).
torch.Tensor: Mask tensor (#batch, time).
"""
if isinstance(x_input, tuple):
x, pos_emb = x_input[0], x_input[1]
else:
x, pos_emb = x_input, None
# whether to use macaron style
if self.feed_forward_macaron is not None:
residual = x
if self.normalize_before:
x = self.norm_ff_macaron(x)
x = residual + self.ff_scale * self.dropout(self.feed_forward_macaron(x))
if not self.normalize_before:
x = self.norm_ff_macaron(x)
# multi-headed self-attention module
residual = x
if self.normalize_before:
x = self.norm_mha(x)
if cache is None:
x_q = x
else:
assert cache.shape == (x.shape[0], x.shape[1] - 1, self.size)
x_q = x[:, -1:, :]
residual = residual[:, -1:, :]
mask = None if mask is None else mask[:, -1:, :]
if pos_emb is not None:
x_att = self.self_attn(x_q, x, x, pos_emb, mask)
else:
x_att = self.self_attn(x_q, x, x, mask)
if self.concat_after:
x_concat = torch.cat((x, x_att), dim=-1)
x = residual + self.concat_linear(x_concat)
else:
x = residual + self.dropout(x_att)
if not self.normalize_before:
x = self.norm_mha(x)
# convolution module
if self.conv_module is not None:
residual = x
if self.normalize_before:
x = self.norm_conv(x)
x = residual + self.dropout(self.conv_module(x))
if not self.normalize_before:
x = self.norm_conv(x)
# feed forward module
residual = x
if self.normalize_before:
x = self.norm_ff(x)
x = residual + self.ff_scale * self.dropout(self.feed_forward(x))
if not self.normalize_before:
x = self.norm_ff(x)
if self.conv_module is not None:
x = self.norm_final(x)
if cache is not None:
x = torch.cat([cache, x], dim=1)
if pos_emb is not None:
return (x, pos_emb), mask
return x, mask
@@ -0,0 +1,175 @@
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from .layers import LayerNorm, Embedding
class LambdaLayer(nn.Module):
def __init__(self, lambd):
super(LambdaLayer, self).__init__()
self.lambd = lambd
def forward(self, x):
return self.lambd(x)
def init_weights_func(m):
classname = m.__class__.__name__
if classname.find("Conv1d") != -1:
torch.nn.init.xavier_uniform_(m.weight)
def get_norm_builder(norm_type, channels, ln_eps=1e-6):
if norm_type == 'bn':
norm_builder = lambda: nn.BatchNorm1d(channels)
elif norm_type == 'in':
norm_builder = lambda: nn.InstanceNorm1d(channels, affine=True)
elif norm_type == 'gn':
norm_builder = lambda: nn.GroupNorm(8, channels)
elif norm_type == 'ln':
norm_builder = lambda: LayerNorm(channels, dim=1, eps=ln_eps)
else:
norm_builder = lambda: nn.Identity()
return norm_builder
def get_act_builder(act_type):
if act_type == 'gelu':
act_builder = lambda: nn.GELU()
elif act_type == 'relu':
act_builder = lambda: nn.ReLU(inplace=True)
elif act_type == 'leakyrelu':
act_builder = lambda: nn.LeakyReLU(negative_slope=0.01, inplace=True)
elif act_type == 'swish':
act_builder = lambda: nn.SiLU(inplace=True)
else:
act_builder = lambda: nn.Identity()
return act_builder
class ResidualBlock(nn.Module):
"""Implements conv->PReLU->norm n-times"""
def __init__(self, channels, kernel_size, dilation, n=2, norm_type='bn', dropout=0.0,
c_multiple=2, ln_eps=1e-12, act_type='gelu'):
super(ResidualBlock, self).__init__()
norm_builder = get_norm_builder(norm_type, channels, ln_eps)
act_builder = get_act_builder(act_type)
self.blocks = [
nn.Sequential(
norm_builder(),
nn.Conv1d(channels, c_multiple * channels, kernel_size, dilation=dilation,
padding=(dilation * (kernel_size - 1)) // 2),
LambdaLayer(lambda x: x * kernel_size ** -0.5),
act_builder(),
nn.Conv1d(c_multiple * channels, channels, 1, dilation=dilation),
)
for i in range(n)
]
self.blocks = nn.ModuleList(self.blocks)
self.dropout = dropout
def forward(self, x):
nonpadding = (x.abs().sum(1) > 0).float()[:, None, :]
for b in self.blocks:
x_ = b(x)
if self.dropout > 0 and self.training:
x_ = F.dropout(x_, self.dropout, training=self.training)
x = x + x_
x = x * nonpadding
return x
class ConvBlocks(nn.Module):
"""Decodes the expanded phoneme encoding into spectrograms"""
def __init__(self, hidden_size, out_dims, dilations, kernel_size,
norm_type='ln', layers_in_block=2, c_multiple=2,
dropout=0.0, ln_eps=1e-5,
init_weights=True, is_BTC=True, num_layers=None, post_net_kernel=3, act_type='gelu'):
super(ConvBlocks, self).__init__()
self.is_BTC = is_BTC
if num_layers is not None:
dilations = [1] * num_layers
self.res_blocks = nn.Sequential(
*[ResidualBlock(hidden_size, kernel_size, d,
n=layers_in_block, norm_type=norm_type, c_multiple=c_multiple,
dropout=dropout, ln_eps=ln_eps, act_type=act_type)
for d in dilations],
)
norm = get_norm_builder(norm_type, hidden_size, ln_eps)()
self.last_norm = norm
self.post_net1 = nn.Conv1d(hidden_size, out_dims, kernel_size=post_net_kernel,
padding=post_net_kernel // 2)
if init_weights:
self.apply(init_weights_func)
def forward(self, x, nonpadding=None):
"""
:param x: [B, T, H]
:return: [B, T, H]
"""
if self.is_BTC:
x = x.transpose(1, 2)
if nonpadding is None:
nonpadding = (x.abs().sum(1) > 0).float()[:, None, :]
elif self.is_BTC:
nonpadding = nonpadding.transpose(1, 2)
x = self.res_blocks(x) * nonpadding
x = self.last_norm(x) * nonpadding
x = self.post_net1(x) * nonpadding
if self.is_BTC:
x = x.transpose(1, 2)
return x
class TextConvEncoder(ConvBlocks):
def __init__(self, dict_size, hidden_size, out_dims, dilations, kernel_size,
norm_type='ln', layers_in_block=2, c_multiple=2,
dropout=0.0, ln_eps=1e-5, init_weights=True, num_layers=None, post_net_kernel=3):
super().__init__(hidden_size, out_dims, dilations, kernel_size,
norm_type, layers_in_block, c_multiple,
dropout, ln_eps, init_weights, num_layers=num_layers,
post_net_kernel=post_net_kernel)
self.embed_tokens = Embedding(dict_size, hidden_size, 0)
self.embed_scale = math.sqrt(hidden_size)
def forward(self, txt_tokens):
"""
:param txt_tokens: [B, T]
:return: {
'encoder_out': [B x T x C]
}
"""
x = self.embed_scale * self.embed_tokens(txt_tokens)
return super().forward(x)
class ConditionalConvBlocks(ConvBlocks):
def __init__(self, hidden_size, c_cond, c_out, dilations, kernel_size,
norm_type='ln', layers_in_block=2, c_multiple=2,
dropout=0.0, ln_eps=1e-5, init_weights=True, is_BTC=True, num_layers=None):
super().__init__(hidden_size, c_out, dilations, kernel_size,
norm_type, layers_in_block, c_multiple,
dropout, ln_eps, init_weights, is_BTC=False, num_layers=num_layers)
self.g_prenet = nn.Conv1d(c_cond, hidden_size, 3, padding=1)
self.is_BTC_ = is_BTC
if init_weights:
self.g_prenet.apply(init_weights_func)
def forward(self, x, cond, nonpadding=None):
if self.is_BTC_:
x = x.transpose(1, 2)
cond = cond.transpose(1, 2)
if nonpadding is not None:
nonpadding = nonpadding.transpose(1, 2)
if nonpadding is None:
nonpadding = x.abs().sum(1)[:, None]
x = x + self.g_prenet(cond)
x = x * nonpadding
x = super(ConditionalConvBlocks, self).forward(x) # input needs to be BTC
if self.is_BTC_:
x = x.transpose(1, 2)
return x
@@ -0,0 +1,85 @@
import torch
from torch import nn
from torch.autograd import Function
class LayerNorm(torch.nn.LayerNorm):
"""Layer normalization module.
:param int nout: output dim size
:param int dim: dimension to be normalized
"""
def __init__(self, nout, dim=-1, eps=1e-5):
"""Construct an LayerNorm object."""
super(LayerNorm, self).__init__(nout, eps=eps)
self.dim = dim
def forward(self, x):
"""Apply layer normalization.
:param torch.Tensor x: input tensor
:return: layer normalized tensor
:rtype torch.Tensor
"""
if self.dim == -1:
return super(LayerNorm, self).forward(x)
return super(LayerNorm, self).forward(x.transpose(1, -1)).transpose(1, -1)
class Reshape(nn.Module):
def __init__(self, *args):
super(Reshape, self).__init__()
self.shape = args
def forward(self, x):
return x.view(self.shape)
class Permute(nn.Module):
def __init__(self, *args):
super(Permute, self).__init__()
self.args = args
def forward(self, x):
return x.permute(self.args)
def Linear(in_features, out_features, bias=True, init_type='xavier'):
m = nn.Linear(in_features, out_features, bias)
if init_type == 'xavier':
nn.init.xavier_uniform_(m.weight)
elif init_type == 'kaiming':
nn.init.kaiming_normal_(m.weight, mode='fan_in')
if bias:
nn.init.constant_(m.bias, 0.)
return m
def Embedding(num_embeddings, embedding_dim, padding_idx=None, init_type='normal'):
m = nn.Embedding(num_embeddings, embedding_dim, padding_idx=padding_idx)
if init_type == 'normal':
nn.init.normal_(m.weight, mean=0, std=embedding_dim ** -0.5)
elif init_type == 'kaiming':
nn.init.kaiming_normal_(m.weight, mode='fan_in')
if padding_idx is not None:
nn.init.constant_(m.weight[padding_idx], 0)
return m
class GradientReverseFunction(Function):
@staticmethod
def forward(ctx, input, coeff=1.):
ctx.coeff = coeff
output = input * 1.0
return output
@staticmethod
def backward(ctx, grad_output):
return grad_output.neg() * ctx.coeff, None
class GRL(nn.Module):
def __init__(self):
super(GRL, self).__init__()
def forward(self, *input):
return GradientReverseFunction.apply(*input)
@@ -0,0 +1,378 @@
import math
import torch
from torch import nn
from torch.nn import functional as F
from .layers import Embedding
def convert_pad_shape(pad_shape):
l = pad_shape[::-1]
pad_shape = [item for sublist in l for item in sublist]
return pad_shape
def shift_1d(x):
x = F.pad(x, convert_pad_shape([[0, 0], [0, 0], [1, 0]]))[:, :, :-1]
return x
def sequence_mask(length, max_length=None):
if max_length is None:
max_length = length.max()
x = torch.arange(max_length, dtype=length.dtype, device=length.device)
return x.unsqueeze(0) < length.unsqueeze(1)
class Encoder(nn.Module):
def __init__(self, hidden_channels, filter_channels, n_heads, n_layers, kernel_size=1, p_dropout=0.,
window_size=None, block_length=None, pre_ln=False, **kwargs):
super().__init__()
self.hidden_channels = hidden_channels
self.filter_channels = filter_channels
self.n_heads = n_heads
self.n_layers = n_layers
self.kernel_size = kernel_size
self.p_dropout = p_dropout
self.window_size = window_size
self.block_length = block_length
self.pre_ln = pre_ln
self.drop = nn.Dropout(p_dropout)
self.attn_layers = nn.ModuleList()
self.norm_layers_1 = nn.ModuleList()
self.ffn_layers = nn.ModuleList()
self.norm_layers_2 = nn.ModuleList()
for i in range(self.n_layers):
self.attn_layers.append(
MultiHeadAttention(hidden_channels, hidden_channels, n_heads, window_size=window_size,
p_dropout=p_dropout, block_length=block_length))
self.norm_layers_1.append(LayerNorm(hidden_channels))
self.ffn_layers.append(
FFN(hidden_channels, hidden_channels, filter_channels, kernel_size, p_dropout=p_dropout))
self.norm_layers_2.append(LayerNorm(hidden_channels))
if pre_ln:
self.last_ln = LayerNorm(hidden_channels)
def forward(self, x, x_mask):
attn_mask = x_mask.unsqueeze(2) * x_mask.unsqueeze(-1)
for i in range(self.n_layers):
x = x * x_mask
x_ = x
if self.pre_ln:
x = self.norm_layers_1[i](x)
y = self.attn_layers[i](x, x, attn_mask)
y = self.drop(y)
x = x_ + y
if not self.pre_ln:
x = self.norm_layers_1[i](x)
x_ = x
if self.pre_ln:
x = self.norm_layers_2[i](x)
y = self.ffn_layers[i](x, x_mask)
y = self.drop(y)
x = x_ + y
if not self.pre_ln:
x = self.norm_layers_2[i](x)
if self.pre_ln:
x = self.last_ln(x)
x = x * x_mask
return x
class MultiHeadAttention(nn.Module):
def __init__(self, channels, out_channels, n_heads, window_size=None, heads_share=True, p_dropout=0.,
block_length=None, proximal_bias=False, proximal_init=False):
super().__init__()
assert channels % n_heads == 0
self.channels = channels
self.out_channels = out_channels
self.n_heads = n_heads
self.window_size = window_size
self.heads_share = heads_share
self.block_length = block_length
self.proximal_bias = proximal_bias
self.p_dropout = p_dropout
self.attn = None
self.k_channels = channels // n_heads
self.conv_q = nn.Conv1d(channels, channels, 1)
self.conv_k = nn.Conv1d(channels, channels, 1)
self.conv_v = nn.Conv1d(channels, channels, 1)
if window_size is not None:
n_heads_rel = 1 if heads_share else n_heads
rel_stddev = self.k_channels ** -0.5
self.emb_rel_k = nn.Parameter(torch.randn(n_heads_rel, window_size * 2 + 1, self.k_channels) * rel_stddev)
self.emb_rel_v = nn.Parameter(torch.randn(n_heads_rel, window_size * 2 + 1, self.k_channels) * rel_stddev)
self.conv_o = nn.Conv1d(channels, out_channels, 1)
self.drop = nn.Dropout(p_dropout)
nn.init.xavier_uniform_(self.conv_q.weight)
nn.init.xavier_uniform_(self.conv_k.weight)
if proximal_init:
self.conv_k.weight.data.copy_(self.conv_q.weight.data)
self.conv_k.bias.data.copy_(self.conv_q.bias.data)
nn.init.xavier_uniform_(self.conv_v.weight)
def forward(self, x, c, attn_mask=None):
q = self.conv_q(x)
k = self.conv_k(c)
v = self.conv_v(c)
x, self.attn = self.attention(q, k, v, mask=attn_mask)
x = self.conv_o(x)
return x
def attention(self, query, key, value, mask=None):
# reshape [b, d, t] -> [b, n_h, t, d_k]
b, d, t_s, t_t = (*key.size(), query.size(2))
query = query.view(b, self.n_heads, self.k_channels, t_t).transpose(2, 3)
key = key.view(b, self.n_heads, self.k_channels, t_s).transpose(2, 3)
value = value.view(b, self.n_heads, self.k_channels, t_s).transpose(2, 3)
scores = torch.matmul(query, key.transpose(-2, -1)) / math.sqrt(self.k_channels)
if self.window_size is not None:
assert t_s == t_t, "Relative attention is only available for self-attention."
key_relative_embeddings = self._get_relative_embeddings(self.emb_rel_k, t_s)
rel_logits = self._matmul_with_relative_keys(query, key_relative_embeddings)
rel_logits = self._relative_position_to_absolute_position(rel_logits)
scores_local = rel_logits / math.sqrt(self.k_channels)
scores = scores + scores_local
if self.proximal_bias:
assert t_s == t_t, "Proximal bias is only available for self-attention."
scores = scores + self._attention_bias_proximal(t_s).to(device=scores.device, dtype=scores.dtype)
if mask is not None:
scores = scores.masked_fill(mask == 0, -1e4)
if self.block_length is not None:
block_mask = torch.ones_like(scores).triu(-self.block_length).tril(self.block_length)
scores = scores * block_mask + -1e4 * (1 - block_mask)
p_attn = F.softmax(scores, dim=-1) # [b, n_h, t_t, t_s]
p_attn = self.drop(p_attn)
output = torch.matmul(p_attn, value)
if self.window_size is not None:
relative_weights = self._absolute_position_to_relative_position(p_attn)
value_relative_embeddings = self._get_relative_embeddings(self.emb_rel_v, t_s)
output = output + self._matmul_with_relative_values(relative_weights, value_relative_embeddings)
output = output.transpose(2, 3).contiguous().view(b, d, t_t) # [b, n_h, t_t, d_k] -> [b, d, t_t]
return output, p_attn
def _matmul_with_relative_values(self, x, y):
"""
x: [b, h, l, m]
y: [h or 1, m, d]
ret: [b, h, l, d]
"""
ret = torch.matmul(x, y.unsqueeze(0))
return ret
def _matmul_with_relative_keys(self, x, y):
"""
x: [b, h, l, d]
y: [h or 1, m, d]
ret: [b, h, l, m]
"""
ret = torch.matmul(x, y.unsqueeze(0).transpose(-2, -1))
return ret
def _get_relative_embeddings(self, relative_embeddings, length):
max_relative_position = 2 * self.window_size + 1
# Pad first before slice to avoid using cond ops.
pad_length = max(length - (self.window_size + 1), 0)
slice_start_position = max((self.window_size + 1) - length, 0)
slice_end_position = slice_start_position + 2 * length - 1
if pad_length > 0:
padded_relative_embeddings = F.pad(
relative_embeddings,
convert_pad_shape([[0, 0], [pad_length, pad_length], [0, 0]]))
else:
padded_relative_embeddings = relative_embeddings
used_relative_embeddings = padded_relative_embeddings[:, slice_start_position:slice_end_position]
return used_relative_embeddings
def _relative_position_to_absolute_position(self, x):
"""
x: [b, h, l, 2*l-1]
ret: [b, h, l, l]
"""
batch, heads, length, _ = x.size()
# Concat columns of pad to shift from relative to absolute indexing.
x = F.pad(x, convert_pad_shape([[0, 0], [0, 0], [0, 0], [0, 1]]))
# Concat extra elements so to add up to shape (len+1, 2*len-1).
x_flat = x.view([batch, heads, length * 2 * length])
x_flat = F.pad(x_flat, convert_pad_shape([[0, 0], [0, 0], [0, length - 1]]))
# Reshape and slice out the padded elements.
x_final = x_flat.view([batch, heads, length + 1, 2 * length - 1])[:, :, :length, length - 1:]
return x_final
def _absolute_position_to_relative_position(self, x):
"""
x: [b, h, l, l]
ret: [b, h, l, 2*l-1]
"""
batch, heads, length, _ = x.size()
# padd along column
x = F.pad(x, convert_pad_shape([[0, 0], [0, 0], [0, 0], [0, length - 1]]))
x_flat = x.view([batch, heads, length ** 2 + length * (length - 1)])
# add 0's in the beginning that will skew the elements after reshape
x_flat = F.pad(x_flat, convert_pad_shape([[0, 0], [0, 0], [length, 0]]))
x_final = x_flat.view([batch, heads, length, 2 * length])[:, :, :, 1:]
return x_final
def _attention_bias_proximal(self, length):
"""Bias for self-attention to encourage attention to close positions.
Args:
length: an integer scalar.
Returns:
a Tensor with shape [1, 1, length, length]
"""
r = torch.arange(length, dtype=torch.float32)
diff = torch.unsqueeze(r, 0) - torch.unsqueeze(r, 1)
return torch.unsqueeze(torch.unsqueeze(-torch.log1p(torch.abs(diff)), 0), 0)
class FFN(nn.Module):
def __init__(self, in_channels, out_channels, filter_channels, kernel_size, p_dropout=0., activation=None):
super().__init__()
self.in_channels = in_channels
self.out_channels = out_channels
self.filter_channels = filter_channels
self.kernel_size = kernel_size
self.p_dropout = p_dropout
self.activation = activation
self.conv_1 = nn.Conv1d(in_channels, filter_channels, kernel_size, padding=kernel_size // 2)
self.conv_2 = nn.Conv1d(filter_channels, out_channels, 1)
self.drop = nn.Dropout(p_dropout)
def forward(self, x, x_mask):
x = self.conv_1(x * x_mask)
if self.activation == "gelu":
x = x * torch.sigmoid(1.702 * x)
else:
x = torch.relu(x)
x = self.drop(x)
x = self.conv_2(x * x_mask)
return x * x_mask
class LayerNorm(nn.Module):
def __init__(self, channels, eps=1e-4):
super().__init__()
self.channels = channels
self.eps = eps
self.gamma = nn.Parameter(torch.ones(channels))
self.beta = nn.Parameter(torch.zeros(channels))
def forward(self, x):
n_dims = len(x.shape)
mean = torch.mean(x, 1, keepdim=True)
variance = torch.mean((x - mean) ** 2, 1, keepdim=True)
x = (x - mean) * torch.rsqrt(variance + self.eps)
shape = [1, -1] + [1] * (n_dims - 2)
x = x * self.gamma.view(*shape) + self.beta.view(*shape)
return x
class ConvReluNorm(nn.Module):
def __init__(self, in_channels, hidden_channels, out_channels, kernel_size, n_layers, p_dropout):
super().__init__()
self.in_channels = in_channels
self.hidden_channels = hidden_channels
self.out_channels = out_channels
self.kernel_size = kernel_size
self.n_layers = n_layers
self.p_dropout = p_dropout
assert n_layers > 1, "Number of layers should be larger than 0."
self.conv_layers = nn.ModuleList()
self.norm_layers = nn.ModuleList()
self.conv_layers.append(nn.Conv1d(in_channels, hidden_channels, kernel_size, padding=kernel_size // 2))
self.norm_layers.append(LayerNorm(hidden_channels))
self.relu_drop = nn.Sequential(
nn.ReLU(),
nn.Dropout(p_dropout))
for _ in range(n_layers - 1):
self.conv_layers.append(nn.Conv1d(hidden_channels, hidden_channels, kernel_size, padding=kernel_size // 2))
self.norm_layers.append(LayerNorm(hidden_channels))
self.proj = nn.Conv1d(hidden_channels, out_channels, 1)
self.proj.weight.data.zero_()
self.proj.bias.data.zero_()
def forward(self, x, x_mask):
x_org = x
for i in range(self.n_layers):
x = self.conv_layers[i](x * x_mask)
x = self.norm_layers[i](x)
x = self.relu_drop(x)
x = x_org + self.proj(x)
return x * x_mask
class RelTransformerEncoder(nn.Module):
def __init__(self,
n_vocab,
out_channels,
hidden_channels,
filter_channels,
n_heads,
n_layers,
kernel_size,
p_dropout=0.0,
window_size=4,
block_length=None,
prenet=True,
pre_ln=True,
):
super().__init__()
self.n_vocab = n_vocab
self.out_channels = out_channels
self.hidden_channels = hidden_channels
self.filter_channels = filter_channels
self.n_heads = n_heads
self.n_layers = n_layers
self.kernel_size = kernel_size
self.p_dropout = p_dropout
self.window_size = window_size
self.block_length = block_length
self.prenet = prenet
if n_vocab > 0:
self.emb = Embedding(n_vocab, hidden_channels, padding_idx=0)
if prenet:
self.pre = ConvReluNorm(hidden_channels, hidden_channels, hidden_channels,
kernel_size=5, n_layers=3, p_dropout=0)
self.encoder = Encoder(
hidden_channels,
filter_channels,
n_heads,
n_layers,
kernel_size,
p_dropout,
window_size=window_size,
block_length=block_length,
pre_ln=pre_ln,
)
def forward(self, x, x_mask=None):
if self.n_vocab > 0:
x_lengths = (x > 0).long().sum(-1)
x = self.emb(x) * math.sqrt(self.hidden_channels) # [b, t, h]
else:
x_lengths = (x.abs().sum(-1) > 0).long().sum(-1)
x = torch.transpose(x, 1, -1) # [b, h, t]
x_mask = torch.unsqueeze(sequence_mask(x_lengths, x.size(2)), 1).to(x.dtype)
if self.prenet:
x = self.pre(x, x_mask)
x = self.encoder(x, x_mask)
return x.transpose(1, 2)
@@ -0,0 +1,261 @@
import torch
from torch import nn
import torch.nn.functional as F
class PreNet(nn.Module):
def __init__(self, in_dims, fc1_dims=256, fc2_dims=128, dropout=0.5):
super().__init__()
self.fc1 = nn.Linear(in_dims, fc1_dims)
self.fc2 = nn.Linear(fc1_dims, fc2_dims)
self.p = dropout
def forward(self, x):
x = self.fc1(x)
x = F.relu(x)
x = F.dropout(x, self.p, training=self.training)
x = self.fc2(x)
x = F.relu(x)
x = F.dropout(x, self.p, training=self.training)
return x
class HighwayNetwork(nn.Module):
def __init__(self, size):
super().__init__()
self.W1 = nn.Linear(size, size)
self.W2 = nn.Linear(size, size)
self.W1.bias.data.fill_(0.)
def forward(self, x):
x1 = self.W1(x)
x2 = self.W2(x)
g = torch.sigmoid(x2)
y = g * F.relu(x1) + (1. - g) * x
return y
class BatchNormConv(nn.Module):
def __init__(self, in_channels, out_channels, kernel, relu=True):
super().__init__()
self.conv = nn.Conv1d(in_channels, out_channels, kernel, stride=1, padding=kernel // 2, bias=False)
self.bnorm = nn.BatchNorm1d(out_channels)
self.relu = relu
def forward(self, x):
x = self.conv(x)
x = F.relu(x) if self.relu is True else x
return self.bnorm(x)
class ConvNorm(torch.nn.Module):
def __init__(self, in_channels, out_channels, kernel_size=1, stride=1,
padding=None, dilation=1, bias=True, w_init_gain='linear'):
super(ConvNorm, self).__init__()
if padding is None:
assert (kernel_size % 2 == 1)
padding = int(dilation * (kernel_size - 1) / 2)
self.conv = torch.nn.Conv1d(in_channels, out_channels,
kernel_size=kernel_size, stride=stride,
padding=padding, dilation=dilation,
bias=bias)
torch.nn.init.xavier_uniform_(
self.conv.weight, gain=torch.nn.init.calculate_gain(w_init_gain))
def forward(self, signal):
conv_signal = self.conv(signal)
return conv_signal
class CBHG(nn.Module):
def __init__(self, K, in_channels, channels, proj_channels, num_highways):
super().__init__()
# List of all rnns to call `flatten_parameters()` on
self._to_flatten = []
self.bank_kernels = [i for i in range(1, K + 1)]
self.conv1d_bank = nn.ModuleList()
for k in self.bank_kernels:
conv = BatchNormConv(in_channels, channels, k)
self.conv1d_bank.append(conv)
self.maxpool = nn.MaxPool1d(kernel_size=2, stride=1, padding=1)
self.conv_project1 = BatchNormConv(len(self.bank_kernels) * channels, proj_channels[0], 3)
self.conv_project2 = BatchNormConv(proj_channels[0], proj_channels[1], 3, relu=False)
# Fix the highway input if necessary
if proj_channels[-1] != channels:
self.highway_mismatch = True
self.pre_highway = nn.Linear(proj_channels[-1], channels, bias=False)
else:
self.highway_mismatch = False
self.highways = nn.ModuleList()
for i in range(num_highways):
hn = HighwayNetwork(channels)
self.highways.append(hn)
self.rnn = nn.GRU(channels, channels, batch_first=True, bidirectional=True)
self._to_flatten.append(self.rnn)
# Avoid fragmentation of RNN parameters and associated warning
self._flatten_parameters()
def forward(self, x):
# Although we `_flatten_parameters()` on init, when using DataParallel
# the model gets replicated, making it no longer guaranteed that the
# weights are contiguous in GPU memory. Hence, we must call it again
self._flatten_parameters()
# Save these for later
residual = x
seq_len = x.size(-1)
conv_bank = []
# Convolution Bank
for conv in self.conv1d_bank:
c = conv(x) # Convolution
conv_bank.append(c[:, :, :seq_len])
# Stack along the channel axis
conv_bank = torch.cat(conv_bank, dim=1)
# dump the last padding to fit residual
x = self.maxpool(conv_bank)[:, :, :seq_len]
# Conv1d projections
x = self.conv_project1(x)
x = self.conv_project2(x)
# Residual Connect
x = x + residual
# Through the highways
x = x.transpose(1, 2)
if self.highway_mismatch is True:
x = self.pre_highway(x)
for h in self.highways:
x = h(x)
# And then the RNN
x, _ = self.rnn(x)
return x
def _flatten_parameters(self):
"""Calls `flatten_parameters` on all the rnns used by the WaveRNN. Used
to improve efficiency and avoid PyTorch yelling at us."""
[m.flatten_parameters() for m in self._to_flatten]
class TacotronEncoder(nn.Module):
def __init__(self, embed_dims, num_chars, cbhg_channels, K, num_highways, dropout):
super().__init__()
self.embedding = nn.Embedding(num_chars, embed_dims)
self.pre_net = PreNet(embed_dims, embed_dims, embed_dims, dropout=dropout)
self.cbhg = CBHG(K=K, in_channels=cbhg_channels, channels=cbhg_channels,
proj_channels=[cbhg_channels, cbhg_channels],
num_highways=num_highways)
self.proj_out = nn.Linear(cbhg_channels * 2, cbhg_channels)
def forward(self, x):
x = self.embedding(x)
x = self.pre_net(x)
x.transpose_(1, 2)
x = self.cbhg(x)
x = self.proj_out(x)
return x
class RNNEncoder(nn.Module):
def __init__(self, num_chars, embedding_dim, n_convolutions=3, kernel_size=5):
super(RNNEncoder, self).__init__()
self.embedding = nn.Embedding(num_chars, embedding_dim, padding_idx=0)
convolutions = []
for _ in range(n_convolutions):
conv_layer = nn.Sequential(
ConvNorm(embedding_dim,
embedding_dim,
kernel_size=kernel_size, stride=1,
padding=int((kernel_size - 1) / 2),
dilation=1, w_init_gain='relu'),
nn.BatchNorm1d(embedding_dim))
convolutions.append(conv_layer)
self.convolutions = nn.ModuleList(convolutions)
self.lstm = nn.LSTM(embedding_dim, int(embedding_dim / 2), 1,
batch_first=True, bidirectional=True)
def forward(self, x):
input_lengths = (x > 0).sum(-1)
input_lengths = input_lengths.cpu().numpy()
x = self.embedding(x)
x = x.transpose(1, 2) # [B, H, T]
for conv in self.convolutions:
x = F.dropout(F.relu(conv(x)), 0.5, self.training) + x
x = x.transpose(1, 2) # [B, T, H]
# pytorch tensor are not reversible, hence the conversion
x = nn.utils.rnn.pack_padded_sequence(x, input_lengths, batch_first=True, enforce_sorted=False)
self.lstm.flatten_parameters()
outputs, _ = self.lstm(x)
outputs, _ = nn.utils.rnn.pad_packed_sequence(outputs, batch_first=True)
return outputs
class DecoderRNN(torch.nn.Module):
def __init__(self, hidden_size, decoder_rnn_dim, dropout):
super(DecoderRNN, self).__init__()
self.in_conv1d = nn.Sequential(
torch.nn.Conv1d(
in_channels=hidden_size,
out_channels=hidden_size,
kernel_size=9, padding=4,
),
torch.nn.ReLU(),
torch.nn.Conv1d(
in_channels=hidden_size,
out_channels=hidden_size,
kernel_size=9, padding=4,
),
)
self.ln = nn.LayerNorm(hidden_size)
if decoder_rnn_dim == 0:
decoder_rnn_dim = hidden_size * 2
self.rnn = torch.nn.LSTM(
input_size=hidden_size,
hidden_size=decoder_rnn_dim,
num_layers=1,
batch_first=True,
bidirectional=True,
dropout=dropout
)
self.rnn.flatten_parameters()
self.conv1d = torch.nn.Conv1d(
in_channels=decoder_rnn_dim * 2,
out_channels=hidden_size,
kernel_size=3,
padding=1,
)
def forward(self, x):
input_masks = x.abs().sum(-1).ne(0).data[:, :, None]
input_lengths = input_masks.sum([-1, -2])
input_lengths = input_lengths.cpu().numpy()
x = self.in_conv1d(x.transpose(1, 2)).transpose(1, 2)
x = self.ln(x)
x = nn.utils.rnn.pack_padded_sequence(x, input_lengths, batch_first=True, enforce_sorted=False)
self.rnn.flatten_parameters()
x, _ = self.rnn(x) # [B, T, C]
x, _ = nn.utils.rnn.pad_packed_sequence(x, batch_first=True)
x = x * input_masks
pre_mel = self.conv1d(x.transpose(1, 2)).transpose(1, 2) # [B, T, C]
pre_mel = pre_mel * input_masks
return pre_mel
@@ -0,0 +1,751 @@
import math
import torch
from torch import nn
from torch.nn import Parameter, Linear
from .layers import LayerNorm, Embedding
from ...utils.nn.seq_utils import (
get_incremental_state,
set_incremental_state,
softmax,
make_positions,
)
import torch.nn.functional as F
DEFAULT_MAX_SOURCE_POSITIONS = 2000
DEFAULT_MAX_TARGET_POSITIONS = 2000
class SinusoidalPositionalEmbedding(nn.Module):
"""This module produces sinusoidal positional embeddings of any length.
Padding symbols are ignored.
"""
def __init__(self, embedding_dim, padding_idx, init_size=1024):
super().__init__()
self.embedding_dim = embedding_dim
self.padding_idx = padding_idx
self.weights = SinusoidalPositionalEmbedding.get_embedding(
init_size,
embedding_dim,
padding_idx,
)
self.register_buffer('_float_tensor', torch.FloatTensor(1))
@staticmethod
def get_embedding(num_embeddings, embedding_dim, padding_idx=None):
"""Build sinusoidal embeddings.
This matches the implementation in tensor2tensor, but differs slightly
from the description in Section 3.5 of "Attention Is All You Need".
"""
half_dim = embedding_dim // 2
emb = math.log(10000) / (half_dim - 1)
emb = torch.exp(torch.arange(half_dim, dtype=torch.float) * -emb)
emb = torch.arange(num_embeddings, dtype=torch.float).unsqueeze(1) * emb.unsqueeze(0)
emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=1).view(num_embeddings, -1)
if embedding_dim % 2 == 1:
# zero pad
emb = torch.cat([emb, torch.zeros(num_embeddings, 1)], dim=1)
if padding_idx is not None:
emb[padding_idx, :] = 0
return emb
def forward(self, input, incremental_state=None, timestep=None, positions=None, **kwargs):
"""Input is expected to be of size [bsz x seqlen]."""
bsz, seq_len = input.shape[:2]
max_pos = self.padding_idx + 1 + seq_len
if self.weights is None or max_pos > self.weights.size(0):
# recompute/expand embeddings if needed
self.weights = SinusoidalPositionalEmbedding.get_embedding(
max_pos,
self.embedding_dim,
self.padding_idx,
)
self.weights = self.weights.to(self._float_tensor)
if incremental_state is not None:
# positions is the same for every token when decoding a single step
pos = timestep.view(-1)[0] + 1 if timestep is not None else seq_len
return self.weights[self.padding_idx + pos, :].expand(bsz, 1, -1)
positions = make_positions(input, self.padding_idx) if positions is None else positions
return self.weights.index_select(0, positions.view(-1)).view(bsz, seq_len, -1).detach()
def max_positions(self):
"""Maximum number of supported positions."""
return int(1e5) # an arbitrary large number
class TransformerFFNLayer(nn.Module):
def __init__(self, hidden_size, filter_size, padding="SAME", kernel_size=1, dropout=0., act='gelu'):
super().__init__()
self.kernel_size = kernel_size
self.dropout = dropout
self.act = act
if padding == 'SAME':
self.ffn_1 = nn.Conv1d(hidden_size, filter_size, kernel_size, padding=kernel_size // 2)
elif padding == 'LEFT':
self.ffn_1 = nn.Sequential(
nn.ConstantPad1d((kernel_size - 1, 0), 0.0),
nn.Conv1d(hidden_size, filter_size, kernel_size)
)
self.ffn_2 = Linear(filter_size, hidden_size)
def forward(self, x, incremental_state=None):
# x: T x B x C
if incremental_state is not None:
saved_state = self._get_input_buffer(incremental_state)
if 'prev_input' in saved_state:
prev_input = saved_state['prev_input']
x = torch.cat((prev_input, x), dim=0)
x = x[-self.kernel_size:]
saved_state['prev_input'] = x
self._set_input_buffer(incremental_state, saved_state)
x = self.ffn_1(x.permute(1, 2, 0)).permute(2, 0, 1)
x = x * self.kernel_size ** -0.5
if incremental_state is not None:
x = x[-1:]
if self.act == 'gelu':
x = F.gelu(x)
if self.act == 'relu':
x = F.relu(x)
x = F.dropout(x, self.dropout, training=self.training)
x = self.ffn_2(x)
return x
def _get_input_buffer(self, incremental_state):
return get_incremental_state(
self,
incremental_state,
'f',
) or {}
def _set_input_buffer(self, incremental_state, buffer):
set_incremental_state(
self,
incremental_state,
'f',
buffer,
)
def clear_buffer(self, incremental_state):
if incremental_state is not None:
saved_state = self._get_input_buffer(incremental_state)
if 'prev_input' in saved_state:
del saved_state['prev_input']
self._set_input_buffer(incremental_state, saved_state)
class MultiheadAttention(nn.Module):
def __init__(self, embed_dim, num_heads, kdim=None, vdim=None, dropout=0., bias=True,
add_bias_kv=False, add_zero_attn=False, self_attention=False,
encoder_decoder_attention=False):
super().__init__()
self.embed_dim = embed_dim
self.kdim = kdim if kdim is not None else embed_dim
self.vdim = vdim if vdim is not None else embed_dim
self.qkv_same_dim = self.kdim == embed_dim and self.vdim == embed_dim
self.num_heads = num_heads
self.dropout = dropout
self.head_dim = embed_dim // num_heads
assert self.head_dim * num_heads == self.embed_dim, "embed_dim must be divisible by num_heads"
self.scaling = self.head_dim ** -0.5
self.self_attention = self_attention
self.encoder_decoder_attention = encoder_decoder_attention
assert not self.self_attention or self.qkv_same_dim, 'Self-attention requires query, key and ' \
'value to be of the same size'
if self.qkv_same_dim:
self.in_proj_weight = Parameter(torch.Tensor(3 * embed_dim, embed_dim))
else:
self.k_proj_weight = Parameter(torch.Tensor(embed_dim, self.kdim))
self.v_proj_weight = Parameter(torch.Tensor(embed_dim, self.vdim))
self.q_proj_weight = Parameter(torch.Tensor(embed_dim, embed_dim))
if bias:
self.in_proj_bias = Parameter(torch.Tensor(3 * embed_dim))
else:
self.register_parameter('in_proj_bias', None)
self.out_proj = nn.Linear(embed_dim, embed_dim, bias=bias)
if add_bias_kv:
self.bias_k = Parameter(torch.Tensor(1, 1, embed_dim))
self.bias_v = Parameter(torch.Tensor(1, 1, embed_dim))
else:
self.bias_k = self.bias_v = None
self.add_zero_attn = add_zero_attn
self.reset_parameters()
self.enable_torch_version = False
if hasattr(F, "multi_head_attention_forward"):
self.enable_torch_version = True
else:
self.enable_torch_version = False
self.last_attn_probs = None
def reset_parameters(self):
if self.qkv_same_dim:
nn.init.xavier_uniform_(self.in_proj_weight)
else:
nn.init.xavier_uniform_(self.k_proj_weight)
nn.init.xavier_uniform_(self.v_proj_weight)
nn.init.xavier_uniform_(self.q_proj_weight)
nn.init.xavier_uniform_(self.out_proj.weight)
if self.in_proj_bias is not None:
nn.init.constant_(self.in_proj_bias, 0.)
nn.init.constant_(self.out_proj.bias, 0.)
if self.bias_k is not None:
nn.init.xavier_normal_(self.bias_k)
if self.bias_v is not None:
nn.init.xavier_normal_(self.bias_v)
def forward(
self,
query, key, value,
key_padding_mask=None,
incremental_state=None,
need_weights=True,
static_kv=False,
attn_mask=None,
before_softmax=False,
need_head_weights=False,
enc_dec_attn_constraint_mask=None,
reset_attn_weight=None
):
"""Input shape: Time x Batch x Channel
Args:
key_padding_mask (ByteTensor, optional): mask to exclude
keys that are pads, of shape `(batch, src_len)`, where
padding elements are indicated by 1s.
need_weights (bool, optional): return the attention weights,
averaged over heads (default: False).
attn_mask (ByteTensor, optional): typically used to
implement causal attention, where the mask prevents the
attention from looking forward in time (default: None).
before_softmax (bool, optional): return the raw attention
weights and values before the attention softmax.
need_head_weights (bool, optional): return the attention
weights for each head. Implies *need_weights*. Default:
return the average attention weights over all heads.
"""
if need_head_weights:
need_weights = True
tgt_len, bsz, embed_dim = query.size()
assert embed_dim == self.embed_dim
assert list(query.size()) == [tgt_len, bsz, embed_dim]
if self.enable_torch_version and incremental_state is None and not static_kv and reset_attn_weight is None:
if self.qkv_same_dim:
return F.multi_head_attention_forward(query, key, value,
self.embed_dim, self.num_heads,
self.in_proj_weight,
self.in_proj_bias, self.bias_k, self.bias_v,
self.add_zero_attn, self.dropout,
self.out_proj.weight, self.out_proj.bias,
self.training, key_padding_mask, need_weights,
attn_mask)
else:
return F.multi_head_attention_forward(query, key, value,
self.embed_dim, self.num_heads,
torch.empty([0]),
self.in_proj_bias, self.bias_k, self.bias_v,
self.add_zero_attn, self.dropout,
self.out_proj.weight, self.out_proj.bias,
self.training, key_padding_mask, need_weights,
attn_mask, use_separate_proj_weight=True,
q_proj_weight=self.q_proj_weight,
k_proj_weight=self.k_proj_weight,
v_proj_weight=self.v_proj_weight)
if incremental_state is not None:
saved_state = self._get_input_buffer(incremental_state)
if 'prev_key' in saved_state:
# previous time steps are cached - no need to recompute
# key and value if they are static
if static_kv:
assert self.encoder_decoder_attention and not self.self_attention
key = value = None
else:
saved_state = None
if self.self_attention:
# self-attention
q, k, v = self.in_proj_qkv(query)
elif self.encoder_decoder_attention:
# encoder-decoder attention
q = self.in_proj_q(query)
if key is None:
assert value is None
k = v = None
else:
k = self.in_proj_k(key)
v = self.in_proj_v(key)
else:
q = self.in_proj_q(query)
k = self.in_proj_k(key)
v = self.in_proj_v(value)
q *= self.scaling
if self.bias_k is not None:
assert self.bias_v is not None
k = torch.cat([k, self.bias_k.repeat(1, bsz, 1)])
v = torch.cat([v, self.bias_v.repeat(1, bsz, 1)])
if attn_mask is not None:
attn_mask = torch.cat([attn_mask, attn_mask.new_zeros(attn_mask.size(0), 1)], dim=1)
if key_padding_mask is not None:
key_padding_mask = torch.cat(
[key_padding_mask, key_padding_mask.new_zeros(key_padding_mask.size(0), 1)], dim=1)
q = q.contiguous().view(tgt_len, bsz * self.num_heads, self.head_dim).transpose(0, 1)
if k is not None:
k = k.contiguous().view(-1, bsz * self.num_heads, self.head_dim).transpose(0, 1)
if v is not None:
v = v.contiguous().view(-1, bsz * self.num_heads, self.head_dim).transpose(0, 1)
if saved_state is not None:
# saved states are stored with shape (bsz, num_heads, seq_len, head_dim)
if 'prev_key' in saved_state:
prev_key = saved_state['prev_key'].view(bsz * self.num_heads, -1, self.head_dim)
if static_kv:
k = prev_key
else:
k = torch.cat((prev_key, k), dim=1)
if 'prev_value' in saved_state:
prev_value = saved_state['prev_value'].view(bsz * self.num_heads, -1, self.head_dim)
if static_kv:
v = prev_value
else:
v = torch.cat((prev_value, v), dim=1)
if 'prev_key_padding_mask' in saved_state and saved_state['prev_key_padding_mask'] is not None:
prev_key_padding_mask = saved_state['prev_key_padding_mask']
if static_kv:
key_padding_mask = prev_key_padding_mask
else:
key_padding_mask = torch.cat((prev_key_padding_mask, key_padding_mask), dim=1)
saved_state['prev_key'] = k.view(bsz, self.num_heads, -1, self.head_dim)
saved_state['prev_value'] = v.view(bsz, self.num_heads, -1, self.head_dim)
saved_state['prev_key_padding_mask'] = key_padding_mask
self._set_input_buffer(incremental_state, saved_state)
src_len = k.size(1)
# This is part of a workaround to get around fork/join parallelism
# not supporting Optional types.
if key_padding_mask is not None and key_padding_mask.shape == torch.Size([]):
key_padding_mask = None
if key_padding_mask is not None:
assert key_padding_mask.size(0) == bsz
assert key_padding_mask.size(1) == src_len
if self.add_zero_attn:
src_len += 1
k = torch.cat([k, k.new_zeros((k.size(0), 1) + k.size()[2:])], dim=1)
v = torch.cat([v, v.new_zeros((v.size(0), 1) + v.size()[2:])], dim=1)
if attn_mask is not None:
attn_mask = torch.cat([attn_mask, attn_mask.new_zeros(attn_mask.size(0), 1)], dim=1)
if key_padding_mask is not None:
key_padding_mask = torch.cat(
[key_padding_mask, torch.zeros(key_padding_mask.size(0), 1).type_as(key_padding_mask)], dim=1)
attn_weights = torch.bmm(q, k.transpose(1, 2))
attn_weights = self.apply_sparse_mask(attn_weights, tgt_len, src_len, bsz)
assert list(attn_weights.size()) == [bsz * self.num_heads, tgt_len, src_len]
if attn_mask is not None:
if len(attn_mask.shape) == 2:
attn_mask = attn_mask.unsqueeze(0)
elif len(attn_mask.shape) == 3:
attn_mask = attn_mask[:, None].repeat([1, self.num_heads, 1, 1]).reshape(
bsz * self.num_heads, tgt_len, src_len)
attn_weights = attn_weights + attn_mask
if enc_dec_attn_constraint_mask is not None: # bs x head x L_kv
attn_weights = attn_weights.view(bsz, self.num_heads, tgt_len, src_len)
attn_weights = attn_weights.masked_fill(
enc_dec_attn_constraint_mask.unsqueeze(2).bool(),
-1e8,
)
attn_weights = attn_weights.view(bsz * self.num_heads, tgt_len, src_len)
if key_padding_mask is not None:
# don't attend to padding symbols
attn_weights = attn_weights.view(bsz, self.num_heads, tgt_len, src_len)
attn_weights = attn_weights.masked_fill(
key_padding_mask.unsqueeze(1).unsqueeze(2),
-1e8,
)
attn_weights = attn_weights.view(bsz * self.num_heads, tgt_len, src_len)
attn_logits = attn_weights.view(bsz, self.num_heads, tgt_len, src_len)
if before_softmax:
return attn_weights, v
attn_weights_float = softmax(attn_weights, dim=-1)
attn_weights = attn_weights_float.type_as(attn_weights)
attn_probs = F.dropout(attn_weights_float.type_as(attn_weights), p=self.dropout, training=self.training)
if reset_attn_weight is not None:
if reset_attn_weight:
self.last_attn_probs = attn_probs.detach()
else:
assert self.last_attn_probs is not None
attn_probs = self.last_attn_probs
attn = torch.bmm(attn_probs, v)
assert list(attn.size()) == [bsz * self.num_heads, tgt_len, self.head_dim]
attn = attn.transpose(0, 1).contiguous().view(tgt_len, bsz, embed_dim)
attn = self.out_proj(attn)
if need_weights:
attn_weights = attn_weights_float.view(bsz, self.num_heads, tgt_len, src_len).transpose(1, 0)
if not need_head_weights:
# average attention weights over heads
attn_weights = attn_weights.mean(dim=0)
else:
attn_weights = None
return attn, (attn_weights, attn_logits)
def in_proj_qkv(self, query):
return self._in_proj(query).chunk(3, dim=-1)
def in_proj_q(self, query):
if self.qkv_same_dim:
return self._in_proj(query, end=self.embed_dim)
else:
bias = self.in_proj_bias
if bias is not None:
bias = bias[:self.embed_dim]
return F.linear(query, self.q_proj_weight, bias)
def in_proj_k(self, key):
if self.qkv_same_dim:
return self._in_proj(key, start=self.embed_dim, end=2 * self.embed_dim)
else:
weight = self.k_proj_weight
bias = self.in_proj_bias
if bias is not None:
bias = bias[self.embed_dim:2 * self.embed_dim]
return F.linear(key, weight, bias)
def in_proj_v(self, value):
if self.qkv_same_dim:
return self._in_proj(value, start=2 * self.embed_dim)
else:
weight = self.v_proj_weight
bias = self.in_proj_bias
if bias is not None:
bias = bias[2 * self.embed_dim:]
return F.linear(value, weight, bias)
def _in_proj(self, input, start=0, end=None):
weight = self.in_proj_weight
bias = self.in_proj_bias
weight = weight[start:end, :]
if bias is not None:
bias = bias[start:end]
return F.linear(input, weight, bias)
def _get_input_buffer(self, incremental_state):
return get_incremental_state(
self,
incremental_state,
'attn_state',
) or {}
def _set_input_buffer(self, incremental_state, buffer):
set_incremental_state(
self,
incremental_state,
'attn_state',
buffer,
)
def apply_sparse_mask(self, attn_weights, tgt_len, src_len, bsz):
return attn_weights
def clear_buffer(self, incremental_state=None):
if incremental_state is not None:
saved_state = self._get_input_buffer(incremental_state)
if 'prev_key' in saved_state:
del saved_state['prev_key']
if 'prev_value' in saved_state:
del saved_state['prev_value']
self._set_input_buffer(incremental_state, saved_state)
class EncSALayer(nn.Module):
def __init__(self, c, num_heads, dropout, attention_dropout=0.1,
relu_dropout=0.1, kernel_size=9, padding='SAME', act='gelu'):
super().__init__()
self.c = c
self.dropout = dropout
self.num_heads = num_heads
if num_heads > 0:
self.layer_norm1 = LayerNorm(c)
self.self_attn = MultiheadAttention(
self.c, num_heads, self_attention=True, dropout=attention_dropout, bias=False)
self.layer_norm2 = LayerNorm(c)
self.ffn = TransformerFFNLayer(
c, 4 * c, kernel_size=kernel_size, dropout=relu_dropout, padding=padding, act=act)
def forward(self, x, encoder_padding_mask=None, **kwargs):
layer_norm_training = kwargs.get('layer_norm_training', None)
if layer_norm_training is not None:
self.layer_norm1.training = layer_norm_training
self.layer_norm2.training = layer_norm_training
if self.num_heads > 0:
residual = x
x = self.layer_norm1(x)
x, _, = self.self_attn(
query=x,
key=x,
value=x,
key_padding_mask=encoder_padding_mask
)
x = F.dropout(x, self.dropout, training=self.training)
x = residual + x
x = x * (1 - encoder_padding_mask.float()).transpose(0, 1)[..., None]
residual = x
x = self.layer_norm2(x)
x = self.ffn(x)
x = F.dropout(x, self.dropout, training=self.training)
x = residual + x
x = x * (1 - encoder_padding_mask.float()).transpose(0, 1)[..., None]
return x
class DecSALayer(nn.Module):
def __init__(self, c, num_heads, dropout, attention_dropout=0.1, relu_dropout=0.1,
kernel_size=9, act='gelu'):
super().__init__()
self.c = c
self.dropout = dropout
self.layer_norm1 = LayerNorm(c)
self.self_attn = MultiheadAttention(
c, num_heads, self_attention=True, dropout=attention_dropout, bias=False
)
self.layer_norm2 = LayerNorm(c)
self.encoder_attn = MultiheadAttention(
c, num_heads, encoder_decoder_attention=True, dropout=attention_dropout, bias=False,
)
self.layer_norm3 = LayerNorm(c)
self.ffn = TransformerFFNLayer(
c, 4 * c, padding='LEFT', kernel_size=kernel_size, dropout=relu_dropout, act=act)
def forward(
self,
x,
encoder_out=None,
encoder_padding_mask=None,
incremental_state=None,
self_attn_mask=None,
self_attn_padding_mask=None,
attn_out=None,
reset_attn_weight=None,
**kwargs,
):
layer_norm_training = kwargs.get('layer_norm_training', None)
if layer_norm_training is not None:
self.layer_norm1.training = layer_norm_training
self.layer_norm2.training = layer_norm_training
self.layer_norm3.training = layer_norm_training
residual = x
x = self.layer_norm1(x)
x, _ = self.self_attn(
query=x,
key=x,
value=x,
key_padding_mask=self_attn_padding_mask,
incremental_state=incremental_state,
attn_mask=self_attn_mask
)
x = F.dropout(x, self.dropout, training=self.training)
x = residual + x
attn_logits = None
if encoder_out is not None or attn_out is not None:
residual = x
x = self.layer_norm2(x)
if encoder_out is not None:
x, attn = self.encoder_attn(
query=x,
key=encoder_out,
value=encoder_out,
key_padding_mask=encoder_padding_mask,
incremental_state=incremental_state,
static_kv=True,
enc_dec_attn_constraint_mask=get_incremental_state(self, incremental_state,
'enc_dec_attn_constraint_mask'),
reset_attn_weight=reset_attn_weight
)
attn_logits = attn[1]
elif attn_out is not None:
x = self.encoder_attn.in_proj_v(attn_out)
if encoder_out is not None or attn_out is not None:
x = F.dropout(x, self.dropout, training=self.training)
x = residual + x
residual = x
x = self.layer_norm3(x)
x = self.ffn(x, incremental_state=incremental_state)
x = F.dropout(x, self.dropout, training=self.training)
x = residual + x
return x, attn_logits
def clear_buffer(self, input, encoder_out=None, encoder_padding_mask=None, incremental_state=None):
self.encoder_attn.clear_buffer(incremental_state)
self.ffn.clear_buffer(incremental_state)
def set_buffer(self, name, tensor, incremental_state):
return set_incremental_state(self, incremental_state, name, tensor)
class TransformerEncoderLayer(nn.Module):
def __init__(self, hidden_size, dropout, kernel_size=9, num_heads=2):
super().__init__()
self.hidden_size = hidden_size
self.dropout = dropout
self.num_heads = num_heads
self.op = EncSALayer(
hidden_size, num_heads, dropout=dropout,
attention_dropout=0.0, relu_dropout=dropout,
kernel_size=kernel_size)
def forward(self, x, **kwargs):
return self.op(x, **kwargs)
class TransformerDecoderLayer(nn.Module):
def __init__(self, hidden_size, dropout, kernel_size=9, num_heads=2):
super().__init__()
self.hidden_size = hidden_size
self.dropout = dropout
self.num_heads = num_heads
self.op = DecSALayer(
hidden_size, num_heads, dropout=dropout,
attention_dropout=0.0, relu_dropout=dropout,
kernel_size=kernel_size)
def forward(self, x, **kwargs):
return self.op(x, **kwargs)
def clear_buffer(self, *args):
return self.op.clear_buffer(*args)
def set_buffer(self, *args):
return self.op.set_buffer(*args)
class FFTBlocks(nn.Module):
def __init__(self, hidden_size, num_layers, ffn_kernel_size=9, dropout=0.0,
num_heads=2, use_pos_embed=True, use_last_norm=True,
use_pos_embed_alpha=True):
super().__init__()
self.num_layers = num_layers
embed_dim = self.hidden_size = hidden_size
self.dropout = dropout
self.use_pos_embed = use_pos_embed
self.use_last_norm = use_last_norm
if use_pos_embed:
self.max_source_positions = DEFAULT_MAX_TARGET_POSITIONS
self.padding_idx = 0
self.pos_embed_alpha = nn.Parameter(torch.Tensor([1])) if use_pos_embed_alpha else 1
self.embed_positions = SinusoidalPositionalEmbedding(
embed_dim, self.padding_idx, init_size=DEFAULT_MAX_TARGET_POSITIONS,
)
self.layers = nn.ModuleList([])
self.layers.extend([
TransformerEncoderLayer(self.hidden_size, self.dropout,
kernel_size=ffn_kernel_size, num_heads=num_heads)
for _ in range(self.num_layers)
])
if self.use_last_norm:
self.layer_norm = nn.LayerNorm(embed_dim)
else:
self.layer_norm = None
def forward(self, x, padding_mask=None, attn_mask=None, return_hiddens=False):
"""
:param x: [B, T, C]
:param padding_mask: [B, T]
:return: [B, T, C] or [L, B, T, C]
"""
padding_mask = x.abs().sum(-1).eq(0).data if padding_mask is None else padding_mask
nonpadding_mask_TB = 1 - padding_mask.transpose(0, 1).float()[:, :, None] # [T, B, 1]
if self.use_pos_embed:
positions = self.pos_embed_alpha * self.embed_positions(x[..., 0])
x = x + positions
x = F.dropout(x, p=self.dropout, training=self.training)
# B x T x C -> T x B x C
x = x.transpose(0, 1) * nonpadding_mask_TB
hiddens = []
for layer in self.layers:
x = layer(x, encoder_padding_mask=padding_mask, attn_mask=attn_mask) * nonpadding_mask_TB
hiddens.append(x)
if self.use_last_norm:
x = self.layer_norm(x) * nonpadding_mask_TB
if return_hiddens:
x = torch.stack(hiddens, 0) # [L, T, B, C]
x = x.transpose(1, 2) # [L, B, T, C]
else:
x = x.transpose(0, 1) # [B, T, C]
return x
class FastSpeechEncoder(FFTBlocks):
def __init__(self, dict_size, hidden_size=256, num_layers=4, kernel_size=9, num_heads=2,
dropout=0.0):
super().__init__(hidden_size, num_layers, kernel_size, num_heads=num_heads,
use_pos_embed=False, dropout=dropout) # use_pos_embed_alpha for compatibility
self.embed_tokens = Embedding(dict_size, hidden_size, 0)
self.embed_scale = math.sqrt(hidden_size)
self.padding_idx = 0
self.embed_positions = SinusoidalPositionalEmbedding(
hidden_size, self.padding_idx, init_size=DEFAULT_MAX_TARGET_POSITIONS,
)
def forward(self, txt_tokens, attn_mask=None):
"""
:param txt_tokens: [B, T]
:return: {
'encoder_out': [B x T x C]
}
"""
encoder_padding_mask = txt_tokens.eq(self.padding_idx).data
x = self.forward_embedding(txt_tokens) # [B, T, H]
if self.num_layers > 0:
x = super(FastSpeechEncoder, self).forward(x, encoder_padding_mask, attn_mask=attn_mask)
return x
def forward_embedding(self, txt_tokens):
# embed tokens and positions
x = self.embed_scale * self.embed_tokens(txt_tokens)
positions = self.embed_positions(txt_tokens)
x = x + positions
x = F.dropout(x, p=self.dropout, training=self.training)
return x
class FastSpeechDecoder(FFTBlocks):
def __init__(self, hidden_size=256, num_layers=4, kernel_size=9, num_heads=2):
super().__init__(hidden_size, num_layers, kernel_size, num_heads=num_heads)
@@ -0,0 +1,109 @@
import torch
from torch import nn
from packaging import version
def fused_add_tanh_sigmoid_multiply(input_a, input_b, n_channels):
n_channels_int = n_channels[0]
in_act = input_a + input_b
t_act = torch.tanh(in_act[:, :n_channels_int, :])
s_act = torch.sigmoid(in_act[:, n_channels_int:, :])
acts = t_act * s_act
return acts
jit_fused_add_tanh_sigmoid_multiply = fused_add_tanh_sigmoid_multiply
def script_function():
if version.parse(torch.__version__) >= version.parse('2.0'):
global jit_fused_add_tanh_sigmoid_multiply
jit_fused_add_tanh_sigmoid_multiply = torch.jit.script(fused_add_tanh_sigmoid_multiply)
class WN(torch.nn.Module):
def __init__(self, hidden_size, kernel_size, dilation_rate, n_layers, c_cond=0,
p_dropout=0, share_cond_layers=False, is_BTC=False):
super(WN, self).__init__()
assert (kernel_size % 2 == 1)
assert (hidden_size % 2 == 0)
self.is_BTC = is_BTC
self.hidden_size = hidden_size
self.kernel_size = kernel_size
self.dilation_rate = dilation_rate
self.n_layers = n_layers
self.gin_channels = c_cond
self.p_dropout = p_dropout
self.share_cond_layers = share_cond_layers
self.in_layers = torch.nn.ModuleList()
self.res_skip_layers = torch.nn.ModuleList()
self.drop = nn.Dropout(p_dropout)
if c_cond != 0 and not share_cond_layers:
cond_layer = torch.nn.Conv1d(c_cond, 2 * hidden_size * n_layers, 1)
self.cond_layer = torch.nn.utils.weight_norm(cond_layer, name='weight')
for i in range(n_layers):
dilation = dilation_rate ** i
padding = int((kernel_size * dilation - dilation) / 2)
in_layer = torch.nn.Conv1d(hidden_size, 2 * hidden_size, kernel_size,
dilation=dilation, padding=padding)
in_layer = torch.nn.utils.weight_norm(in_layer, name='weight')
self.in_layers.append(in_layer)
# last one is not necessary
if i < n_layers - 1:
res_skip_channels = 2 * hidden_size
else:
res_skip_channels = hidden_size
res_skip_layer = torch.nn.Conv1d(hidden_size, res_skip_channels, 1)
res_skip_layer = torch.nn.utils.weight_norm(res_skip_layer, name='weight')
self.res_skip_layers.append(res_skip_layer)
script_function()
def forward(self, x, nonpadding=None, cond=None):
if self.is_BTC:
x = x.transpose(1, 2)
cond = cond.transpose(1, 2) if cond is not None else None
nonpadding = nonpadding.transpose(1, 2) if nonpadding is not None else None
if nonpadding is None:
nonpadding = 1
output = torch.zeros_like(x)
n_channels_tensor = torch.IntTensor([self.hidden_size])
if cond is not None and not self.share_cond_layers:
cond = self.cond_layer(cond)
for i in range(self.n_layers):
x_in = self.in_layers[i](x)
x_in = self.drop(x_in)
if cond is not None:
cond_offset = i * 2 * self.hidden_size
cond_l = cond[:, cond_offset:cond_offset + 2 * self.hidden_size, :]
else:
cond_l = torch.zeros_like(x_in)
if version.parse(torch.__version__) >= version.parse('2.0'):
acts = jit_fused_add_tanh_sigmoid_multiply(x_in, cond_l, n_channels_tensor)
else:
acts = fused_add_tanh_sigmoid_multiply(x_in, cond_l, n_channels_tensor)
res_skip_acts = self.res_skip_layers[i](acts)
if i < self.n_layers - 1:
x = (x + res_skip_acts[:, :self.hidden_size, :]) * nonpadding
output = output + res_skip_acts[:, self.hidden_size:, :]
else:
output = output + res_skip_acts
output = output * nonpadding
if self.is_BTC:
output = output.transpose(1, 2)
return output
def remove_weight_norm(self):
def remove_weight_norm(m):
try:
nn.utils.remove_weight_norm(m)
except ValueError: # this module didn't have weight norm
return
self.apply(remove_weight_norm)
@@ -0,0 +1 @@
"""Pitch extractor modules for ROSVOT."""
@@ -0,0 +1,6 @@
from .constants import *
from .model import E2E0
from .utils import to_local_average_f0, to_viterbi_f0
from .inference import RMVPE
from .spec import MelSpectrogram
from .extractor import extract
@@ -0,0 +1,9 @@
SAMPLE_RATE = 16000
N_CLASS = 360
N_MELS = 128
MEL_FMIN = 30
MEL_FMAX = 8000
WINDOW_LENGTH = 1024
CONST = 1997.3794084376191
@@ -0,0 +1,173 @@
import torch
import torch.nn as nn
from .constants import N_MELS
class ConvBlockRes(nn.Module):
def __init__(self, in_channels, out_channels, momentum=0.01):
super(ConvBlockRes, self).__init__()
self.conv = nn.Sequential(
nn.Conv2d(in_channels=in_channels,
out_channels=out_channels,
kernel_size=(3, 3),
stride=(1, 1),
padding=(1, 1),
bias=False),
nn.BatchNorm2d(out_channels, momentum=momentum),
nn.ReLU(),
nn.Conv2d(in_channels=out_channels,
out_channels=out_channels,
kernel_size=(3, 3),
stride=(1, 1),
padding=(1, 1),
bias=False),
nn.BatchNorm2d(out_channels, momentum=momentum),
nn.ReLU(),
)
if in_channels != out_channels:
self.shortcut = nn.Conv2d(in_channels, out_channels, (1, 1))
self.is_shortcut = True
else:
self.is_shortcut = False
def forward(self, x):
if self.is_shortcut:
return self.conv(x) + self.shortcut(x)
else:
return self.conv(x) + x
class ResEncoderBlock(nn.Module):
def __init__(self, in_channels, out_channels, kernel_size, n_blocks=1, momentum=0.01):
super(ResEncoderBlock, self).__init__()
self.n_blocks = n_blocks
self.conv = nn.ModuleList()
self.conv.append(ConvBlockRes(in_channels, out_channels, momentum))
for i in range(n_blocks - 1):
self.conv.append(ConvBlockRes(out_channels, out_channels, momentum))
self.kernel_size = kernel_size
if self.kernel_size is not None:
self.pool = nn.AvgPool2d(kernel_size=kernel_size)
def forward(self, x):
for i in range(self.n_blocks):
x = self.conv[i](x)
if self.kernel_size is not None:
return x, self.pool(x)
else:
return x
class ResDecoderBlock(nn.Module):
def __init__(self, in_channels, out_channels, stride, n_blocks=1, momentum=0.01):
super(ResDecoderBlock, self).__init__()
out_padding = (0, 1) if stride == (1, 2) else (1, 1)
self.n_blocks = n_blocks
self.conv1 = nn.Sequential(
nn.ConvTranspose2d(in_channels=in_channels,
out_channels=out_channels,
kernel_size=(3, 3),
stride=stride,
padding=(1, 1),
output_padding=out_padding,
bias=False),
nn.BatchNorm2d(out_channels, momentum=momentum),
nn.ReLU(),
)
self.conv2 = nn.ModuleList()
self.conv2.append(ConvBlockRes(out_channels * 2, out_channels, momentum))
for i in range(n_blocks-1):
self.conv2.append(ConvBlockRes(out_channels, out_channels, momentum))
def forward(self, x, concat_tensor):
x = self.conv1(x)
x = torch.cat((x, concat_tensor), dim=1)
for i in range(self.n_blocks):
x = self.conv2[i](x)
return x
class Encoder(nn.Module):
def __init__(self, in_channels, in_size, n_encoders, kernel_size, n_blocks, out_channels=16, momentum=0.01):
super(Encoder, self).__init__()
self.n_encoders = n_encoders
self.bn = nn.BatchNorm2d(in_channels, momentum=momentum)
self.layers = nn.ModuleList()
self.latent_channels = []
for i in range(self.n_encoders):
self.layers.append(ResEncoderBlock(in_channels, out_channels, kernel_size, n_blocks, momentum=momentum))
self.latent_channels.append([out_channels, in_size])
in_channels = out_channels
out_channels *= 2
in_size //= 2
self.out_size = in_size
self.out_channel = out_channels
def forward(self, x):
concat_tensors = []
x = self.bn(x)
for i in range(self.n_encoders):
_, x = self.layers[i](x)
concat_tensors.append(_)
return x, concat_tensors
class Intermediate(nn.Module):
def __init__(self, in_channels, out_channels, n_inters, n_blocks, momentum=0.01):
super(Intermediate, self).__init__()
self.n_inters = n_inters
self.layers = nn.ModuleList()
self.layers.append(ResEncoderBlock(in_channels, out_channels, None, n_blocks, momentum))
for i in range(self.n_inters-1):
self.layers.append(ResEncoderBlock(out_channels, out_channels, None, n_blocks, momentum))
def forward(self, x):
for i in range(self.n_inters):
x = self.layers[i](x)
return x
class Decoder(nn.Module):
def __init__(self, in_channels, n_decoders, stride, n_blocks, momentum=0.01):
super(Decoder, self).__init__()
self.layers = nn.ModuleList()
self.n_decoders = n_decoders
for i in range(self.n_decoders):
out_channels = in_channels // 2
self.layers.append(ResDecoderBlock(in_channels, out_channels, stride, n_blocks, momentum))
in_channels = out_channels
def forward(self, x, concat_tensors):
for i in range(self.n_decoders):
x = self.layers[i](x, concat_tensors[-1-i])
return x
class TimbreFilter(nn.Module):
def __init__(self, latent_rep_channels):
super(TimbreFilter, self).__init__()
self.layers = nn.ModuleList()
for latent_rep in latent_rep_channels:
self.layers.append(ConvBlockRes(latent_rep[0], latent_rep[0]))
def forward(self, x_tensors):
out_tensors = []
for i, layer in enumerate(self.layers):
out_tensors.append(layer(x_tensors[i]))
return out_tensors
class DeepUnet0(nn.Module):
def __init__(self, kernel_size, n_blocks, en_de_layers=5, inter_layers=4, in_channels=1, en_out_channels=16):
super(DeepUnet0, self).__init__()
self.encoder = Encoder(in_channels, N_MELS, en_de_layers, kernel_size, n_blocks, en_out_channels)
self.intermediate = Intermediate(self.encoder.out_channel // 2, self.encoder.out_channel, inter_layers, n_blocks)
self.tf = TimbreFilter(self.encoder.latent_channels)
self.decoder = Decoder(self.encoder.out_channel, en_de_layers, kernel_size, n_blocks)
def forward(self, x):
x, concat_tensors = self.encoder(x)
x = self.intermediate(x)
x = self.decoder(x, concat_tensors)
return x
@@ -0,0 +1,183 @@
import math
import os
from tqdm import tqdm
import librosa
import numpy as np
import torch
from torch.utils.data import Dataset, DataLoader, DistributedSampler
import torch.multiprocessing as mp
from torch.distributed import init_process_group
import torch.distributed as dist
from .inference import RMVPE
from ....utils.commons.dataset_utils import batch_by_size, build_dataloader
# import utils
from ....utils.audio import get_wav_num_frames
"""
A convenient API for batch inference
update: add ddp
"""
class RMVPEInferDataset(Dataset):
def __init__(self, wav_fns: list, id_and_sizes=None, sr=24000, hop_size=128, num_workers=0):
if id_and_sizes is None:
id_and_sizes = []
if type(wav_fns[0]) == str: # wav_paths
for idx, wav_path in enumerate(wav_fns):
total_frames = get_wav_num_frames(wav_path, sr)
id_and_sizes.append((idx, round(total_frames / hop_size)))
else: # numpy arrays, mono wavs
for idx, wav in enumerate(wav_fns):
id_and_sizes.append((idx, round(wav.shape[-1] / hop_size)))
self.wav_fns = wav_fns
self.id_and_sizes = id_and_sizes
self.sr = sr
self.num_workers = num_workers
def __getitem__(self, idx):
if type(self.wav_fns[idx]) == str:
wav_fn = self.wav_fns[idx]
wav, _ = librosa.core.load(wav_fn, sr=self.sr)
else:
wav = self.wav_fns[idx]
return idx, wav
def collater(self, samples: list):
return samples
def __len__(self):
return len(self.wav_fns)
def ordered_indices(self):
"""Return an ordered list of indices. Batches will be constructed based
on this order."""
return np.arange(len(self))
def num_tokens(self, index):
return self.id_and_sizes[index][1]
@torch.no_grad()
def extract(wav_fns: list, id_and_sizes=None, ckpt=None, sr=24000, hop_size=128, bsz=128, max_tokens=100000,
fmax=900, fmin=50, ds_workers=0):
all_gpu_ids = [int(x) for x in os.environ.get("CUDA_VISIBLE_DEVICES", "").split(",") if x != '']
num_gpus = len(all_gpu_ids)
dist_config = {
"dist_backend": "nccl",
"dist_url": "tcp://localhost:54189",
"world_size": 1
}
# https://discuss.pytorch.org/t/how-to-fix-a-sigsegv-in-pytorch-when-using-distributed-training-e-g-ddp/113518/10#:~:text=Using%20start%20and%20join%20avoids
# https://github.com/pytorch/pytorch/issues/40403#issuecomment-648515174
# mp.set_start_method('spawn')
if num_gpus > 1:
result_queue = mp.Queue()
for rank in range(num_gpus):
mp.Process(target=extract_worker, args=(rank, wav_fns, id_and_sizes, ckpt, sr, hop_size, bsz, max_tokens, fmax,
fmin, dist_config, num_gpus, ds_workers, result_queue,)).start()
f0_res = [None] * len(wav_fns)
for _ in range(num_gpus):
f0_res_dict = result_queue.get()
for idx in f0_res_dict:
f0_res[idx] = f0_res_dict[idx]
del f0_res_dict
else:
# f0_res = extract_one_process(wav_fns, id_and_sizes, ckpt, sr, hop_size, bsz, max_tokens, fmax, fmin)
f0_res_dict = extract_worker(0, wav_fns, id_and_sizes, ckpt, sr, hop_size, bsz, max_tokens, fmax,
fmin, dist_config, num_gpus, ds_workers, None)
f0_res = [None] * len(wav_fns)
for idx in f0_res_dict:
f0_res[idx] = f0_res_dict[idx]
return f0_res
@torch.no_grad()
def extract_worker(rank, wav_fns: list, id_and_sizes=None, ckpt=None, sr=24000, hop_size=128, bsz=128, max_tokens=100000,
fmax=900, fmin=50, dist_config=None, num_gpus=1, ds_workers=0, q=None):
# print(f"rank: {rank}")
if num_gpus > 1:
init_process_group(backend=dist_config['dist_backend'], init_method=dist_config['dist_url'],
world_size=dist_config['world_size'] * num_gpus, rank=rank)
dataset = RMVPEInferDataset(wav_fns, id_and_sizes, sr, hop_size, num_workers=ds_workers)
# ds_sampler = DistributedSampler(dataset, shuffle=False) if num_gpus > 1 else None
# loader = DataLoader(dataset, sampler=ds_sampler, collate_fn=dataset.collator, batch_size=1, num_workers=40, drop_last=False)
loader = build_dataloader(dataset, shuffle=False, max_tokens=max_tokens, max_sentences=bsz, use_ddp=num_gpus > 1)
loader = tqdm(loader, desc=f'| Processing f0 in [n_ranks={num_gpus}; max_tokens={max_tokens}; max_sentences={bsz}]') if rank == 0 else loader
device = torch.device(f"cuda:{int(rank)}")
model = RMVPE(ckpt, device=device)
f0_res_dict = {}
for batch in loader:
if batch is None or len(batch) == 0:
continue
idxs = [item[0] for item in batch]
wavs = [item[1] for item in batch]
lengths = [(wav.shape[0] + hop_size - 1) // hop_size for wav in wavs]
with torch.no_grad():
f0s, uvs = model.get_pitch_batch(
wavs, sample_rate=sr,
hop_size=hop_size,
lengths=lengths,
fmax=fmax,
fmin=fmin
)
for i, idx in enumerate(idxs):
f0_res_dict[idx] = f0s[i]
if q is not None:
q.put(f0_res_dict)
else:
return f0_res_dict
# old version
def extract_one_process(wav_fns: list, id_and_sizes=None, ckpt=None, sr=24000, hop_size=128, bsz=128, max_tokens=100000,
fmax=900, fmin=50, device='cuda'):
assert ckpt is not None
rmvpe = RMVPE(ckpt, device=device)
if id_and_sizes is None:
id_and_sizes = []
if type(wav_fns[0]) == str: # wav_paths
for idx, wav_path in enumerate(wav_fns):
total_frames = get_wav_num_frames(wav_path, sr)
id_and_sizes.append((idx, round(total_frames / hop_size)))
else: # numpy arrays, mono wavs
for idx, wav in enumerate(wav_fns):
id_and_sizes.append((idx, round(wav.shape[-1] / hop_size)))
get_size = lambda x: x[1]
bs = batch_by_size(id_and_sizes, get_size, max_tokens=max_tokens, max_sentences=bsz)
for i in range(len(bs)):
bs[i] = [bs[i][j][0] for j in range(len(bs[i]))]
f0_res = [None] * len(wav_fns)
for batch in tqdm(bs, total=len(bs), desc=f'| Processing f0 in [max_tokens={max_tokens}; max_sentences={bsz}]'):
wavs, mel_lengths, lengths = [], [], []
for idx in batch:
if type(wav_fns[idx]) == str:
wav_fn = wav_fns[idx]
wav, _ = librosa.core.load(wav_fn, sr=sr)
else:
wav = wav_fns[idx]
wavs.append(wav)
mel_lengths.append(math.ceil((wav.shape[0] + 1) / hop_size))
lengths.append((wav.shape[0] + hop_size - 1) // hop_size)
with torch.no_grad():
f0s, uvs = rmvpe.get_pitch_batch(
wavs, sample_rate=sr,
hop_size=hop_size,
lengths=lengths,
fmax=fmax,
fmin=fmin
)
for i, idx in enumerate(batch):
f0_res[idx] = f0s[i]
if rmvpe is not None:
rmvpe.release_cuda()
torch.cuda.empty_cache()
rmvpe = None
return f0_res
@@ -0,0 +1,134 @@
import math
import numpy as np
import torch
import torch.nn.functional as F
from torchaudio.transforms import Resample
import pyworld as pw
from ....utils.audio.pitch_utils import interp_f0, resample_align_curve
from .constants import *
from .model import E2E0
from .spec import MelSpectrogram
from .utils import to_local_average_f0, to_viterbi_f0
class RMVPE:
def __init__(self, model_path, hop_length=160, device=None):
self.resample_kernel = {}
if device is None:
self.device = 'cuda' if torch.cuda.is_available() else 'cpu'
else:
self.device = device
self.model = E2E0(4, 1, (2, 2)).eval().to(self.device)
ckpt = torch.load(model_path, map_location=self.device)
self.model.load_state_dict(ckpt['model'], strict=False)
self.mel_extractor = MelSpectrogram(
N_MELS, SAMPLE_RATE, WINDOW_LENGTH, hop_length, None, MEL_FMIN, MEL_FMAX
).to(self.device)
self.hop_length = hop_length
@torch.no_grad()
def mel2hidden(self, mel):
n_frames = mel.shape[-1]
mel = F.pad(mel, (0, 32 * ((n_frames - 1) // 32 + 1) - n_frames), mode='constant')
hidden = self.model(mel)
return hidden[:, :n_frames]
def decode(self, hidden, thred=0.03, use_viterbi=False):
if use_viterbi:
f0 = to_viterbi_f0(hidden, thred=thred)
else:
f0 = to_local_average_f0(hidden, thred=thred)
return f0
def postprocess(self, f0, fmin=50, fmax=1000, audio=None, min_gap=2):
if audio is not None:
# this doesn't work. deprecated
t = np.arange(0, f0.shape[0] * self.hop_length / 16000, self.hop_length / 16000)
f0 = pw.stonemask(audio.astype(np.float64), f0.astype(np.float64), t, 16000).astype(float)
f0[f0 < fmin] = 0
f0[f0 > fmax] = 0
# eliminate glitch
# min_gap: if successive positive f0 positions < min_gap, zero these positions
# eg: if min_gap=2, [0, 500, 500, 0] => [0, 0, 0, 0]
for idx in range(f0.shape[0] - min_gap - 1):
if f0[idx] == 0 and f0[idx + min_gap + 1] == 0 and np.sum(f0[idx: idx + min_gap + 2]) > 0:
f0[idx: idx + min_gap + 2] = 0
return f0
def infer_from_audio(self, audio, sample_rate=16000, thred=0.03, use_viterbi=False):
audio = torch.from_numpy(audio).float().unsqueeze(0).to(self.device)
if sample_rate == 16000:
audio_res = audio
else:
key_str = str(sample_rate)
if key_str not in self.resample_kernel:
self.resample_kernel[key_str] = Resample(sample_rate, 16000, lowpass_filter_width=128)
self.resample_kernel[key_str] = self.resample_kernel[key_str].to(self.device)
audio_res = self.resample_kernel[key_str](audio)
mel = self.mel_extractor(audio_res, center=True)
hidden = self.mel2hidden(mel)
f0 = self.decode(hidden, thred=thred, use_viterbi=use_viterbi).squeeze(0)
return f0
def get_pitch(self, waveform, sample_rate, hop_size, length, interp_uv=False, fmin=50, fmax=1000):
f0 = self.infer_from_audio(waveform, sample_rate=sample_rate)
f0 = self.postprocess(f0, fmin, fmax)
uv = f0 == 0
time_step = hop_size / sample_rate
f0_res = resample_align_curve(f0, 0.01, time_step, length)
uv_res = resample_align_curve(uv.astype(np.float32), 0.01, time_step, length) > 0.5
if not interp_uv:
f0_res[uv_res] = 0
return f0_res, uv_res
def infer_from_audio_batch(self, audios, sample_rate=16000, thred=0.03, use_viterbi=False):
from ....utils.commons.dataset_utils import collate_1d_or_2d
if isinstance(audios, list):
audios = [torch.from_numpy(audio).float() for audio in audios]
sizes = [math.ceil((audio.shape[0] + 1) / self.hop_length) for audio in audios]
audios = collate_1d_or_2d(audios, 0.0).to(self.device)
elif isinstance(audios, torch.Tensor):
sizes = None
if audios.device != self.device:
audios = audios.to(self.device)
else:
raise NotImplementedError
if sample_rate == 16000:
audios_res = audios
else:
key_str = str(sample_rate)
if key_str not in self.resample_kernel:
self.resample_kernel[key_str] = Resample(sample_rate, 16000, lowpass_filter_width=128)
self.resample_kernel[key_str] = self.resample_kernel[key_str].to(self.device)
audios_res = self.resample_kernel[key_str](audios)
mels = self.mel_extractor(audios_res, center=True)
hiddens = self.mel2hidden(mels)
f0 = self.decode(hiddens, thred=thred, use_viterbi=use_viterbi)
f0s = []
for i in range(f0.shape[0]):
f = f0[i, :sizes[i]] if sizes is not None else f0[i, :]
f0s.append(f)
return f0s
def get_pitch_batch(self, waveforms, sample_rate, hop_size, lengths, interp_uv=False, fmin=50, fmax=1000):
# hop_size, sample_rate: tgt params
f0s = self.infer_from_audio_batch(waveforms, sample_rate=sample_rate)
f0s_res, uvs_res = [], []
for idx, f0 in enumerate(f0s):
f0 = self.postprocess(f0, fmin, fmax, min_gap=6)
uv = f0 == 0
length = lengths[idx]
time_step = hop_size / sample_rate
f0_res = resample_align_curve(f0, 0.01, time_step, length)
uv_res = resample_align_curve(uv.astype(np.float32), 0.01, time_step, length) > 0.5
if not interp_uv:
f0_res[uv_res] = 0
f0s_res.append(f0_res)
uvs_res.append(uv_res)
return f0s_res, uvs_res
def release_cuda(self):
self.model = self.model.cpu()
self.mel_extractor = self.mel_extractor.cpu()
@@ -0,0 +1,32 @@
from torch import nn
from .constants import *
from .deepunet import DeepUnet0
from .seq import BiGRU
class E2E0(nn.Module):
def __init__(self, n_blocks, n_gru, kernel_size, en_de_layers=5, inter_layers=4, in_channels=1,
en_out_channels=16):
super(E2E0, self).__init__()
self.unet = DeepUnet0(kernel_size, n_blocks, en_de_layers, inter_layers, in_channels, en_out_channels)
self.cnn = nn.Conv2d(en_out_channels, 3, (3, 3), padding=(1, 1))
if n_gru:
self.fc = nn.Sequential(
BiGRU(3 * N_MELS, 256, n_gru),
nn.Linear(512, N_CLASS),
nn.Dropout(0.25),
nn.Sigmoid()
)
else:
self.fc = nn.Sequential(
nn.Linear(3 * N_MELS, N_CLASS),
nn.Dropout(0.25),
nn.Sigmoid()
)
def forward(self, mel):
mel = mel.transpose(-1, -2).unsqueeze(1)
x = self.cnn(self.unet(mel)).transpose(1, 2).flatten(-2)
x = self.fc(x)
return x
@@ -0,0 +1,10 @@
import torch.nn as nn
class BiGRU(nn.Module):
def __init__(self, input_features, hidden_features, num_layers):
super(BiGRU, self).__init__()
self.gru = nn.GRU(input_features, hidden_features, num_layers=num_layers, batch_first=True, bidirectional=True)
def forward(self, x):
return self.gru(x)[0]
@@ -0,0 +1,72 @@
import torch
import numpy as np
import torch.nn.functional as F
from librosa.filters import mel
class MelSpectrogram(torch.nn.Module):
def __init__(
self,
n_mel_channels,
sampling_rate,
win_length,
hop_length,
n_fft=None,
mel_fmin=0,
mel_fmax=None,
clamp=1e-5
):
super().__init__()
n_fft = win_length if n_fft is None else n_fft
self.hann_window = {}
mel_basis = mel(
sr=sampling_rate,
n_fft=n_fft,
n_mels=n_mel_channels,
fmin=mel_fmin,
fmax=mel_fmax,
htk=True)
mel_basis = torch.from_numpy(mel_basis).float()
self.register_buffer("mel_basis", mel_basis)
self.n_fft = win_length if n_fft is None else n_fft
self.hop_length = hop_length
self.win_length = win_length
self.sampling_rate = sampling_rate
self.n_mel_channels = n_mel_channels
self.clamp = clamp
def forward(self, audio, keyshift=0, speed=1, center=True):
factor = 2 ** (keyshift / 12)
n_fft_new = int(np.round(self.n_fft * factor))
win_length_new = int(np.round(self.win_length * factor))
hop_length_new = int(np.round(self.hop_length * speed))
keyshift_key = str(keyshift) + '_' + str(audio.device)
if keyshift_key not in self.hann_window:
self.hann_window[keyshift_key] = torch.hann_window(win_length_new).to(audio.device)
if center:
pad_left = win_length_new // 2
pad_right = (win_length_new + 1) // 2
audio = F.pad(audio, (pad_left, pad_right))
fft = torch.stft(
audio,
n_fft=n_fft_new,
hop_length=hop_length_new,
win_length=win_length_new,
window=self.hann_window[keyshift_key],
center=False,
return_complex=True
)
magnitude = fft.abs()
if keyshift != 0:
size = self.n_fft // 2 + 1
resize = magnitude.size(1)
if resize < size:
magnitude = F.pad(magnitude, (0, 0, 0, size - resize))
magnitude = magnitude[:, :size, :] * self.win_length / win_length_new
mel_output = torch.matmul(self.mel_basis, magnitude)
log_mel_spec = torch.log(torch.clamp(mel_output, min=self.clamp))
return log_mel_spec
@@ -0,0 +1,43 @@
import librosa
import numpy as np
import torch
from .constants import *
def to_local_average_f0(hidden, center=None, thred=0.03):
idx = torch.arange(N_CLASS, device=hidden.device)[None, None, :] # [B=1, T=1, N]
idx_cents = idx * 20 + CONST # [B=1, N]
if center is None:
center = torch.argmax(hidden, dim=2, keepdim=True) # [B, T, 1]
start = torch.clip(center - 4, min=0) # [B, T, 1]
end = torch.clip(center + 5, max=N_CLASS) # [B, T, 1]
idx_mask = (idx >= start) & (idx < end) # [B, T, N]
weights = hidden * idx_mask # [B, T, N]
product_sum = torch.sum(weights * idx_cents, dim=2) # [B, T]
weight_sum = torch.sum(weights, dim=2) # [B, T]
cents = product_sum / (weight_sum + (weight_sum == 0)) # avoid dividing by zero, [B, T]
f0 = 10 * 2 ** (cents / 1200)
uv = hidden.max(dim=2)[0] < thred # [B, T]
f0 = f0 * ~uv
return f0.cpu().numpy()
def to_viterbi_f0(hidden, thred=0.03):
# Create viterbi transition matrix
if not hasattr(to_viterbi_f0, 'transition'):
xx, yy = np.meshgrid(range(N_CLASS), range(N_CLASS))
transition = np.maximum(30 - abs(xx - yy), 0)
transition = transition / transition.sum(axis=1, keepdims=True)
to_viterbi_f0.transition = transition
# Convert to probability
prob = hidden.squeeze(0).cpu().numpy()
prob = prob.T
prob = prob / prob.sum(axis=0)
# Perform viterbi decoding
path = librosa.sequence.viterbi(prob, to_viterbi_f0.transition).astype(np.int64)
center = torch.from_numpy(path).unsqueeze(0).unsqueeze(-1).to(hidden.device)
return to_local_average_f0(hidden, center=center, thred=thred)
@@ -0,0 +1 @@
"""Core ROSVOT model components."""
@@ -0,0 +1,295 @@
from copy import deepcopy
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
from ...utils.commons.hparams import hparams
from ...utils.commons.gpu_mem_track import MemTracker
from ..commons.layers import Embedding
from ..commons.conv import ResidualBlock, ConvBlocks
from ..commons.conformer.conformer import ConformerLayers
from .unet import Unet
def regulate_boundary(bd_logits, threshold, min_gap=18, ref_bd=None, ref_bd_min_gap=8, non_padding=None):
# this doesn't preserve gradient
device = bd_logits.device
bd_logits = torch.sigmoid(bd_logits).data.cpu()
# bd_logits[0] = bd_logits[-1] = 1e-5 # avoid itv invalid problem
bd = (bd_logits > threshold).long()
bd_res = torch.zeros_like(bd).long()
for i in range(bd.shape[0]):
bd_i = bd[i]
last_bd_idx = -1
start = -1
for j in range(bd_i.shape[0]):
if bd_i[j] == 1:
if 0 <= start < j:
continue
elif start < 0:
start = j
else:
if 0 <= start < j:
if j - 1 > start:
bd_idx = start + int(torch.argmax(bd_logits[i, start: j]).item())
else:
bd_idx = start
if bd_idx - last_bd_idx < min_gap and last_bd_idx > 0:
bd_idx = round((bd_idx + last_bd_idx) / 2)
bd_res[i, last_bd_idx] = 0
bd_res[i, bd_idx] = 1
last_bd_idx = bd_idx
start = -1
# assert ref_bd_min_gap <= min_gap // 2
if ref_bd is not None and ref_bd_min_gap > 0:
ref = ref_bd.data.cpu()
for i in range(bd_res.shape[0]):
ref_bd_i = ref[i]
ref_bd_i_js = []
for j in range(ref_bd_i.shape[0]):
if ref_bd_i[j] == 1:
ref_bd_i_js.append(j)
seg_sum = torch.sum(bd_res[i, max(0, j - ref_bd_min_gap): j + ref_bd_min_gap])
if seg_sum == 0:
bd_res[i, j] = 1
elif seg_sum == 1 and bd_res[i, j] != 1:
bd_res[i, max(0, j - ref_bd_min_gap): j + ref_bd_min_gap] = \
ref_bd_i[max(0, j - ref_bd_min_gap): j + ref_bd_min_gap]
elif seg_sum > 1:
for k in range(1, ref_bd_min_gap+1):
if bd_res[i, max(0, j - k)] == 1 and ref_bd_i[max(0, j - k)] != 1:
bd_res[i, max(0, j - k)] = 0
break
if bd_res[i, min(bd_res.shape[1] - 1, j + k)] == 1 and ref_bd_i[min(bd_res.shape[1] - 1, j + k)] != 1:
bd_res[i, min(bd_res.shape[1] - 1, j + k)] = 0
break
bd_res[i, j] = 1
# final check
assert torch.sum(bd_res[i, ref_bd_i_js]) == len(ref_bd_i_js), \
f"{torch.sum(bd_res[i, ref_bd_i_js])} {len(ref_bd_i_js)}"
bd_res = bd_res.to(device)
# force valid begin and end
bd_res[:, 0] = 0
if non_padding is not None:
for i in range(bd_res.shape[0]):
bd_res[i, sum(non_padding[i]) - 1:] = 0
else:
bd_res[:, -1] = 0
return bd_res
class BackboneNet(nn.Module):
def __init__(self, hparams):
super().__init__()
self.hidden_size = hidden_size = hparams['hidden_size']
self.dropout = hparams.get('dropout', 0.0)
updown_rates = [2, 2, 2]
channel_multiples = [1, 1, 1]
if hparams.get('updown_rates', None) is not None:
updown_rates = [int(i) for i in hparams.get('updown_rates', None).split('-')]
if hparams.get('channel_multiples', None) is not None:
channel_multiples = [float(i) for i in hparams.get('channel_multiples', None).split('-')]
assert len(updown_rates) == len(channel_multiples)
# convs
if hparams.get('bkb_net', 'conv') == 'conv':
self.net = Unet(hidden_size, down_layers=len(updown_rates), mid_layers=hparams.get('bkb_layers', 12),
up_layers=len(updown_rates), kernel_size=3, updown_rates=updown_rates,
channel_multiples=channel_multiples, dropout=0, is_BTC=True,
constant_channels=False, mid_net=None, use_skip_layer=hparams.get('unet_skip_layer', False))
# conformer
elif hparams.get('bkb_net', 'conv') == 'conformer':
mid_net = ConformerLayers(
hidden_size, num_layers=hparams.get('bkb_layers', 12), kernel_size=hparams.get('conformer_kernel', 9),
dropout=self.dropout, num_heads=4)
self.net = Unet(hidden_size, down_layers=len(updown_rates), up_layers=len(updown_rates), kernel_size=3,
updown_rates=updown_rates, channel_multiples=channel_multiples, dropout=0,
is_BTC=True, constant_channels=False, mid_net=mid_net,
use_skip_layer=hparams.get('unet_skip_layer', False))
def forward(self, x):
return self.net(x)
class PitchDecoder(nn.Module):
def __init__(self, hparams):
super().__init__()
self.hidden_size = hidden_size = hparams['hidden_size']
self.dropout = hparams.get('dropout', 0.0)
self.note_bd_out = nn.Linear(hidden_size, 1)
self.note_bd_temperature = max(1e-7, hparams.get('note_bd_temperature', 1.0))
# note prediction
self.pitch_attn_num_head = hparams.get('pitch_attn_num_head', 1)
self.multihead_dot_attn = nn.Linear(hidden_size, self.pitch_attn_num_head)
self.post = ConvBlocks(hidden_size, out_dims=hidden_size, dilations=None, kernel_size=3,
layers_in_block=1, c_multiple=1, dropout=self.dropout, num_layers=1,
post_net_kernel=3, act_type='leakyrelu')
self.pitch_out = nn.Linear(hidden_size, hparams.get('note_num', 100) + 4)
self.note_num = hparams.get('note_num', 100)
self.note_start = hparams.get('note_start', 30)
self.pitch_temperature = max(1e-7, hparams.get('note_pitch_temperature', 1.0))
def forward(self, feat, note_bd, train=True):
bsz, T, _ = feat.shape
attn = torch.sigmoid(self.multihead_dot_attn(feat)) # [B, T, C] -> [B, T, num_head]
attn = F.dropout(attn, self.dropout, train)
attn_feat = feat.unsqueeze(3) * attn.unsqueeze(2) # [B, T, C, 1] x [B, T, 1, num_head] -> [B, T, C, num_head]
attn_feat = torch.mean(attn_feat, dim=-1) # [B, T, C, num_head] -> [B, T, C]
mel2note = torch.cumsum(note_bd, 1)
note_length = torch.max(torch.sum(note_bd, dim=1)).item() + 1 # max length
note_lengths = torch.sum(note_bd, dim=1) + 1 # [B]
# print('note_length', note_length)
attn = torch.mean(attn, dim=-1, keepdim=True) # [B, T, num_head] -> [B, T, 1]
denom = mel2note.new_zeros(bsz, note_length, dtype=attn.dtype).scatter_add_(
dim=1, index=mel2note, src=attn.squeeze(-1)
) # [B, T] -> [B, note_length] count the note frames of each note (with padding excluded)
frame2note = mel2note.unsqueeze(-1).repeat(1, 1, self.hidden_size) # [B, T] -> [B, T, C], with padding included
note_aggregate = frame2note.new_zeros(bsz, note_length, self.hidden_size, dtype=attn_feat.dtype).scatter_add_(
dim=1, index=frame2note, src=attn_feat
) # [B, T, C] -> [B, note_length, C]
note_aggregate = note_aggregate / (denom.unsqueeze(-1) + 1e-5)
note_aggregate = F.dropout(note_aggregate, self.dropout, train)
note_logits = self.post(note_aggregate)
note_logits = self.pitch_out(note_logits) / self.pitch_temperature
# note_logits = torch.clamp(note_logits, min=-16., max=16.) # don't know need it or not
note_pred = torch.softmax(note_logits, dim=-1) # [B, note_length, note_num]
note_pred = torch.argmax(note_pred, dim=-1) # [B, note_length]
# for some reason, note idx maybe 130 (why?)
note_pred[note_pred > self.note_num] = 0
note_pred[note_pred < self.note_start] = 0
return note_lengths, note_logits, note_pred
class MidiExtractor(nn.Module):
def __init__(self, hparams):
super(MidiExtractor, self).__init__()
self.hparams = deepcopy(hparams)
self.hidden_size = hidden_size = hparams['hidden_size']
self.dropout = hparams.get('dropout', 0.0)
self.note_bd_threshold = hparams.get('note_bd_threshold', 0.5)
self.note_bd_min_gap = round(hparams.get('note_bd_min_gap', 100) * hparams['audio_sample_rate'] / 1000 / hparams['hop_size'])
self.note_bd_ref_min_gap = round(hparams.get('note_bd_ref_min_gap', 50) * hparams['audio_sample_rate'] / 1000 / hparams['hop_size'])
self.mel_proj = nn.Conv1d(hparams['use_mel_bins'], hidden_size, kernel_size=3, padding=1)
self.mel_encoder = ConvBlocks(hidden_size, out_dims=hidden_size, dilations=None, kernel_size=3,
layers_in_block=2, c_multiple=1, dropout=self.dropout, num_layers=1,
post_net_kernel=3, act_type='leakyrelu')
self.use_pitch = hparams.get('use_pitch_embed', True)
if self.use_pitch:
self.pitch_embed = Embedding(300, hidden_size, 0, 'kaiming')
self.uv_embed = Embedding(3, hidden_size, 0, 'kaiming')
self.use_wbd = hparams.get('use_wbd', True)
if self.use_wbd:
self.word_bd_embed = Embedding(3, hidden_size, 0, 'kaiming')
self.cond_encoder = ConvBlocks(hidden_size, out_dims=hidden_size, dilations=None, kernel_size=3,
layers_in_block=1, c_multiple=1, dropout=self.dropout, num_layers=1,
post_net_kernel=3, act_type='leakyrelu')
# backbone
self.net = BackboneNet(hparams)
# note bd prediction
self.note_bd_out = nn.Linear(hidden_size, 1)
self.note_bd_temperature = max(1e-7, hparams.get('note_bd_temperature', 1.0))
# note prediction
self.pitch_decoder = PitchDecoder(hparams)
self.reset_parameters()
def run_encoder(self, mel=None, word_bd=None, pitch=None, uv=None, non_padding=None):
mel_embed = self.mel_proj(mel.transpose(1, 2)).transpose(1, 2)
mel_embed = self.mel_encoder(mel_embed)
pitch_embed = word_bd_embed = 0
if self.use_pitch and pitch is not None and uv is not None:
pitch_embed = self.pitch_embed(pitch) + self.uv_embed(uv) # [B, T, C]
if self.use_wbd and word_bd is not None:
word_bd_embed = self.word_bd_embed(word_bd)
feat = self.cond_encoder(mel_embed + pitch_embed + word_bd_embed)
return feat
def forward(self, mel=None, word_bd=None, note_bd=None, pitch=None, uv=None, non_padding=None, train=True):
ret = {}
bsz, T, _ = mel.shape
feat = self.run_encoder(mel, word_bd, pitch, uv, non_padding)
feat = self.net(feat) # [B, T, C]
# note bd prediction
note_bd_logits = self.note_bd_out(F.dropout(feat, self.dropout, train)).squeeze(-1) / self.note_bd_temperature
note_bd_logits = torch.clamp(note_bd_logits, min=-16., max=16.)
ret['note_bd_logits'] = note_bd_logits # [B, T]
if note_bd is None or not train:
note_bd = regulate_boundary(note_bd_logits, self.note_bd_threshold, self.note_bd_min_gap,
word_bd, self.note_bd_ref_min_gap, non_padding)
ret['note_bd_pred'] = note_bd # [B, T]
# note pitch prediction
note_lengths, note_logits, note_pred = self.pitch_decoder(feat, note_bd, train)
ret['note_lengths'], ret['note_logits'], ret['note_pred'] = note_lengths, note_logits, note_pred
return ret
def reset_parameters(self):
nn.init.kaiming_normal_(self.pitch_decoder.multihead_dot_attn.weight, mode='fan_in')
nn.init.kaiming_normal_(self.note_bd_out.weight, mode='fan_in')
nn.init.kaiming_normal_(self.pitch_decoder.pitch_out.weight, mode='fan_in')
nn.init.kaiming_normal_(self.mel_proj.weight, mode='fan_in')
nn.init.constant_(self.pitch_decoder.multihead_dot_attn.bias, 0.0)
nn.init.constant_(self.note_bd_out.bias, 0.0)
nn.init.constant_(self.pitch_decoder.pitch_out.bias, 0.0)
class WordbdExtractor(MidiExtractor):
def __init__(self, hparams):
super().__init__(hparams)
self.use_wbd = False
self.word_bd_embed = None
self.note_bd_out = self.note_bd_temperature = self.pitch_decoder = None
self.word_bd_threshold = hparams.get('word_bd_threshold', 0.5)
self.word_bd_min_gap = round(
hparams.get('word_bd_min_gap', 100) * hparams['audio_sample_rate'] / 1000 / hparams['hop_size'])
self.word_bd_out = nn.Linear(self.hidden_size, 1)
self.word_bd_temperature = max(1e-7, hparams.get('word_bd_temperature', 1.0))
nn.init.kaiming_normal_(self.word_bd_out.weight, mode='fan_in')
nn.init.constant_(self.word_bd_out.bias, 0.0)
def forward(self, mel=None, pitch=None, uv=None, non_padding=None, train=True):
# gpu_tracker.track()
ret = {}
bsz, T, _ = mel.shape
feat = self.run_encoder(mel=mel, pitch=pitch, uv=uv, non_padding=non_padding)
feat = self.net(feat) # [B, T, C]
word_bd_logits = self.word_bd_out(F.dropout(feat, self.dropout, train)).squeeze(-1) / self.word_bd_temperature
word_bd_logits = torch.clamp(word_bd_logits, min=-16., max=16.)
ret['word_bd_logits'] = word_bd_logits # [B, T]
if not train:
word_bd = regulate_boundary(word_bd_logits, self.word_bd_threshold, self.word_bd_min_gap,
non_padding=non_padding)
ret['word_bd_pred'] = word_bd # [B, T]
return ret
def reset_parameters(self):
if self.use_pitch:
nn.init.kaiming_normal_(self.pitch_embed.weight, mode='fan_in')
nn.init.kaiming_normal_(self.uv_embed.weight, mode='fan_in')
nn.init.kaiming_normal_(self.mel_proj.weight, mode='fan_in')
if self.use_pitch:
nn.init.constant_(self.pitch_embed.weight[self.pitch_embed.padding_idx], 0.0)
nn.init.constant_(self.uv_embed.weight[self.uv_embed.padding_idx], 0.0)
@@ -0,0 +1,172 @@
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from ..commons.layers import LayerNorm, Embedding
from ..commons.conv import ConvBlocks, ResidualBlock, get_norm_builder, get_act_builder
class UnetDown(nn.Module):
def __init__(self, hidden_size, n_layers, kernel_size, down_rates, channel_multiples=None, dropout=0.0,
is_BTC=True, constant_channels=False):
super(UnetDown, self).__init__()
assert n_layers == len(down_rates) # downs, down sample rate
down_rates = [int(i) for i in down_rates]
self.n_layers = n_layers
self.hidden_size = hidden_size
self.is_BTC = is_BTC
channel_multiples = channel_multiples if channel_multiples is not None else down_rates
self.layers = nn.ModuleList()
self.downs = nn.ModuleList()
in_channels = hidden_size
for i in range(self.n_layers):
out_channels = int(in_channels * channel_multiples[i]) if not constant_channels else in_channels
self.layers.append(nn.Sequential(
ResidualBlock(in_channels, kernel_size, dilation=1, n=1, norm_type='ln', dropout=dropout,
c_multiple=1, ln_eps=1e-5, act_type='leakyrelu'),
nn.Conv1d(in_channels, out_channels, kernel_size, padding=(kernel_size - 1) // 2),
ResidualBlock(out_channels, kernel_size, dilation=1, n=1, norm_type='ln',
dropout=dropout, c_multiple=1, ln_eps=1e-5, act_type='leakyrelu')
))
self.downs.append(nn.Sequential(
nn.AvgPool1d(down_rates[i])
))
in_channels = out_channels
self.last_norm = get_norm_builder('ln', out_channels)()
self.post_net = nn.Conv1d(out_channels, out_channels, kernel_size=kernel_size,
padding=kernel_size // 2)
def forward(self, x, **kwargs):
# x [B, T, C]
if self.is_BTC:
x = x.transpose(1, 2) # [B, C, T]
skip_xs = []
for i in range(self.n_layers):
skip_x = self.layers[i](x)
x = self.downs[i](skip_x)
if self.is_BTC:
skip_xs.append(skip_x.transpose(1, 2)) # [B, T, C]
else:
skip_xs.append(skip_x)
x = self.post_net(self.last_norm(x))
if self.is_BTC:
x = x.transpose(1, 2)
return x, skip_xs
class UnetMid(nn.Module):
def __init__(self, hidden_size, kernel_size, n_layers=None, in_dims=None, out_dims=None,
dropout=0.0, is_BTC=True, net=None):
super(UnetMid, self).__init__()
in_dims = in_dims if in_dims is not None else hidden_size
out_dims = out_dims if out_dims is not None else hidden_size
self.pre = nn.Conv1d(in_dims, hidden_size, kernel_size, padding=kernel_size // 2)
self.post = nn.Conv1d(hidden_size, out_dims, kernel_size, padding=kernel_size // 2)
self.is_BTC = is_BTC
if net is not None:
self.net = net
else:
self.net = ConvBlocks(hidden_size, out_dims=hidden_size, dilations=None, kernel_size=kernel_size,
layers_in_block=2, c_multiple=2, dropout=dropout, num_layers=n_layers,
post_net_kernel=3, act_type='leakyrelu', is_BTC=is_BTC)
def forward(self, x, cond=None, **kwargs):
# x [B, T, C]
if self.is_BTC:
x = self.pre(x.transpose(1, 2)).transpose(1, 2)
else:
x = self.pre(x)
if cond is None:
cond = 0
x = self.net(x + cond)
if self.is_BTC:
x = self.post(x.transpose(1, 2)).transpose(1, 2)
else:
x = self.post(x)
return x
class UnetUp(nn.Module):
def __init__(self, hidden_size, n_layers, kernel_size, up_rates, channel_multiples=None, dropout=0.0,
is_BTC=True, constant_channels=False, use_skip_layer=False, skip_scale=1.0):
super(UnetUp, self).__init__()
assert n_layers == len(up_rates) # this is reversed in up module, from the output to the interface with middle
up_rates = [int(i) for i in up_rates]
self.n_layers = n_layers
self.hidden_size = hidden_size
self.is_BTC = is_BTC
self.skip_scale = skip_scale
channel_multiples = channel_multiples if channel_multiples is not None else up_rates
# in_channels = int(np.cumprod(channel_multiples)[-1] * hidden_size) if not constant_channels else hidden_size
self.in_channels_lst = (np.cumprod([1] + channel_multiples) * hidden_size).astype(int) if not constant_channels \
else [hidden_size for _ in range(self.n_layers + 1)]
in_channels = self.in_channels_lst[-1]
self.ups = nn.ModuleList()
self.skip_layers = nn.ModuleList()
self.layers = nn.ModuleList()
for i in range(self.n_layers-1, -1, -1):
out_channels = self.in_channels_lst[i] if not constant_channels else in_channels
self.ups.append(nn.Sequential(
nn.ConvTranspose1d(in_channels, in_channels, kernel_size=kernel_size, stride=up_rates[i],
padding=kernel_size//2, output_padding=up_rates[i]-1),
get_norm_builder('ln', in_channels)(),
get_act_builder('leakyrelu')()
))
self.layers.append(nn.Sequential(
# ResidualBlock(in_channels*2, kernel_size, dilation=1, n=1, norm_type='ln', dropout=dropout,
# c_multiple=1, ln_eps=1e-5, act_type='leakyrelu'),
nn.Conv1d(in_channels*2, out_channels, kernel_size, padding=(kernel_size - 1) // 2),
ResidualBlock(out_channels, kernel_size, dilation=1, n=1, norm_type='ln',
dropout=dropout, c_multiple=1, ln_eps=1e-5, act_type='leakyrelu')
))
if use_skip_layer:
self.skip_layers.append(
ResidualBlock(in_channels, kernel_size, dilation=1, n=1, norm_type='ln', dropout=dropout,
c_multiple=1, ln_eps=1e-5, act_type='leakyrelu')
)
else:
self.skip_layers.append(nn.Identity())
in_channels = out_channels
self.out_channels = out_channels
self.last_norm = get_norm_builder('ln', out_channels)()
self.post_net = nn.Conv1d(out_channels, out_channels, kernel_size=kernel_size,
padding=kernel_size // 2)
def forward(self, x, skips, **kwargs):
# x [B, T, C]
if self.is_BTC:
x = x.transpose(1, 2) # [B, C, T]
for i in range(self.n_layers):
x = self.ups[i](x)
skip_x = skips[self.n_layers - i - 1] if not self.is_BTC \
else skips[self.n_layers - i - 1].transpose(1, 2) # [B, T, C] -> [B, C, T]
skip_x = self.skip_layers[i](skip_x) * self.skip_scale
x = torch.cat((x, skip_x), dim=1) # [B, C, T]
x = self.layers[i](x)
x = self.post_net(self.last_norm(x))
if self.is_BTC:
x = x.transpose(1, 2)
return x
class Unet(nn.Module):
def __init__(self, hidden_size, down_layers, up_layers, kernel_size,
updown_rates, mid_layers=None, channel_multiples=None, dropout=0.0,
is_BTC=True, constant_channels=False, mid_net=None, use_skip_layer=False, skip_scale=1.0):
super(Unet, self).__init__()
assert len(updown_rates) == down_layers == up_layers, f"{len(updown_rates)}, {down_layers}, {up_layers}"
if channel_multiples is not None:
assert len(channel_multiples) == len(updown_rates)
else:
channel_multiples = updown_rates
self.down = UnetDown(hidden_size, down_layers, kernel_size, updown_rates,
channel_multiples, dropout, is_BTC, constant_channels)
down_out_dims = int(np.cumprod(channel_multiples)[-1] * hidden_size) if not constant_channels else hidden_size
self.mid = UnetMid(hidden_size, kernel_size, mid_layers,
in_dims=down_out_dims, out_dims=down_out_dims, dropout=dropout, is_BTC=is_BTC, net=mid_net)
self.up = UnetUp(hidden_size, up_layers, kernel_size, updown_rates,
channel_multiples, dropout, is_BTC, constant_channels, use_skip_layer, skip_scale)
def forward(self, x, mid_cond=None, **kwargs):
x, skips = self.down(x)
x = self.mid(x, mid_cond)
x = self.up(x, skips)
return x
@@ -0,0 +1,15 @@
def seed_everything(seed: int, seed_cudnn=False):
import random, os
import numpy as np
import torch
random.seed(seed)
os.environ['PYTHONHASHSEED'] = str(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed(seed)
if seed_cudnn:
torch.backends.cudnn.deterministic = True
torch.backends.cudnn.benchmark = True
@@ -0,0 +1,100 @@
import librosa
import numpy as np
import wave
import soundfile as sf
def librosa_pad_lr(x, fsize, fshift, pad_sides=1):
'''compute right padding (final frame) or both sides padding (first and final frames)
'''
assert pad_sides in (1, 2)
# return int(fsize // 2)
pad = (x.shape[0] // fshift + 1) * fshift - x.shape[0]
if pad_sides == 1:
return 0, pad
else:
return pad // 2, pad // 2 + pad % 2
def amp_to_db(x):
return 20 * np.log10(np.maximum(1e-5, x))
def db_to_amp(x):
return 10.0 ** (x * 0.05)
def normalize(S, min_level_db):
return (S - min_level_db) / -min_level_db
def denormalize(D, min_level_db):
return (D * -min_level_db) + min_level_db
def librosa_wav2spec(wav_path,
fft_size=1024,
hop_size=256,
win_length=1024,
window="hann",
num_mels=80,
fmin=80,
fmax=-1,
eps=1e-6,
sample_rate=22050,
loud_norm=False,
trim_long_sil=False):
import pyloudnorm as pyln
if isinstance(wav_path, str):
if trim_long_sil:
from .vad import trim_long_silences
wav, _, _ = trim_long_silences(wav_path, sample_rate)
else:
wav, _ = librosa.core.load(wav_path, sr=sample_rate)
else:
wav = wav_path
wav_orig = np.copy(wav)
if loud_norm:
meter = pyln.Meter(sample_rate) # create BS.1770 meter
loudness = meter.integrated_loudness(wav)
wav = pyln.normalize.loudness(wav, loudness, -22.0)
if np.abs(wav).max() > 1:
wav = wav / np.abs(wav).max()
# get amplitude spectrogram
x_stft = librosa.stft(wav, n_fft=fft_size, hop_length=hop_size,
win_length=win_length, window=window, pad_mode="constant")
linear_spc = np.abs(x_stft) # (n_bins, T)
# get mel basis
fmin = 0 if fmin == -1 else fmin
fmax = sample_rate / 2 if fmax == -1 else fmax
mel_basis = librosa.filters.mel(sr=sample_rate, n_fft=fft_size, n_mels=num_mels, fmin=fmin, fmax=fmax)
# calculate mel spec
mel = mel_basis @ linear_spc
mel = np.log10(np.maximum(eps, mel)) # (n_mel_bins, T)
l_pad, r_pad = librosa_pad_lr(wav, fft_size, hop_size, 1)
wav = np.pad(wav, (l_pad, r_pad), mode='constant', constant_values=0.0)
wav = wav[:mel.shape[1] * hop_size]
# log linear spec
linear_spc = np.log10(np.maximum(eps, linear_spc))
return {'wav': wav, 'mel': mel.T, 'linear': linear_spc.T, 'mel_basis': mel_basis, 'wav_orig': wav_orig}
def get_wav_num_frames(path, sr=None):
try:
with wave.open(path, 'rb') as f:
sr_ = f.getframerate()
if sr is None:
sr = sr_
return int(f.getnframes() / (sr_ / sr))
except wave.Error:
wav_file, sr_ = sf.read(path, dtype='float32')
if sr is None:
sr = sr_
return int(len(wav_file) / (sr_ / sr))
except:
wav_file, sr_ = librosa.core.load(path, sr=sr)
return len(wav_file)
@@ -0,0 +1,90 @@
import re
import torch
import numpy as np
from ..text.text_encoder import is_sil_phoneme
def get_mel2ph(tg_fn, ph, mel, hop_size, audio_sample_rate, min_sil_duration=0):
from textgrid import TextGrid
ph_list = ph.split(" ")
itvs = TextGrid.fromFile(tg_fn)[1]
itvs_ = []
for i in range(len(itvs)):
if itvs[i].maxTime - itvs[i].minTime < min_sil_duration and i > 0 and is_sil_phoneme(itvs[i].mark):
itvs_[-1].maxTime = itvs[i].maxTime
else:
itvs_.append(itvs[i])
itvs.intervals = itvs_
itv_marks = [itv.mark for itv in itvs]
tg_len = len([x for x in itvs if not is_sil_phoneme(x.mark)])
ph_len = len([x for x in ph_list if not is_sil_phoneme(x)])
assert tg_len == ph_len, (tg_len, ph_len, itv_marks, ph_list, tg_fn)
mel2ph = np.zeros([mel.shape[0]], int)
i_itv = 0
i_ph = 0
while i_itv < len(itvs):
itv = itvs[i_itv]
ph = ph_list[i_ph]
itv_ph = itv.mark
start_frame = int(itv.minTime * audio_sample_rate / hop_size + 0.5)
end_frame = int(itv.maxTime * audio_sample_rate / hop_size + 0.5)
if is_sil_phoneme(itv_ph) and not is_sil_phoneme(ph):
mel2ph[start_frame:end_frame] = i_ph
i_itv += 1
elif not is_sil_phoneme(itv_ph) and is_sil_phoneme(ph):
i_ph += 1
else:
if not ((is_sil_phoneme(itv_ph) and is_sil_phoneme(ph)) \
or re.sub(r'\d+', '', itv_ph.lower()) == re.sub(r'\d+', '', ph.lower())):
print(f"| WARN: {tg_fn} phs are not same: ", itv_ph, ph, itv_marks, ph_list)
mel2ph[start_frame:end_frame] = i_ph + 1
i_ph += 1
i_itv += 1
mel2ph[-1] = mel2ph[-2]
assert not np.any(mel2ph == 0)
T_t = len(ph_list)
dur = mel2token_to_dur(mel2ph, T_t)
return mel2ph.tolist(), dur.tolist()
def split_audio_by_mel2ph(audio, mel2ph, hop_size, audio_num_mel_bins):
if isinstance(audio, torch.Tensor):
audio = audio.numpy()
if isinstance(mel2ph, torch.Tensor):
mel2ph = mel2ph.numpy()
assert len(audio.shape) == 1, len(mel2ph.shape) == 1
split_locs = []
for i in range(1, len(mel2ph)):
if mel2ph[i] != mel2ph[i - 1]:
split_loc = i * hop_size
split_locs.append(split_loc)
new_audio = []
for i in range(len(split_locs) - 1):
new_audio.append(audio[split_locs[i]:split_locs[i + 1]])
new_audio.append(np.zeros([0.5 * audio_num_mel_bins]))
return np.concatenate(new_audio)
def mel2token_to_dur(mel2token, T_txt=None, max_dur=None):
is_torch = isinstance(mel2token, torch.Tensor)
has_batch_dim = True
if not is_torch:
mel2token = torch.LongTensor(mel2token)
if T_txt is None:
T_txt = mel2token.max()
if len(mel2token.shape) == 1:
mel2token = mel2token[None, ...]
has_batch_dim = False
B, _ = mel2token.shape
dur = mel2token.new_zeros(B, T_txt + 1).scatter_add(1, mel2token, torch.ones_like(mel2token))
dur = dur[:, 1:]
if max_dur is not None:
dur = dur.clamp(max=max_dur)
if not is_torch:
dur = dur.numpy()
if not has_batch_dim:
dur = dur[0]
return dur
@@ -0,0 +1,22 @@
import subprocess
import numpy as np
from scipy.io import wavfile
def save_wav(wav, path, sr, norm=False):
if norm:
wav = wav / np.abs(wav).max()
wav = wav * 32767
wavfile.write(path[:-4] + '.wav', sr, wav.astype(np.int16))
if path[-4:] == '.mp3':
to_mp3(path[:-4])
def to_mp3(out_path):
if out_path[-4:] == '.wav':
out_path = out_path[:-4]
subprocess.check_call(
f'ffmpeg -threads 1 -loglevel error -i "{out_path}.wav" -vn -b:a 192k -y -hide_banner -async 1 "{out_path}.mp3"',
shell=True, stdin=subprocess.PIPE)
subprocess.check_call(f'rm -f "{out_path}.wav"', shell=True)
@@ -0,0 +1,139 @@
import math
import numpy as np
import torch
import torch.utils.data
from librosa.filters import mel as librosa_mel_fn
from scipy.io.wavfile import read
import torch
import torch.nn as nn
MAX_WAV_VALUE = 32768.0
def load_wav(full_path):
sampling_rate, data = read(full_path)
return data, sampling_rate
def dynamic_range_compression(x, C=1, clip_val=1e-5):
return np.log10(np.clip(x, a_min=clip_val, a_max=None) * C)
def dynamic_range_decompression(x, C=1):
return np.exp(x) / C
def dynamic_range_compression_torch(x, C=1, clip_val=1e-5):
return torch.log10(torch.clamp(x, min=clip_val) * C)
def dynamic_range_decompression_torch(x, C=1):
return torch.exp(x) / C
def spectral_normalize_torch(magnitudes):
output = dynamic_range_compression_torch(magnitudes)
return output
def spectral_de_normalize_torch(magnitudes):
output = dynamic_range_decompression_torch(magnitudes)
return output
class MelNet(nn.Module):
def __init__(self, hparams, device='cpu') -> None:
super().__init__()
self.n_fft = hparams['fft_size']
self.num_mels = hparams['audio_num_mel_bins']
self.sampling_rate = hparams['audio_sample_rate']
self.hop_size = hparams['hop_size']
self.win_size = hparams['win_size']
self.fmin = hparams['fmin']
self.fmax = hparams['fmax']
self.device = device
mel = librosa_mel_fn(sr=self.sampling_rate, n_fft=self.n_fft, n_mels=self.num_mels, fmin=self.fmin,
fmax=self.fmax)
self.mel_basis = torch.from_numpy(mel).float().to(self.device)
self.hann_window = torch.hann_window(self.win_size).to(self.device)
def to(self, device, **kwagrs):
super().to(device=device, **kwagrs)
self.mel_basis = self.mel_basis.to(device)
self.hann_window = self.hann_window.to(device)
self.device = device
def forward(self, y, center=False, complex=False):
if isinstance(y, np.ndarray):
y = torch.FloatTensor(y)
if len(y.shape) == 1:
y = y.unsqueeze(0)
y = y.clamp(min=-1., max=1.).to(self.device)
pad_length = math.ceil(y.shape[1] / self.hop_size) * self.hop_size - y.shape[1]
y = torch.nn.functional.pad(y.unsqueeze(1),
[int((self.n_fft - self.hop_size) / 2),
int((self.n_fft - self.hop_size) / 2 + pad_length)],
mode='reflect')
y = y.squeeze(1)
spec = torch.stft(y, self.n_fft, hop_length=self.hop_size, win_length=self.win_size, window=self.hann_window,
center=center, pad_mode='reflect', normalized=False, onesided=True, return_complex=True)
if not complex:
spec = torch.view_as_real(spec)
spec = torch.sqrt(spec.pow(2).sum(-1) + (1e-9)) # [B, n_fft, T]
spec = torch.matmul(self.mel_basis, spec)
spec = spectral_normalize_torch(spec)
spec = spec.transpose(1, 2) # [B, T, n_fft]
else:
B, C, T, _ = spec.shape
spec = spec.transpose(1, 2) # [B, T, n_fft, 2]
return spec
## below can be used in one gpu, but not ddp
mel_basis = {}
hann_window = {}
def mel_spectrogram(y, hparams, center=False, complex=False): # y should be a tensor with shape (b,wav_len)
# hop_size: 512 # For 22050Hz, 275 ~= 12.5 ms (0.0125 * sample_rate)
# win_size: 2048 # For 22050Hz, 1100 ~= 50 ms (If None, win_size: fft_size) (0.05 * sample_rate)
# fmin: 55 # Set this to 55 if your speaker is male! if female, 95 should help taking off noise. (To test depending on dataset. Pitch info: male~[65, 260], female~[100, 525])
# fmax: 10000 # To be increased/reduced depending on data.
# fft_size: 2048 # Extra window size is filled with 0 paddings to match this parameter
# n_fft, num_mels, sampling_rate, hop_size, win_size, fmin, fmax,
n_fft = hparams['fft_size']
num_mels = hparams['audio_num_mel_bins']
sampling_rate = hparams['audio_sample_rate']
hop_size = hparams['hop_size']
win_size = hparams['win_size']
fmin = hparams['fmin']
fmax = hparams['fmax']
if isinstance(y, np.ndarray):
y = torch.FloatTensor(y)
if len(y.shape) == 1:
y = y.unsqueeze(0)
y = y.clamp(min=-1., max=1.)
global mel_basis, hann_window
if fmax not in mel_basis:
mel = librosa_mel_fn(sampling_rate, n_fft, num_mels, fmin, fmax)
mel_basis[str(fmax) + '_' + str(y.device)] = torch.from_numpy(mel).float().to(y.device)
hann_window[str(y.device)] = torch.hann_window(win_size).to(y.device)
y = torch.nn.functional.pad(y.unsqueeze(1), [int((n_fft - hop_size) / 2), int((n_fft - hop_size) / 2)],
mode='reflect')
y = y.squeeze(1)
spec = torch.stft(y, n_fft, hop_length=hop_size, win_length=win_size, window=hann_window[str(y.device)],
center=center, pad_mode='reflect', normalized=False, onesided=True, return_complex=complex)
if not complex:
spec = torch.sqrt(spec.pow(2).sum(-1) + (1e-9))
spec = torch.matmul(mel_basis[str(fmax) + '_' + str(y.device)], spec)
spec = spectral_normalize_torch(spec)
else:
B, C, T, _ = spec.shape
spec = spec.transpose(1, 2) # [B, T, n_fft, 2]
return spec
@@ -0,0 +1,60 @@
import math
import numpy as np
PITCH_EXTRACTOR = {}
def register_pitch_extractor(name):
def register_pitch_extractor_(cls):
PITCH_EXTRACTOR[name] = cls
return cls
return register_pitch_extractor_
def get_pitch_extractor(name):
return PITCH_EXTRACTOR[name]
def extract_pitch_simple(wav):
from ..commons.hparams import hparams
return extract_pitch(hparams['pitch_extractor'], wav,
hparams['hop_size'], hparams['audio_sample_rate'],
f0_min=hparams['f0_min'], f0_max=hparams['f0_max'])
def extract_pitch(extractor_name, wav_data, hop_size, audio_sample_rate, f0_min=75, f0_max=800, **kwargs):
return get_pitch_extractor(extractor_name)(wav_data, hop_size, audio_sample_rate, f0_min, f0_max, **kwargs)
@register_pitch_extractor('parselmouth')
def parselmouth_pitch(wav_data, hop_size, audio_sample_rate, f0_min, f0_max,
voicing_threshold=0.6, *args, **kwargs):
import parselmouth
time_step = hop_size / audio_sample_rate * 1000
n_mel_frames = int(len(wav_data) // hop_size)
f0_pm = parselmouth.Sound(wav_data, audio_sample_rate).to_pitch_ac(
time_step=time_step / 1000, voicing_threshold=voicing_threshold,
pitch_floor=f0_min, pitch_ceiling=f0_max).selected_array['frequency']
pad_size = (n_mel_frames - len(f0_pm) + 1) // 2
f0 = np.pad(f0_pm, [[pad_size, n_mel_frames - len(f0_pm) - pad_size]], mode='constant')
return f0
@register_pitch_extractor('pyworld')
def pyworld_pitch(wav_data, hop_size, audio_sample_rate, f0_min, f0_max,
voicing_threshold=0.6, *args, **kwargs):
import pyworld as pw
# f0, _ = pw.harvest(wav_data.astype(np.double), audio_sample_rate, f0_floor=f0_min, f0_ceil=f0_max,
# frame_period=hop_size * 1000 / audio_sample_rate)
f0, _ = pw.dio(wav_data.astype(np.double), audio_sample_rate, f0_floor=f0_min, f0_ceil=f0_max, frame_period=hop_size * 1000 / audio_sample_rate)
f0[f0 < f0_min] = 0.0
f0[f0 > f0_max] = 0.0
n_mel_frames = math.ceil(len(wav_data) / hop_size)
if n_mel_frames > len(f0):
pad_size = (n_mel_frames - len(f0) + 1) // 2
f0 = np.pad(f0, [[pad_size, n_mel_frames - len(f0) - pad_size]], mode='constant')
elif n_mel_frames < len(f0):
left_del = (len(f0) - n_mel_frames + 1) // 2
right_del = len(f0) - n_mel_frames - left_del
f0 = f0[left_del: (-right_del if right_del > 0 else len(f0))]
return f0
@@ -0,0 +1,303 @@
import numpy as np
import torch
import pretty_midi
def to_lf0(f0):
f0[f0 < 1.0e-5] = 1.0e-6
lf0 = f0.log() if isinstance(f0, torch.Tensor) else np.log(f0)
lf0[f0 < 1.0e-5] = - 1.0E+10
return lf0
def to_f0(lf0):
f0 = np.where(lf0 <= 0, 0.0, np.exp(lf0))
return f0.flatten()
def f0_to_coarse(f0, f0_bin=256, f0_max=900.0, f0_min=50.0):
f0_mel_min = 1127 * np.log(1 + f0_min / 700)
f0_mel_max = 1127 * np.log(1 + f0_max / 700)
is_torch = isinstance(f0, torch.Tensor)
f0_mel = 1127 * (1 + f0 / 700).log() if is_torch else 1127 * np.log(1 + f0 / 700)
f0_mel[f0_mel > 0] = (f0_mel[f0_mel > 0] - f0_mel_min) * (f0_bin - 2) / (f0_mel_max - f0_mel_min) + 1
f0_mel[f0_mel <= 1] = 1
f0_mel[f0_mel > f0_bin - 1] = f0_bin - 1
f0_coarse = (f0_mel + 0.5).long() if is_torch else np.rint(f0_mel).astype(int)
assert f0_coarse.max() <= f0_bin-1 and f0_coarse.min() >= 1, (f0_coarse.max(), f0_coarse.min(), f0.min(), f0.max())
return f0_coarse
def coarse_to_f0(f0_coarse, f0_bin=256, f0_max=900.0, f0_min=50.0):
f0_mel_min = 1127 * np.log(1 + f0_min / 700)
f0_mel_max = 1127 * np.log(1 + f0_max / 700)
uv = f0_coarse == 1
f0 = f0_mel_min + (f0_coarse - 1) * (f0_mel_max - f0_mel_min) / (f0_bin - 2)
f0 = ((f0 / 1127).exp() - 1) * 700
f0[uv] = 0
return f0
def norm_f0(f0, uv, pitch_norm='log', f0_mean=400, f0_std=100):
is_torch = isinstance(f0, torch.Tensor)
if pitch_norm == 'standard':
f0 = (f0 - f0_mean) / f0_std
if pitch_norm == 'log':
f0 = torch.log2(f0 + 1e-8) if is_torch else np.log2(f0 + 1e-8)
if uv is not None:
f0[uv > 0] = 0
return f0
def norm_interp_f0(f0, pitch_norm='log', f0_mean=None, f0_std=None):
is_torch = isinstance(f0, torch.Tensor)
if is_torch:
device = f0.device
f0 = f0.data.cpu().numpy()
uv = f0 == 0
f0 = norm_f0(f0, uv, pitch_norm, f0_mean, f0_std)
if sum(uv) == len(f0):
f0[uv] = 0
elif sum(uv) > 0:
f0[uv] = np.interp(np.where(uv)[0], np.where(~uv)[0], f0[~uv])
if is_torch:
uv = torch.FloatTensor(uv)
f0 = torch.FloatTensor(f0)
f0 = f0.to(device)
uv = uv.to(device)
return f0, uv
def denorm_f0(f0, uv, pitch_norm='log', f0_mean=400, f0_std=100, pitch_padding=None, min=50, max=900):
is_torch = isinstance(f0, torch.Tensor)
if pitch_norm == 'standard':
f0 = f0 * f0_std + f0_mean
if pitch_norm == 'log':
f0 = 2 ** f0
f0 = f0.clamp(min=min, max=max) if is_torch else np.clip(f0, a_min=min, a_max=max)
if uv is not None:
f0[uv > 0] = 0
if pitch_padding is not None:
f0[pitch_padding] = 0
return f0
def interp_f0(f0, uv=None):
if uv is None:
uv = f0 == 0
f0 = norm_f0(f0, uv)
if uv.any() and not uv.all():
f0[uv] = np.interp(np.where(uv)[0], np.where(~uv)[0], f0[~uv])
return denorm_f0(f0, uv=None), uv
def resample_align_curve(points: np.ndarray, original_timestep: float, target_timestep: float, align_length=-1):
t_max = (len(points) - 1) * original_timestep
curve_interp = np.interp(
np.arange(0, t_max, target_timestep),
original_timestep * np.arange(len(points)),
points
).astype(points.dtype)
if align_length > 0:
delta_l = align_length - len(curve_interp)
if delta_l < 0:
curve_interp = curve_interp[:align_length]
elif delta_l > 0:
curve_interp = np.concatenate((curve_interp, np.full(delta_l, fill_value=curve_interp[-1])), axis=0)
return curve_interp
def midi_to_hz(midi):
if type(midi) == np.ndarray:
non_mask = midi == 0
freq_hz = 440.0 * 2.0 ** ((midi - 69.0) / 12.0)
freq_hz[non_mask] = 0
else:
freq_hz = 440.0 * 2.0 ** ((midi - 69.0) / 12.0)
return freq_hz
def hz_to_midi(hz):
if type(hz) == torch.Tensor:
non_mask = hz == 0
midi = 69.0 + 12.0 * (torch.log2(hz) - torch.log2(torch.Tensor(440.0)))
midi[non_mask] = 0
elif type(hz) == np.ndarray:
non_mask = hz == 0
midi = 69.0 + 12.0 * (np.log2(hz) - np.log2(440.0))
midi[non_mask] = 0
else:
midi = 69.0 + 12.0 * (np.log2(hz) - np.log2(440.0))
if hz == 0:
midi = 0
return midi
def boundary2Interval(bd):
# bd has a shape of [T] with T frames
is_torch = isinstance(bd, torch.Tensor)
if is_torch:
device = bd.device
bd = bd.data.cpu().numpy()
assert len(bd.shape) == 1
# force valid begin and end
# bd[0] = 0 # took care of in regulate_boundary()
# bd[-1] = 0
ret = np.zeros(shape=(bd.sum() + 1, 2), dtype=int)
ret_idx = 0
ret[0, 0] = 0
for i, u in enumerate(bd):
if i == 0:
continue
if u == 1:
ret[ret_idx, 1] = i
ret[ret_idx+1, 0] = i
ret_idx += 1
ret[-1, 1] = bd.shape[0] - 1
if is_torch:
ret = torch.LongTensor(ret).to(device)
return ret
def validate_pitch_and_itv(notes, note_itv):
# notes [T]
# note_itv [T, 2]
assert notes.shape[0] == note_itv.shape[0]
res_notes = []
res_note_itv = []
for idx in range(notes.shape[0]):
pitch, itv = notes[idx], note_itv[idx]
if itv[0] >= itv[1]:
raise RuntimeError("The note duration should be positive")
if pitch == 0:
continue
res_notes.append(pitch)
res_note_itv.append([itv[0], itv[1]])
res_notes = np.array(res_notes)
res_note_itv = np.array(res_note_itv)
return res_notes, res_note_itv
def save_midi(notes, note_itv, midi_path):
# notes [T]
# note_itv [T, 2]
notes, note_itv = validate_pitch_and_itv(notes, note_itv)
if notes.shape == (0,):
return None
assert notes.shape[0] == note_itv.shape[0]
piano_chord = pretty_midi.PrettyMIDI()
piano_program = pretty_midi.instrument_name_to_program('Acoustic Grand Piano')
piano = pretty_midi.Instrument(program=piano_program)
for idx in range(notes.shape[0]):
pitch, itv = notes[idx], note_itv[idx]
note = pretty_midi.Note(velocity=120, pitch=pitch, start=itv[0], end=itv[1])
piano.notes.append(note)
piano_chord.remove_invalid_notes()
piano_chord.instruments.append(piano)
piano_chord.write(midi_path)
return piano_chord
def midi2NoteInterval(mid):
assert type(mid) == pretty_midi.PrettyMIDI
if len(mid.instruments) == 0 or len(mid.instruments[0].notes) == 0:
return None
ret = np.zeros(shape=(len(mid.instruments[0].notes), 2))
for i, note in enumerate(mid.instruments[0].notes):
ret[i, 0] = note.start
ret[i, 1] = note.end
return ret
def midi2NotePitch(mid):
assert type(mid) == pretty_midi.PrettyMIDI
if len(mid.instruments) == 0 or len(mid.instruments[0].notes) == 0:
return None
ret = np.zeros(shape=len(mid.instruments[0].notes))
for i, note in enumerate(mid.instruments[0].notes):
ret[i] = note.pitch
return ret
def midi_onset_eval(mid_gt, mid_pred):
import mir_eval
interval_true = midi2NoteInterval(mid_gt)
if interval_true is None:
raise RuntimeError('Midi ground truth is None')
interval_pred = midi2NoteInterval(mid_pred)
if interval_pred is None:
return 0, 0, 0
onset_p, onset_r, onset_f = mir_eval.transcription.onset_precision_recall_f1(
interval_true, interval_pred, onset_tolerance=0.05, strict=False, beta=1.0)
return onset_p, onset_r, onset_f
def midi_offset_eval(mid_gt, mid_pred):
import mir_eval
interval_true = midi2NoteInterval(mid_gt)
if interval_true is None:
raise RuntimeError('Midi ground truth is None')
interval_pred = midi2NoteInterval(mid_pred)
if interval_pred is None:
return 0, 0, 0
offset_p, offset_r, offset_f = mir_eval.transcription.offset_precision_recall_f1(
interval_true, interval_pred, offset_ratio=0.2, offset_min_tolerance=0.05, strict=False, beta=1.0)
return offset_p, offset_r, offset_f
def midi_pitch_eval(mid_gt, mid_pred, offset_ratio=0.2):
import mir_eval
interval_true = midi2NoteInterval(mid_gt)
pitch_true = midi_to_hz(midi2NotePitch(mid_gt))
if interval_true is None or pitch_true is None:
raise RuntimeError('Midi ground truth is None')
interval_pred = midi2NoteInterval(mid_pred)
pitch_pred = midi2NotePitch(mid_pred)
if interval_pred is None:
return 0, 0, 0, 0
if pitch_pred is None:
pitch_pred = np.zeros(interval_pred.shape[0])
pitch_pred = midi_to_hz(pitch_pred)
overlap_p, overlap_r, overlap_f, avg_overlap_ratio = mir_eval.transcription.precision_recall_f1_overlap(
interval_true, pitch_true, interval_pred, pitch_pred, onset_tolerance=0.05, pitch_tolerance=50.0,
offset_ratio=offset_ratio, offset_min_tolerance=0.05, strict=False, beta=1.0)
return overlap_p, overlap_r, overlap_f, avg_overlap_ratio
def midi_COn_eval(mid_gt, mid_pred):
return midi_onset_eval(mid_gt, mid_pred)
def midi_COnP_eval(mid_gt, mid_pred):
return midi_pitch_eval(mid_gt, mid_pred, offset_ratio=None)
def midi_COnPOff_eval(mid_gt, mid_pred):
return midi_pitch_eval(mid_gt, mid_pred)
def midi_melody_eval(mid_gt, mid_pred, hop_size=256, sample_rate=48000):
interval_true = midi2NoteInterval(mid_gt)
pitch_true = midi_to_hz(midi2NotePitch(mid_gt))
if interval_true is None or pitch_true is None:
raise RuntimeError('Midi ground truth is None')
interval_pred = midi2NoteInterval(mid_pred)
pitch_pred = midi2NotePitch(mid_pred)
if interval_pred is None:
return 0, 0, 0, 0
if pitch_pred is None:
pitch_pred = np.zeros(interval_pred.shape[0])
pitch_pred = midi_to_hz(pitch_pred)
vr, vfa, rpa, rca, oa = melody_eval_pitch_and_itv(
pitch_true, interval_true, pitch_pred, interval_pred, hop_size, sample_rate)
return vr, vfa, rpa, rca, oa
def melody_eval_pitch_and_itv(pitch_true, interval_true, pitch_pred, interval_pred, hop_size=256, sample_rate=48000):
import mir_eval
t_gt = np.arange(0, interval_true[-1][1], hop_size / sample_rate)
freq_gt = np.zeros_like(t_gt)
for idx in range(len(pitch_true)):
freq_gt[min(len(freq_gt) - 1, round(interval_true[idx][0] * sample_rate / hop_size)): round(
interval_true[idx][1] * sample_rate / hop_size)] = pitch_true[idx]
t_pred = np.arange(0, interval_pred[-1][1], hop_size / sample_rate)
freq_pred = np.zeros_like(t_pred)
for idx in range(len(pitch_pred)):
freq_pred[min(len(freq_pred) - 1, round(interval_pred[idx][0] * sample_rate / hop_size)): round(
interval_pred[idx][1] * sample_rate / hop_size)] = pitch_pred[idx]
ref_voicing, ref_cent, est_voicing, est_cent = mir_eval.melody.to_cent_voicing(t_gt, freq_gt,
t_pred, freq_pred)
vr, vfa = mir_eval.melody.voicing_measures(ref_voicing,
est_voicing) # voicing recall, voicing false alarm
rpa = mir_eval.melody.raw_pitch_accuracy(ref_voicing, ref_cent, est_voicing, est_cent)
rca = mir_eval.melody.raw_chroma_accuracy(ref_voicing, ref_cent, est_voicing, est_cent)
oa = mir_eval.melody.overall_accuracy(ref_voicing, ref_cent, est_voicing, est_cent)
return vr, vfa, rpa, rca, oa
@@ -0,0 +1,78 @@
from skimage.transform import resize
import struct
import webrtcvad
from scipy.ndimage.morphology import binary_dilation
import librosa
import numpy as np
import pyloudnorm as pyln
import warnings
warnings.filterwarnings("ignore", message="Possible clipped samples in output")
int16_max = (2 ** 15) - 1
def trim_long_silences(path, sr=None, return_raw_wav=False, norm=True, vad_max_silence_length=12):
"""
Ensures that segments without voice in the waveform remain no longer than a
threshold determined by the VAD parameters in params.py.
:param wav: the raw waveform as a numpy array of floats
:param vad_max_silence_length: Maximum number of consecutive silent frames a segment can have.
:return: the same waveform with silences trimmed away (length <= original wav length)
"""
## Voice Activation Detection
# Window size of the VAD. Must be either 10, 20 or 30 milliseconds.
# This sets the granularity of the VAD. Should not need to be changed.
sampling_rate = 16000
wav_raw, sr = librosa.core.load(path, sr=sr)
if norm:
meter = pyln.Meter(sr) # create BS.1770 meter
loudness = meter.integrated_loudness(wav_raw)
wav_raw = pyln.normalize.loudness(wav_raw, loudness, -20.0)
if np.abs(wav_raw).max() > 1.0:
wav_raw = wav_raw / np.abs(wav_raw).max()
wav = librosa.resample(wav_raw, sr, sampling_rate, res_type='kaiser_best')
vad_window_length = 30 # In milliseconds
# Number of frames to average together when performing the moving average smoothing.
# The larger this value, the larger the VAD variations must be to not get smoothed out.
vad_moving_average_width = 8
# Compute the voice detection window size
samples_per_window = (vad_window_length * sampling_rate) // 1000
# Trim the end of the audio to have a multiple of the window size
wav = wav[:len(wav) - (len(wav) % samples_per_window)]
# Convert the float waveform to 16-bit mono PCM
pcm_wave = struct.pack("%dh" % len(wav), *(np.round(wav * int16_max)).astype(np.int16))
# Perform voice activation detection
voice_flags = []
vad = webrtcvad.Vad(mode=3)
for window_start in range(0, len(wav), samples_per_window):
window_end = window_start + samples_per_window
voice_flags.append(vad.is_speech(pcm_wave[window_start * 2:window_end * 2],
sample_rate=sampling_rate))
voice_flags = np.array(voice_flags)
# Smooth the voice detection with a moving average
def moving_average(array, width):
array_padded = np.concatenate((np.zeros((width - 1) // 2), array, np.zeros(width // 2)))
ret = np.cumsum(array_padded, dtype=float)
ret[width:] = ret[width:] - ret[:-width]
return ret[width - 1:] / width
audio_mask = moving_average(voice_flags, vad_moving_average_width)
audio_mask = np.round(audio_mask).astype(np.bool)
# Dilate the voiced regions
audio_mask = binary_dilation(audio_mask, np.ones(vad_max_silence_length + 1))
audio_mask = np.repeat(audio_mask, samples_per_window)
audio_mask = resize(audio_mask, (len(wav_raw),)) > 0
if return_raw_wav:
return wav_raw, audio_mask, sr
return wav_raw[audio_mask], audio_mask, sr
@@ -0,0 +1,235 @@
import logging
import os
import random
import subprocess
import sys
from datetime import datetime
import numpy as np
import torch.utils.data
from torch import nn
from torch.utils.tensorboard import SummaryWriter
from .dataset_utils import data_loader
from .hparams import hparams
from .meters import AvgrageMeter
from .tensor_utils import tensors_to_scalars
from .trainer import Trainer
torch.multiprocessing.set_sharing_strategy(os.getenv('TORCH_SHARE_STRATEGY', 'file_system'))
log_format = '%(asctime)s %(message)s'
logging.basicConfig(stream=sys.stdout, level=logging.INFO,
format=log_format, datefmt='%m/%d %I:%M:%S %p')
class BaseTask(nn.Module):
def __init__(self, *args, **kwargs):
super(BaseTask, self).__init__()
self.current_epoch = 0
self.global_step = 0
self.trainer = None
self.use_ddp = False
self.gradient_clip_norm = hparams['clip_grad_norm']
self.gradient_clip_val = hparams.get('clip_grad_value', 0)
self.model = None
self.training_losses_meter = None
self.logger: SummaryWriter = None
######################
# build model, dataloaders, optimizer, scheduler and tensorboard
######################
def build_model(self):
raise NotImplementedError
@data_loader
def train_dataloader(self):
raise NotImplementedError
@data_loader
def test_dataloader(self):
raise NotImplementedError
@data_loader
def val_dataloader(self):
raise NotImplementedError
def build_scheduler(self, optimizer):
return None
def build_optimizer(self, model):
raise NotImplementedError
def configure_optimizers(self):
optm = self.build_optimizer(self.model)
self.scheduler = self.build_scheduler(optm)
if isinstance(optm, (list, tuple)):
return optm
return [optm]
def build_tensorboard(self, save_dir, name, **kwargs):
log_dir = os.path.join(save_dir, name)
os.makedirs(log_dir, exist_ok=True)
self.logger = SummaryWriter(log_dir=log_dir, **kwargs)
######################
# training
######################
def on_train_start(self):
pass
def on_train_end(self):
pass
def on_epoch_start(self):
self.training_losses_meter = {'total_loss': AvgrageMeter()}
def on_epoch_end(self):
loss_outputs = {k: round(v.avg, 4) for k, v in self.training_losses_meter.items()}
print(f"Epoch {self.current_epoch} ended. Steps: {self.global_step}. {loss_outputs}")
def _training_step(self, sample, batch_idx, optimizer_idx):
"""
:param sample:
:param batch_idx:
:return: total loss: torch.Tensor, loss_log: dict
"""
raise NotImplementedError
def training_step(self, sample, batch_idx, optimizer_idx=-1):
"""
:param sample:
:param batch_idx:
:param optimizer_idx:
:return: {'loss': torch.Tensor, 'progress_bar': dict, 'tb_log': dict}
"""
loss_ret = self._training_step(sample, batch_idx, optimizer_idx)
if loss_ret is None:
return {'loss': None}
total_loss, log_outputs = loss_ret
log_outputs = tensors_to_scalars(log_outputs)
for k, v in log_outputs.items():
if k not in self.training_losses_meter:
self.training_losses_meter[k] = AvgrageMeter()
if not np.isnan(v):
self.training_losses_meter[k].update(v)
self.training_losses_meter['total_loss'].update(total_loss.item())
if optimizer_idx >= 0:
log_outputs[f'lr_{optimizer_idx}'] = self.trainer.optimizers[optimizer_idx].param_groups[0]['lr']
progress_bar_log = log_outputs
tb_log = {f'tr/{k}': v for k, v in log_outputs.items()}
return {
'loss': total_loss,
'progress_bar': progress_bar_log,
'tb_log': tb_log
}
def on_before_optimization(self, opt_idx):
if self.gradient_clip_norm > 0:
torch.nn.utils.clip_grad_norm_(self.parameters(), self.gradient_clip_norm)
if self.gradient_clip_val > 0:
torch.nn.utils.clip_grad_value_(self.parameters(), self.gradient_clip_val)
def on_after_optimization(self, epoch, batch_idx, optimizer, optimizer_idx):
if self.scheduler is not None:
# self.scheduler.step(self.global_step // hparams['accumulate_grad_batches'])
# the code above causes EPOCH_DEPRECATION_WARNING, changed it and changed the optimizer init with
# step_size divided by accumulate_grad_batches
self.scheduler.step()
######################
# validation
######################
def validation_start(self):
pass
def validation_step(self, sample, batch_idx):
"""
:param sample:
:param batch_idx:
:return: output: {"losses": {...}, "total_loss": float, ...} or (total loss: torch.Tensor, loss_log: dict)
"""
raise NotImplementedError
def validation_end(self, outputs):
"""
:param outputs:
:return: loss_output: dict
"""
all_losses_meter = {'total_loss': AvgrageMeter()}
for output in outputs:
if len(output) == 0 or output is None:
continue
if isinstance(output, dict):
assert 'losses' in output, 'Key "losses" should exist in validation output.'
n = output.pop('nsamples', 1)
losses = tensors_to_scalars(output['losses'])
total_loss = output.get('total_loss', sum(losses.values()))
else:
assert len(output) == 2, 'Validation output should only consist of two elements: (total_loss, losses)'
n = 1
total_loss, losses = output
losses = tensors_to_scalars(losses)
if isinstance(total_loss, torch.Tensor):
total_loss = total_loss.item()
for k, v in losses.items():
if k not in all_losses_meter:
all_losses_meter[k] = AvgrageMeter()
all_losses_meter[k].update(v, n)
all_losses_meter['total_loss'].update(total_loss, n)
loss_output = {k: round(v.avg, 4) for k, v in all_losses_meter.items()}
print(f"| Validation results@{self.global_step}: {loss_output}")
return {
'tb_log': {f'val/{k}': v for k, v in loss_output.items()},
'val_loss': loss_output['total_loss']
}
######################
# testing
######################
def test_start(self):
pass
def test_step(self, sample, batch_idx):
return self.validation_step(sample, batch_idx)
def test_end(self, outputs):
return self.validation_end(outputs)
######################
# start training/testing
######################
@classmethod
def start(cls):
os.environ['MASTER_PORT'] = str(random.randint(15000, 30000))
random.seed(hparams['seed'])
np.random.seed(hparams['seed'])
work_dir = hparams['work_dir']
trainer = Trainer(
work_dir=work_dir,
val_check_interval=hparams['val_check_interval'],
tb_log_interval=hparams['tb_log_interval'],
max_updates=hparams['max_updates'],
num_sanity_val_steps=hparams['num_sanity_val_steps'] if not hparams['validate'] else 10000,
accumulate_grad_batches=hparams['accumulate_grad_batches'],
print_nan_grads=hparams['print_nan_grads'],
resume_from_checkpoint=hparams.get('resume_from_checkpoint', 0),
amp=hparams['amp'],
monitor_key=hparams['valid_monitor_key'],
monitor_mode=hparams['valid_monitor_mode'],
num_ckpt_keep=hparams['num_ckpt_keep'],
save_best=hparams['save_best'],
seed=hparams['seed'],
debug=hparams['debug']
)
if not hparams['infer']: # train
trainer.fit(cls)
else:
trainer.test(cls)
def on_keyboard_interrupt(self):
pass
@@ -0,0 +1,68 @@
import glob
import os
import re
import torch
def get_last_checkpoint(work_dir, steps=None):
checkpoint = None
last_ckpt_path = None
ckpt_paths = get_all_ckpts(work_dir, steps)
if len(ckpt_paths) > 0:
last_ckpt_path = ckpt_paths[0]
checkpoint = torch.load(last_ckpt_path, map_location='cpu')
return checkpoint, last_ckpt_path
def get_all_ckpts(work_dir, steps=None):
if steps is None:
ckpt_path_pattern = f'{work_dir}/model_ckpt_steps_*.ckpt'
else:
ckpt_path_pattern = f'{work_dir}/model_ckpt_steps_{steps}.ckpt'
return sorted(glob.glob(ckpt_path_pattern),
key=lambda x: -int(re.findall('.*steps\_(\d+)\.ckpt', x)[0]))
def load_ckpt(cur_model, ckpt_base_dir, model_name='model', force=True, strict=True, verbose=True):
if os.path.isfile(ckpt_base_dir):
base_dir = os.path.dirname(ckpt_base_dir)
ckpt_path = ckpt_base_dir
checkpoint = torch.load(ckpt_base_dir, map_location='cpu')
else:
base_dir = ckpt_base_dir
checkpoint, ckpt_path = get_last_checkpoint(ckpt_base_dir)
if checkpoint is not None:
state_dict = checkpoint["state_dict"]
if len([k for k in state_dict.keys() if '.' in k]) > 0:
state_dict = {k[len(model_name) + 1:]: v for k, v in state_dict.items()
if k.startswith(f'{model_name}.')}
else:
if '.' not in model_name:
state_dict = state_dict[model_name]
else:
base_model_name = model_name.split('.')[0]
rest_model_name = model_name[len(base_model_name) + 1:]
state_dict = {
k[len(rest_model_name) + 1:]: v for k, v in state_dict[base_model_name].items()
if k.startswith(f'{rest_model_name}.')}
if not strict:
cur_model_state_dict = cur_model.state_dict()
unmatched_keys = []
for key, param in state_dict.items():
if key in cur_model_state_dict:
new_param = cur_model_state_dict[key]
if new_param.shape != param.shape:
unmatched_keys.append(key)
print("| Unmatched keys: ", key, new_param.shape, param.shape)
for key in unmatched_keys:
del state_dict[key]
# print(state_dict)
cur_model.load_state_dict(state_dict, strict=strict)
if verbose:
print(f"| load '{model_name}' from '{ckpt_path}'.")
else:
e_msg = f"| ckpt not found in {base_dir}."
if force:
assert False, e_msg
else:
print(e_msg)
@@ -0,0 +1,372 @@
import os
import sys
import traceback
import types
from functools import wraps
from itertools import chain
import numpy as np
import torch.utils.data
import torch.nn.functional as F
from torch.utils.data import ConcatDataset
from .hparams import hparams
def collate_1d_or_2d(values, pad_idx=0, left_pad=False, shift_right=False, max_len=None, shift_id=1):
if len(values[0].shape) == 1:
return collate_1d(values, pad_idx, left_pad, shift_right, max_len, shift_id)
else:
return collate_2d(values, pad_idx, left_pad, shift_right, max_len)
def collate_1d(values, pad_idx=0, left_pad=False, shift_right=False, max_len=None, shift_id=1):
"""Convert a list of 1d tensors into a padded 2d tensor."""
size = max(v.size(0) for v in values) if max_len is None else max_len
res = values[0].new(len(values), size).fill_(pad_idx)
def copy_tensor(src, dst):
assert dst.numel() == src.numel()
if shift_right:
dst[1:] = src[:-1]
dst[0] = shift_id
else:
dst.copy_(src)
for i, v in enumerate(values):
copy_tensor(v, res[i][size - len(v):] if left_pad else res[i][:len(v)])
return res
def collate_2d(values, pad_idx=0, left_pad=False, shift_right=False, max_len=None):
"""Convert a list of 2d tensors into a padded 3d tensor."""
size = max(v.size(0) for v in values) if max_len is None else max_len
res = values[0].new(len(values), size, values[0].shape[1]).fill_(pad_idx)
def copy_tensor(src, dst):
assert dst.numel() == src.numel()
if shift_right:
dst[1:] = src[:-1]
else:
dst.copy_(src)
for i, v in enumerate(values):
copy_tensor(v, res[i][size - len(v):] if left_pad else res[i][:len(v)])
return res
def collate_xd(values, pad_value=0, max_len=None):
size = ((max(v.size(0) for v in values) if max_len is None else max_len), *values[0].shape[1:])
res = torch.full((len(values), *size), fill_value=pad_value, dtype=values[0].dtype, device=values[0].device)
for i, v in enumerate(values):
res[i, :len(v), ...] = v
return res
def pad_or_cut_1d(values: torch.tensor, tgt_len, pad_value=0):
src_len = values.shape[0]
if src_len < tgt_len:
res = F.pad(values, [0, tgt_len - src_len], value=pad_value)
else:
res = values[:tgt_len]
return res
def pad_or_cut_2d(values: torch.tensor, tgt_len, dim=-1, pad_value=0):
if dim == 0 or dim == -2:
src_len = values.shape[0]
if src_len < tgt_len:
res = F.pad(values, [0, 0, 0, tgt_len - src_len], value=pad_value)
else:
res = values[:tgt_len]
elif dim == 1 or dim == -1:
src_len = values.shape[1]
if src_len < tgt_len:
res = F.pad(values, [0, tgt_len - src_len], value=pad_value)
else:
res = values[:, :tgt_len]
else:
raise RuntimeError(f"Wrong dim number {dim} while the tensor only has {len(values.shape)} dimensions.")
return res
def pad_or_cut_3d(values: torch.tensor, tgt_len, dim=-1, pad_value=0):
if dim == 0 or dim == -3:
src_len = values.shape[0]
if src_len < tgt_len:
res = F.pad(values, [0, 0, 0, 0, 0, tgt_len - src_len], value=pad_value)
else:
res = values[:tgt_len]
elif dim == 1 or dim == -2:
src_len = values.shape[1]
if src_len < tgt_len:
res = F.pad(values, [0, 0, 0, tgt_len - src_len], value=pad_value)
else:
res = values[:, :tgt_len]
elif dim == 2 or dim == -1:
src_len = values.shape[2]
if src_len < tgt_len:
res = F.pad(values, [0, tgt_len - src_len], value=pad_value)
else:
res = values[:, :, :tgt_len]
else:
raise RuntimeError(f"Wrong dim number {dim} while the tensor only has {len(values.shape)} dimensions.")
return res
def pad_or_cut_xd(values, tgt_len, dim=-1, pad_value=0):
if len(values.shape) == 1:
return pad_or_cut_1d(values, tgt_len, pad_value)
elif len(values.shape) == 2:
return pad_or_cut_2d(values, tgt_len, dim, pad_value)
elif len(values.shape) == 3:
return pad_or_cut_3d(values, tgt_len, dim, pad_value)
else:
raise NotImplementedError
def _is_batch_full(batch, num_tokens, max_tokens, max_sentences):
if len(batch) == 0:
return 0
if len(batch) == max_sentences:
return 1
if num_tokens > max_tokens:
return 1
return 0
def batch_by_size(
indices, num_tokens_fn, max_tokens=None, max_sentences=None,
required_batch_size_multiple=1, distributed=False
):
"""
Yield mini-batches of indices bucketed by size. Batches may contain
sequences of different lengths.
Args:
indices (List[int]): ordered list of dataset indices
num_tokens_fn (callable): function that returns the number of tokens at
a given index
max_tokens (int, optional): max number of tokens in each batch
(default: None).
max_sentences (int, optional): max number of sentences in each
batch (default: None).
required_batch_size_multiple (int, optional): require batch size to
be a multiple of N (default: 1).
"""
max_tokens = max_tokens if max_tokens is not None else sys.maxsize
max_sentences = max_sentences if max_sentences is not None else sys.maxsize
bsz_mult = required_batch_size_multiple
if isinstance(indices, types.GeneratorType):
indices = np.fromiter(indices, dtype=np.int64, count=-1)
sample_len = 0
sample_lens = []
batch = []
batches = []
for i in range(len(indices)):
idx = indices[i]
num_tokens = num_tokens_fn(idx)
sample_lens.append(num_tokens)
sample_len = max(sample_len, num_tokens)
assert sample_len <= max_tokens, (
"sentence at index {} of size {} exceeds max_tokens "
"limit of {}!".format(idx, sample_len, max_tokens)
)
num_tokens = (len(batch) + 1) * sample_len
if _is_batch_full(batch, num_tokens, max_tokens, max_sentences):
mod_len = max(
bsz_mult * (len(batch) // bsz_mult),
len(batch) % bsz_mult,
)
batches.append(batch[:mod_len])
batch = batch[mod_len:]
sample_lens = sample_lens[mod_len:]
sample_len = max(sample_lens) if len(sample_lens) > 0 else 0
batch.append(idx)
if len(batch) > 0:
batches.append(batch)
return batches
def build_dataloader(dataset, shuffle, max_tokens=None, max_sentences=None,
required_batch_size_multiple=-1, endless=False, apply_batch_by_size=True, pin_memory=False, use_ddp=False):
import torch.distributed as dist
devices_cnt = torch.cuda.device_count()
if devices_cnt == 0:
devices_cnt = 1
if not use_ddp:
devices_cnt = 1
if required_batch_size_multiple == -1:
required_batch_size_multiple = devices_cnt
def shuffle_batches(batches):
np.random.shuffle(batches)
return batches
if max_tokens is not None:
max_tokens *= devices_cnt
if max_sentences is not None:
max_sentences *= devices_cnt
indices = dataset.ordered_indices()
if apply_batch_by_size:
batch_sampler = batch_by_size(
indices, dataset.num_tokens, max_tokens=max_tokens, max_sentences=max_sentences,
required_batch_size_multiple=required_batch_size_multiple,
)
else:
batch_sampler = []
for i in range(0, len(indices), max_sentences):
batch_sampler.append(indices[i:i + max_sentences])
if shuffle:
batches = shuffle_batches(list(batch_sampler))
if endless:
batches = [b for _ in range(1000) for b in shuffle_batches(list(batch_sampler))]
else:
batches = batch_sampler
if endless:
batches = [b for _ in range(1000) for b in batches]
num_workers = dataset.num_workers
if use_ddp:
num_replicas = dist.get_world_size()
rank = dist.get_rank()
# batches = [x[rank::num_replicas] for x in batches if len(x) % num_replicas == 0]
# ensure that every sample in the dataset is covered
batches_ = []
for x in batches:
if len(x) % num_replicas == 0:
batches_.append(x[rank::num_replicas])
else:
x_ = x + [x[-1]] * (len(x) - len(x) // num_replicas * num_replicas)
batches_.append(x_[rank::num_replicas])
batches = batches_
return torch.utils.data.DataLoader(dataset,
collate_fn=dataset.collater,
batch_sampler=batches,
num_workers=num_workers,
pin_memory=pin_memory)
def unpack_dict_to_list(samples):
samples_ = []
bsz = samples.get('outputs').size(0)
for i in range(bsz):
res = {}
for k, v in samples.items():
try:
res[k] = v[i]
except:
pass
samples_.append(res)
return samples_
def remove_padding(x, padding_idx=0):
if x is None:
return None
assert len(x.shape) in [1, 2]
if len(x.shape) == 2: # [T, H]
return x[np.abs(x).sum(-1) != padding_idx]
elif len(x.shape) == 1: # [T]
return x[x != padding_idx]
def data_loader(fn):
"""
Decorator to make any fx with this use the lazy property
:param fn:
:return:
"""
wraps(fn)
attr_name = '_lazy_' + fn.__name__
def _get_data_loader(self):
try:
value = getattr(self, attr_name)
except AttributeError:
try:
value = fn(self) # Lazy evaluation, done only once.
except AttributeError as e:
# Guard against AttributeError suppression. (Issue #142)
traceback.print_exc()
error = f'{fn.__name__}: An AttributeError was encountered: ' + str(e)
raise RuntimeError(error) from e
setattr(self, attr_name, value) # Memoize evaluation.
return value
return _get_data_loader
class BaseDataset(torch.utils.data.Dataset):
def __init__(self, shuffle):
super().__init__()
self.hparams = hparams
self.shuffle = shuffle
self.sort_by_len = hparams['sort_by_len']
self.sizes = None
@property
def _sizes(self):
return self.sizes
def __getitem__(self, index):
raise NotImplementedError
def collater(self, samples):
raise NotImplementedError
def __len__(self):
return len(self._sizes)
def num_tokens(self, index):
return self.size(index)
def size(self, index):
"""Return an example's size as a float or tuple. This value is used when
filtering a dataset with ``--max-positions``."""
return min(self._sizes[index], hparams['max_frames'])
def ordered_indices(self):
"""Return an ordered list of indices. Batches will be constructed based
on this order."""
if self.shuffle:
indices = np.random.permutation(len(self))
if self.sort_by_len:
indices = indices[np.argsort(np.array(self._sizes)[indices], kind='mergesort')]
else:
indices = np.arange(len(self))
return indices.tolist()
@property
def num_workers(self):
return int(os.getenv('NUM_WORKERS', hparams['ds_workers']))
class BaseConcatDataset(ConcatDataset):
def collater(self, samples):
return self.datasets[0].collater(samples)
@property
def _sizes(self):
if not hasattr(self, 'sizes'):
self.sizes = list(chain.from_iterable([d._sizes for d in self.datasets]))
return self.sizes
def size(self, index):
return min(self._sizes[index], hparams['max_frames'])
def num_tokens(self, index):
return self.size(index)
def ordered_indices(self):
"""Return an ordered list of indices. Batches will be constructed based
on this order."""
if self.datasets[0].shuffle:
indices = np.random.permutation(len(self))
if self.datasets[0].sort_by_len:
indices = indices[np.argsort(np.array(self._sizes)[indices], kind='mergesort')]
else:
indices = np.arange(len(self))
return indices
@property
def num_workers(self):
return self.datasets[0].num_workers
@@ -0,0 +1,164 @@
from torch.nn.parallel import DistributedDataParallel
from torch.nn.parallel.distributed import _find_tensors
import torch.optim
import torch.utils.data
import torch
from packaging import version
class DDP(DistributedDataParallel):
"""
Override the forward call in lightning so it goes to training and validation step respectively
"""
def forward(self, *inputs, **kwargs): # pragma: no cover
# if version.parse(torch.__version__[:6]) < version.parse("1.11"):
if version.parse(torch.__version__) < version.parse("1.11"): # fix the hard [:6] problem
self._sync_params()
inputs, kwargs = self.scatter(inputs, kwargs, self.device_ids)
assert len(self.device_ids) == 1
if self.module.training:
output = self.module.training_step(*inputs[0], **kwargs[0])
elif self.module.testing:
output = self.module.test_step(*inputs[0], **kwargs[0])
else:
output = self.module.validation_step(*inputs[0], **kwargs[0])
if torch.is_grad_enabled():
# We'll return the output object verbatim since it is a freeform
# object. We need to find any tensors in this object, though,
# because we need to figure out which parameters were used during
# this forward pass, to ensure we short circuit reduction for any
# unused parameters. Only if `find_unused_parameters` is set.
if self.find_unused_parameters:
self.reducer.prepare_for_backward(list(_find_tensors(output)))
else:
self.reducer.prepare_for_backward([])
elif version.parse("1.11") <= version.parse(torch.__version__) < version.parse("2.0"):
from torch.nn.parallel.distributed import \
logging, Join, _DDPSink, _tree_flatten_with_rref, _tree_unflatten_with_rref
with torch.autograd.profiler.record_function("DistributedDataParallel.forward"):
if torch.is_grad_enabled() and self.require_backward_grad_sync:
self.logger.set_runtime_stats_and_log()
self.num_iterations += 1
self.reducer.prepare_for_forward()
# Notify the join context that this process has not joined, if
# needed
work = Join.notify_join_context(self)
if work:
self.reducer._set_forward_pass_work_handle(
work, self._divide_by_initial_world_size
)
# Calling _rebuild_buckets before forward compuation,
# It may allocate new buckets before deallocating old buckets
# inside _rebuild_buckets. To save peak memory usage,
# call _rebuild_buckets before the peak memory usage increases
# during forward computation.
# This should be called only once during whole training period.
if torch.is_grad_enabled() and self.reducer._rebuild_buckets():
logging.info("Reducer buckets have been rebuilt in this iteration.")
self._has_rebuilt_buckets = True
# sync params according to location (before/after forward) user
# specified as part of hook, if hook was specified.
buffer_hook_registered = hasattr(self, 'buffer_hook')
if self._check_sync_bufs_pre_fwd():
self._sync_buffers()
if self._join_config.enable:
# Notify joined ranks whether they should sync in backwards pass or not.
self._check_global_requires_backward_grad_sync(is_joined_rank=False)
# modified part
inputs, kwargs = self.scatter(inputs, kwargs, self.device_ids)
if self.module.training:
output = self.module.training_step(*inputs[0], **kwargs[0])
elif self.module.testing:
output = self.module.test_step(*inputs[0], **kwargs[0])
else:
output = self.module.validation_step(*inputs[0], **kwargs[0])
# sync params according to location (before/after forward) user
# specified as part of hook, if hook was specified.
if self._check_sync_bufs_post_fwd():
self._sync_buffers()
if torch.is_grad_enabled() and self.require_backward_grad_sync:
self.require_forward_param_sync = True
# We'll return the output object verbatim since it is a freeform
# object. We need to find any tensors in this object, though,
# because we need to figure out which parameters were used during
# this forward pass, to ensure we short circuit reduction for any
# unused parameters. Only if `find_unused_parameters` is set.
if self.find_unused_parameters and not self.static_graph:
# Do not need to populate this for static graph.
self.reducer.prepare_for_backward(list(_find_tensors(output)))
else:
self.reducer.prepare_for_backward([])
else:
self.require_forward_param_sync = False
# TODO: DDPSink is currently enabled for unused parameter detection and
# static graph training for first iteration.
if (self.find_unused_parameters and not self.static_graph) or (
self.static_graph and self.num_iterations == 1
):
state_dict = {
'static_graph': self.static_graph,
'num_iterations': self.num_iterations,
}
output_tensor_list, treespec, output_is_rref = _tree_flatten_with_rref(
output
)
output_placeholders = [None for _ in range(len(output_tensor_list))]
# Do not touch tensors that have no grad_fn, which can cause issues
# such as https://github.com/pytorch/pytorch/issues/60733
for i, output in enumerate(output_tensor_list):
if torch.is_tensor(output) and output.grad_fn is None:
output_placeholders[i] = output
# When find_unused_parameters=True, makes tensors which require grad
# run through the DDPSink backward pass. When not all outputs are
# used in loss, this makes those corresponding tensors receive
# undefined gradient which the reducer then handles to ensure
# param.grad field is not touched and we don't error out.
passthrough_tensor_list = _DDPSink.apply(
self.reducer,
state_dict,
*output_tensor_list,
)
for i in range(len(output_placeholders)):
if output_placeholders[i] is None:
output_placeholders[i] = passthrough_tensor_list[i]
# Reconstruct output data structure.
output = _tree_unflatten_with_rref(
output_placeholders, treespec, output_is_rref
)
else:
# now pytorch version >= 2.0
with torch.autograd.profiler.record_function("DistributedDataParallel.forward"):
inputs, kwargs = self._pre_forward(*inputs, **kwargs)
output = (
# self.module.forward(*inputs, **kwargs)
# if self._delay_all_reduce_all_params
# else self._run_ddp_forward(*inputs, **kwargs)
# modified: delete 'delay_all_reduce_named_params' function
self._run_ddp_forward(*inputs, **kwargs)
)
return self._post_forward(output)
return output
def _run_ddp_forward(self, *inputs, **kwargs):
if version.parse(torch.__version__) >= version.parse("2.0"):
with self._inside_ddp_forward():
if self.module.training:
output = self.module.training_step(*inputs, **kwargs)
elif self.module.testing:
output = self.module.test_step(*inputs, **kwargs)
else:
output = self.module.validation_step(*inputs, **kwargs)
return output # type: ignore[index]
else:
return super(DDP, self)._run_ddp_forward(*inputs, **kwargs)
@@ -0,0 +1,113 @@
import gc
import datetime
import inspect
import torch
import numpy as np
dtype_memory_size_dict = {
torch.float64: 64/8,
torch.double: 64/8,
torch.float32: 32/8,
torch.float: 32/8,
torch.float16: 16/8,
torch.half: 16/8,
torch.int64: 64/8,
torch.long: 64/8,
torch.int32: 32/8,
torch.int: 32/8,
torch.int16: 16/8,
torch.short: 16/6,
torch.uint8: 8/8,
torch.int8: 8/8,
}
# compatibility of torch1.0
if getattr(torch, "bfloat16", None) is not None:
dtype_memory_size_dict[torch.bfloat16] = 16/8
if getattr(torch, "bool", None) is not None:
dtype_memory_size_dict[torch.bool] = 8/8 # pytorch use 1 byte for a bool, see https://github.com/pytorch/pytorch/issues/41571
def get_mem_space(x):
try:
ret = dtype_memory_size_dict[x]
except KeyError:
print(f"dtype {x} is not supported!")
return ret
class MemTracker(object):
"""
Class used to track pytorch memory usage
Arguments:
detail(bool, default True): whether the function shows the detail gpu memory usage
path(str): where to save log file
verbose(bool, default False): whether show the trivial exception
device(int): GPU number, default is 0
"""
def __init__(self, detail=True, path='', verbose=False, device=0):
self.print_detail = detail
self.last_tensor_sizes = set()
self.gpu_profile_fn = path + f'{datetime.datetime.now():%d-%b-%y-%H:%M:%S}-gpu_mem_track.txt'
self.verbose = verbose
self.begin = True
self.device = device
def get_tensors(self):
for obj in gc.get_objects():
try:
if torch.is_tensor(obj) or (hasattr(obj, 'data') and torch.is_tensor(obj.data)):
tensor = obj
else:
continue
if tensor.is_cuda:
yield tensor
except Exception as e:
if self.verbose:
print('A trivial exception occured: {}'.format(e))
def get_tensor_usage(self):
sizes = [np.prod(np.array(tensor.size())) * get_mem_space(tensor.dtype) for tensor in self.get_tensors()]
return np.sum(sizes) / 1024**2
def get_allocate_usage(self):
return torch.cuda.memory_allocated() / 1024**2
def clear_cache(self):
gc.collect()
torch.cuda.empty_cache()
def print_all_gpu_tensor(self, file=None):
for x in self.get_tensors():
print(x.size(), x.dtype, np.prod(np.array(x.size()))*get_mem_space(x.dtype)/1024**2, file=file)
def track(self):
"""
Track the GPU memory usage
"""
frameinfo = inspect.stack()[1]
where_str = frameinfo.filename + ' line ' + str(frameinfo.lineno) + ': ' + frameinfo.function
with open(self.gpu_profile_fn, 'a+') as f:
if self.begin:
f.write(f"GPU Memory Track | {datetime.datetime.now():%d-%b-%y-%H:%M:%S} |"
f" Total Tensor Used Memory:{self.get_tensor_usage():<7.1f}Mb"
f" Total Allocated Memory:{self.get_allocate_usage():<7.1f}Mb\n\n")
self.begin = False
if self.print_detail is True:
ts_list = [(tensor.size(), tensor.dtype) for tensor in self.get_tensors()]
new_tensor_sizes = {(type(x),
tuple(x.size()),
ts_list.count((x.size(), x.dtype)),
np.prod(np.array(x.size()))*get_mem_space(x.dtype)/1024**2,
x.dtype) for x in self.get_tensors()}
for t, s, n, m, data_type in new_tensor_sizes - self.last_tensor_sizes:
f.write(f'+ | {str(n)} * Size:{str(s):<20} | Memory: {str(m*n)[:6]} M | {str(t):<20} | {data_type}\n')
for t, s, n, m, data_type in self.last_tensor_sizes - new_tensor_sizes:
f.write(f'- | {str(n)} * Size:{str(s):<20} | Memory: {str(m*n)[:6]} M | {str(t):<20} | {data_type}\n')
self.last_tensor_sizes = new_tensor_sizes
f.write(f"\nAt {where_str:<50}"
f" Total Tensor Used Memory:{self.get_tensor_usage():<7.1f}Mb"
f" Total Allocated Memory:{self.get_allocate_usage():<7.1f}Mb\n\n")
@@ -0,0 +1,131 @@
import argparse
import os
import yaml
global_print_hparams = True
hparams = {}
class Args:
def __init__(self, **kwargs):
for k, v in kwargs.items():
self.__setattr__(k, v)
def override_config(old_config: dict, new_config: dict):
for k, v in new_config.items():
if isinstance(v, dict) and k in old_config:
override_config(old_config[k], new_config[k])
else:
old_config[k] = v
def set_hparams(config='', exp_name='', hparams_str='', print_hparams=True, global_hparams=True, root_dir=''):
if config == '' and exp_name == '':
parser = argparse.ArgumentParser(description='')
parser.add_argument('--config', type=str, default='',
help='location of the data corpus')
parser.add_argument('--exp_name', type=str, default='', help='exp_name')
parser.add_argument('-hp', '--hparams', type=str, default='',
help='location of the data corpus')
parser.add_argument('--infer', action='store_true', help='infer')
parser.add_argument('--validate', action='store_true', help='validate')
parser.add_argument('--reset', action='store_true', help='reset hparams')
parser.add_argument('--remove', action='store_true', help='remove old ckpt')
parser.add_argument('--debug', action='store_true', help='debug')
parser.add_argument('--root_dir', type=str, default='', help='root directory of the project.')
args, unknown = parser.parse_known_args()
print("| Unknow hparams: ", unknown)
else:
args = Args(config=config, exp_name=exp_name, hparams=hparams_str,
infer=False, validate=False, reset=False, debug=False, remove=False, root_dir=root_dir)
global hparams
assert args.config != '' or args.exp_name != ''
root_dir = args.root_dir
if args.config != '':
assert os.path.exists(os.path.join(root_dir, args.config)), f'| Wrong config path! root_dir: {root_dir}, config_path: {args.config}'
config_chains = []
loaded_config = set()
def load_config(config_fn):
# deep first inheritance and avoid the second visit of one node
if not os.path.exists(os.path.join(root_dir, config_fn)):
return {}
with open(os.path.join(root_dir, config_fn)) as f:
hparams_ = yaml.safe_load(f)
loaded_config.add(config_fn)
if 'base_config' in hparams_:
ret_hparams = {}
if not isinstance(hparams_['base_config'], list):
hparams_['base_config'] = [hparams_['base_config']]
for c in hparams_['base_config']:
if c.startswith('.'):
c = f'{os.path.dirname(config_fn)}/{c}'
c = os.path.normpath(c)
if c not in loaded_config:
override_config(ret_hparams, load_config(c))
override_config(ret_hparams, hparams_)
else:
ret_hparams = hparams_
config_chains.append(config_fn)
return ret_hparams
saved_hparams = {}
args_work_dir = ''
if args.exp_name != '':
args_work_dir = os.path.join(root_dir, f'checkpoints/{args.exp_name}')
ckpt_config_path = f'{args_work_dir}/config.yaml'
if os.path.exists(ckpt_config_path):
with open(ckpt_config_path) as f:
saved_hparams_ = yaml.safe_load(f)
if saved_hparams_ is not None:
saved_hparams.update(saved_hparams_)
hparams_ = {}
if args.config != '':
hparams_.update(load_config(args.config))
if not args.reset:
hparams_.update(saved_hparams)
hparams_['work_dir'] = args_work_dir
# Support config overriding in command line. Support list type config overriding.
# Examples: --hparams="a=1,b.c=2,d=[1 1 1]"
if args.hparams != "":
for new_hparam in args.hparams.split(","):
k, v = new_hparam.split("=")
v = v.strip("\'\" ")
config_node = hparams_
for k_ in k.split(".")[:-1]:
config_node = config_node[k_]
k = k.split(".")[-1]
if v in ['True', 'False'] or type(config_node[k]) in [bool, list, dict]:
if type(config_node[k]) == list:
v = v.replace(" ", ",")
config_node[k] = eval(v)
else:
config_node[k] = type(config_node[k])(v)
if args_work_dir != '' and args.remove:
answer = input("REMOVE old checkpoint? Y/N [Default: N]: ")
if answer.lower() == "y":
pass
if args_work_dir != '' and (not os.path.exists(ckpt_config_path) or args.reset) and not args.infer:
os.makedirs(hparams_['work_dir'], exist_ok=True)
with open(ckpt_config_path, 'w') as f:
yaml.safe_dump(hparams_, f)
hparams_['infer'] = args.infer
hparams_['debug'] = args.debug
hparams_['validate'] = args.validate
hparams_['exp_name'] = args.exp_name
global global_print_hparams
if global_hparams:
hparams.clear()
hparams.update(hparams_)
if print_hparams and global_print_hparams and global_hparams:
# print('| Hparams chains: ', config_chains)
# print('| Hparams: ')
# for i, (k, v) in enumerate(sorted(hparams_.items())):
# print(f"\033[;33;m{k}\033[0m: {v}, ", end="\n" if i % 5 == 4 else "")
# print("")
global_print_hparams = False
return hparams_
@@ -0,0 +1,71 @@
import pickle
from copy import deepcopy
import numpy as np
class IndexedDataset:
def __init__(self, path, num_cache=1):
super().__init__()
self.path = path
self.data_file = None
self.data_offsets = np.load(f"{path}.idx", allow_pickle=True).item()['offsets']
self.data_file = open(f"{path}.data", 'rb', buffering=-1)
self.cache = []
self.num_cache = num_cache
def check_index(self, i):
if i < 0 or i >= len(self.data_offsets) - 1:
raise IndexError('index out of range')
def __del__(self):
if self.data_file:
self.data_file.close()
def __getitem__(self, i):
self.check_index(i)
if self.num_cache > 0:
for c in self.cache:
if c[0] == i:
return c[1]
self.data_file.seek(self.data_offsets[i])
b = self.data_file.read(self.data_offsets[i + 1] - self.data_offsets[i])
item = pickle.loads(b)
if self.num_cache > 0:
self.cache = [(i, deepcopy(item))] + self.cache[:-1]
return item
def __len__(self):
return len(self.data_offsets) - 1
class IndexedDatasetBuilder:
def __init__(self, path):
self.path = path
self.out_file = open(f"{path}.data", 'wb')
self.byte_offsets = [0]
def add_item(self, item):
s = pickle.dumps(item)
bytes = self.out_file.write(s)
self.byte_offsets.append(self.byte_offsets[-1] + bytes)
def finalize(self):
self.out_file.close()
np.save(open(f"{self.path}.idx", 'wb'), {'offsets': self.byte_offsets})
if __name__ == "__main__":
import random
from tqdm import tqdm
ds_path = '/tmp/indexed_ds_example'
size = 100
items = [{"a": np.random.normal(size=[10000, 10]),
"b": np.random.normal(size=[10000, 10])} for i in range(size)]
builder = IndexedDatasetBuilder(ds_path)
for i in tqdm(range(size)):
builder.add_item(items[i])
builder.finalize()
ds = IndexedDataset(ds_path)
for i in tqdm(range(10000)):
idx = random.randint(0, size - 1)
assert (ds[idx]['a'] == items[idx]['a']).all()
@@ -0,0 +1,53 @@
import torch
import torch.nn.functional as F
def sigmoid_focal_loss(
inputs: torch.Tensor,
targets: torch.Tensor,
alpha: float = 0.25,
gamma: float = 2,
reduction: str = "none",
) -> torch.Tensor:
"""
Loss used in RetinaNet for dense detection: https://arxiv.org/abs/1708.02002.
Args:
inputs (Tensor): A float tensor of arbitrary shape.
The predictions for each example.
targets (Tensor): A float tensor with the same shape as inputs. Stores the binary
classification label for each element in inputs
(0 for the negative class and 1 for the positive class).
alpha (float): Weighting factor in range (0,1) to balance
positive vs negative examples or -1 for ignore. Default: ``0.25``.
gamma (float): Exponent of the modulating factor (1 - p_t) to
balance easy vs hard examples. Default: ``2``.
reduction (string): ``'none'`` | ``'mean'`` | ``'sum'``
``'none'``: No reduction will be applied to the output.
``'mean'``: The output will be averaged.
``'sum'``: The output will be summed. Default: ``'none'``.
Returns:
Loss tensor with the reduction option applied.
"""
# Original implementation from https://github.com/facebookresearch/fvcore/blob/master/fvcore/nn/focal_loss.py
p = torch.sigmoid(inputs)
ce_loss = F.binary_cross_entropy_with_logits(inputs, targets, reduction="none")
p_t = p * targets + (1 - p) * (1 - targets)
loss = ce_loss * ((1 - p_t) ** gamma)
if alpha >= 0: # decrease the importance of negative samples
alpha_t = alpha * targets + (1 - alpha) * (1 - targets)
loss = alpha_t * loss
# Check reduction option and return loss accordingly
if reduction == "none":
pass
elif reduction == "mean":
loss = loss.mean()
elif reduction == "sum":
loss = loss.sum()
else:
raise ValueError(
f"Invalid Value for arg 'reduction': '{reduction} \n Supported reduction modes: 'none', 'mean', 'sum'"
)
return loss
@@ -0,0 +1,42 @@
import time
import torch
class AvgrageMeter(object):
def __init__(self):
self.reset()
def reset(self):
self.avg = 0
self.sum = 0
self.cnt = 0
def update(self, val, n=1):
self.sum += val * n
self.cnt += n
self.avg = self.sum / self.cnt
class Timer:
timer_map = {}
def __init__(self, name, enable=False):
if name not in Timer.timer_map:
Timer.timer_map[name] = 0
self.name = name
self.enable = enable
def __enter__(self):
if self.enable:
if torch.cuda.is_available():
torch.cuda.synchronize()
self.t = time.time()
def __exit__(self, exc_type, exc_val, exc_tb):
if self.enable:
if torch.cuda.is_available():
torch.cuda.synchronize()
Timer.timer_map[self.name] += time.time() - self.t
if self.enable:
print(f'[Timer] {self.name}: {Timer.timer_map[self.name]}')
@@ -0,0 +1,180 @@
import os
import traceback
from functools import partial
from tqdm import tqdm
import torch
def chunked_worker(worker_id, args_queue=None, results_queue=None, init_ctx_func=None):
ctx = init_ctx_func(worker_id) if init_ctx_func is not None else None
while True:
args = args_queue.get()
if args == '<KILL>':
return
job_idx, map_func, arg = args
try:
map_func_ = partial(map_func, ctx=ctx) if ctx is not None else map_func
if isinstance(arg, dict):
res = map_func_(**arg)
elif isinstance(arg, (list, tuple)):
res = map_func_(*arg)
else:
res = map_func_(arg)
results_queue.put((job_idx, res))
except:
traceback.print_exc()
results_queue.put((job_idx, None))
class MultiprocessManager:
def __init__(self, num_workers=None, init_ctx_func=None, multithread=False, queue_max=-1):
if multithread:
from multiprocessing.dummy import Queue, Process
else:
from multiprocessing import Queue, Process
if num_workers is None:
num_workers = int(os.getenv('N_PROC', os.cpu_count()))
self.num_workers = num_workers
self.results_queue = Queue(maxsize=-1)
self.jobs_pending = []
self.args_queue = Queue(maxsize=queue_max)
self.workers = []
self.total_jobs = 0
self.multithread = multithread
for i in range(num_workers):
if multithread:
p = Process(target=chunked_worker,
args=(i, self.args_queue, self.results_queue, init_ctx_func))
else:
p = Process(target=chunked_worker,
args=(i, self.args_queue, self.results_queue, init_ctx_func),
daemon=True)
self.workers.append(p)
p.start()
def add_job(self, func, args):
if not self.args_queue.full():
self.args_queue.put((self.total_jobs, func, args))
else:
self.jobs_pending.append((self.total_jobs, func, args))
self.total_jobs += 1
def get_results(self):
self.n_finished = 0
while self.n_finished < self.total_jobs:
while len(self.jobs_pending) > 0 and not self.args_queue.full():
self.args_queue.put(self.jobs_pending[0])
self.jobs_pending = self.jobs_pending[1:]
job_id, res = self.results_queue.get()
yield job_id, res
self.n_finished += 1
for w in range(self.num_workers):
self.args_queue.put("<KILL>")
for w in self.workers:
w.join()
def close(self):
if not self.multithread:
for w in self.workers:
w.terminate()
def __len__(self):
return self.total_jobs
def multiprocess_run_tqdm(map_func, args, num_workers=None, ordered=True, init_ctx_func=None,
multithread=False, queue_max=-1, desc=None):
for i, res in tqdm(
multiprocess_run(map_func, args, num_workers, ordered, init_ctx_func, multithread,
queue_max=queue_max),
total=len(args), desc=desc):
yield i, res
def multiprocess_run(map_func, args, num_workers=None, ordered=True, init_ctx_func=None, multithread=False,
queue_max=-1):
"""
Multiprocessing running chunked jobs.
Examples:
>>> for res in tqdm(multiprocess_run(job_func, args):
>>> print(res)
:param map_func:
:param args:
:param num_workers:
:param ordered:
:param init_ctx_func:
:param q_max_size:
:param multithread:
:return:
"""
if num_workers is None:
num_workers = int(os.getenv('N_PROC', os.cpu_count()))
manager = MultiprocessManager(num_workers, init_ctx_func, multithread, queue_max=queue_max)
for arg in args:
manager.add_job(map_func, arg)
if ordered:
n_jobs = len(args)
results = ['<WAIT>' for _ in range(n_jobs)]
i_now = 0
for job_i, res in manager.get_results():
results[job_i] = res
while i_now < n_jobs and (not isinstance(results[i_now], str) or results[i_now] != '<WAIT>'):
yield i_now, results[i_now]
results[i_now] = None
i_now += 1
else:
for job_i, res in manager.get_results():
yield job_i, res
manager.close()
# #### this is the old version of chunked_multiprocess_run
def chunked_worker_old(worker_id, map_func, args, results_queue=None, init_ctx_func=None):
ctx = init_ctx_func(worker_id) if init_ctx_func is not None else None
for job_idx, arg in args:
try:
if not isinstance(arg, tuple) and not isinstance(arg, list):
arg = [arg]
if ctx is not None:
res = map_func(*arg, ctx=ctx)
else:
res = map_func(*arg)
results_queue.put((job_idx, res))
except:
traceback.print_exc()
results_queue.put((job_idx, None))
def chunked_multiprocess_run(
map_func, args, num_workers=None, ordered=True,
init_ctx_func=None, q_max_size=1000, multithread=False):
if multithread:
from multiprocessing.dummy import Queue, Process
else:
from multiprocessing import Queue, Process
args = zip(range(len(args)), args)
args = list(args)
n_jobs = len(args)
if num_workers is None:
num_workers = int(os.getenv('N_PROC', os.cpu_count()))
results_queues = []
if ordered:
for i in range(num_workers):
results_queues.append(Queue(maxsize=q_max_size // num_workers))
else:
results_queue = Queue(maxsize=q_max_size)
for i in range(num_workers):
results_queues.append(results_queue)
workers = []
for i in range(num_workers):
args_worker = args[i::num_workers]
p = Process(target=chunked_worker_old, args=(
i, map_func, args_worker, results_queues[i], init_ctx_func), daemon=True)
workers.append(p)
p.start()
for n_finished in range(n_jobs):
results_queue = results_queues[n_finished % num_workers]
job_idx, res = results_queue.get()
assert job_idx == n_finished or not ordered, (job_idx, n_finished)
yield res
for w in workers:
w.join()
@@ -0,0 +1,74 @@
import numpy as np
import torch
import torch.nn as nn
def get_filter_2d(kernel, kernel_size, channels, no_grad=True):
# Reshape to 2d depthwise convolutional weight
kernel = kernel.view(1, 1, kernel_size, kernel_size)
kernel = kernel.repeat(channels, 1, 1, 1)
filter = nn.Conv2d(in_channels=channels, out_channels=channels, kernel_size=kernel_size, groups=channels,
bias=False, padding=kernel_size // 2)
filter.weight.data = kernel
if no_grad:
filter.weight.requires_grad = False
return filter
def get_filter_1d(kernel, kernel_size, channels, no_grad=True):
kernel = kernel.view(1, 1, kernel_size)
kernel = kernel.repeat(channels, 1, 1)
filter = nn.Conv1d(in_channels=channels, out_channels=channels, kernel_size=kernel_size, groups=channels,
bias=False, padding=kernel_size // 2)
filter.weight.data = kernel
if no_grad:
filter.weight.requires_grad = False
return filter
def get_gaussian_kernel_2d(kernel_size, sigma):
# Create a x, y coordinate grid of shape (kernel_size, kernel_size, 2)
x_coord = torch.arange(kernel_size)
x_grid = x_coord.repeat(kernel_size).view(kernel_size, kernel_size)
y_grid = x_grid.t()
xy_grid = torch.stack([x_grid, y_grid], dim=-1).float()
mean = (kernel_size - 1) / 2.
variance = sigma ** 2.
# Calculate the 2-dimensional gaussian kernel which is
# the product of two gaussian distributions for two different
# variables (in this case called x and y)
gaussian_kernel = (1. / (2. * np.pi * variance)) * torch.exp(
-torch.sum((xy_grid - mean) ** 2., dim=-1) / (2 * variance))
# Make sure sum of values in gaussian kernel equals 1.
gaussian_kernel = gaussian_kernel / torch.sum(gaussian_kernel)
return gaussian_kernel
def get_gaussian_kernel_1d(kernel_size, sigma):
x_grid = torch.arange(kernel_size)
mean = (kernel_size - 1) / 2.
variance = sigma ** 2.
gaussian_kernel = (1. / ((2. * np.pi) ** 0.5 * sigma)) * torch.exp(-(x_grid - mean) ** 2. / (2 * variance))
gaussian_kernel = gaussian_kernel / torch.sum(gaussian_kernel)
return gaussian_kernel
def get_hann_kernel_1d(kernel_size, periodic=False):
# periodic=False gives symmetric kernel, otherwise equivalent to hann(kernel_size + 1)
return torch.hann_window(kernel_size, periodic)
def get_triangle_kernel_1d(kernel_size):
kernel = torch.zeros(kernel_size)
for idx in range(kernel_size):
kernel[idx] = 1 - abs((idx - (kernel_size - 1) / 2) / ((kernel_size - 1) / 2))
return kernel
def add_gaussian_noise(tensor, mean=0, std=1):
noise = torch.randn(tensor.size()) * std + mean
noisy_tensor = tensor + noise
return noisy_tensor
@@ -0,0 +1,5 @@
import os
os.environ["OMP_NUM_THREADS"] = "1"
os.environ['TF_NUM_INTEROP_THREADS'] = '1'
os.environ['TF_NUM_INTRAOP_THREADS'] = '1'

Some files were not shown because too many files have changed in this diff Show More