Initial commit
This commit is contained in:
+38
@@ -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/
|
||||
@@ -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.
@@ -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)
|
||||
@@ -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.
@@ -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.
@@ -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.
@@ -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.
@@ -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.
@@ -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
|
||||
@@ -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
|
||||
@@ -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.
|
||||
@@ -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)
|
||||
@@ -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")
|
||||
@@ -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)
|
||||
@@ -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
|
||||
|
||||
@@ -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"))
|
||||
@@ -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.
|
||||
|
||||
  
|
||||
|
||||
## ✨ 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
|
||||
```
|
||||
@@ -0,0 +1,173 @@
|
||||
# 🎹 MIDI Editor - 网页端歌声 MIDI 编辑器
|
||||
|
||||
[English](README.md) | [简体中文](README_CN.md)
|
||||
|
||||
一个功能完整的网页端歌声 MIDI 文件编辑器,类似 ACE-Studio 和 VOCALOID。支持实时拖拽调整 MIDI 音符、歌词编辑、音频波形对齐,以及导入导出含歌词的 MIDI 文件。
|
||||
|
||||
  
|
||||
|
||||
## ✨ 功能特性
|
||||
|
||||
### 🎼 钢琴卷帘编辑
|
||||
|
||||
- **可视化音符编辑**:支持 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,
|
||||
},
|
||||
},
|
||||
])
|
||||
@@ -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>
|
||||
+4365
File diff suppressed because it is too large
Load Diff
@@ -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 |
@@ -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;
|
||||
}
|
||||
@@ -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' })
|
||||
}
|
||||
@@ -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 })),
|
||||
}))
|
||||
@@ -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()],
|
||||
})
|
||||
@@ -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)
|
||||
+113
@@ -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)
|
||||
+198
@@ -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
Reference in New Issue
Block a user