Initial commit

This commit is contained in:
王新升
2026-02-06 20:31:14 +08:00
parent a0b51be095
commit c589bcb837
145 changed files with 28773 additions and 0 deletions
+125
View File
@@ -0,0 +1,125 @@
# 🎵 SoulX-Singer-Preprocess
This part offers a comprehensive **singing transcription and editing toolkit** for real-world music audio. It provides the pipeline from vocal extraction to high-level annotation optimized for SVS dataset construction. By integrating state-of-the-art models, it transforms raw audio into structured singing data and supports the **customizable creation and editing of lyric-aligned MIDI scores**.
## ✨ Features
The toolkit includes the following core modules:
- 🎤 **Clean Dry Vocal Extraction**
Extracts the lead vocal track from polyphonic music audio and dereverberation.
- 📝 **Lyrics Transcription**
Automatically transcribes lyrics from clean vocal.
- 🎶 **Note Transcription**
Converts singing voice into note-level representations for SVS.
- 🎼 **MIDI Editor**
Supports customizable creation and editing of MIDI scores integrated with lyrics.
## 📁 Data Preparation
To ensure the data processing pipeline runs correctly, please verify that all required checkpoints are correctly placed in `pretrained_models/SoulX-Singer-Preprocess`
Before running the pipeline, prepare the following inputs:
- **Prompt audio**
Reference audio that provides timbre and style
- **Target audio**
Original vocal or music audio to be processed and transcribed.
Configure the corresponding parameters in:
```
example/preprocess.sh
```
Typical configuration includes:
- Input / output paths
- Module enable switches
## 🚀 Usage
After configuring `preprocess.sh`, run the transcription pipeline with:
```bash
bash example/preprocess.sh
```
The script will automatically execute the following steps:
1. **Vocal separation and dereverberation**
2. **F0 extraction and voice activity detection (VAD)**
3. **Lyrics transcription**
4. **Note transcription**
---
After the pipeline completes, you will obtain **SoulX-Singer–style metadata** that can be directly used for Singing Voice Synthesis (SVS).
⚠️ **Important Note**
Transcription errors—especially in **lyrics** and **note annotations**—can significantly affect the final SVS quality. We **strongly recommend manually reviewing and correcting** the generated metadata before inference.
To support this, we provide a **MIDI Editor** for editing lyrics, phoneme alignment, note pitches, and durations. The workflow is:
**Export metadata to MIDI** → edit in the MIDI Editor → **Import edited MIDI back to metadata** for SVS.
---
#### Step 1: Metadata → MIDI (for editing)
Convert SoulX-Singer metadata to a MIDI file so you can open it in the MIDI Editor:
```bash
preprocess_root=example/transcriptions/music
python -m preprocess.tools.midi_parser \
--meta2midi \
--meta "${preprocess_root}/metadata.json" \
--midi "${preprocess_root}/vocal.mid"
```
#### Step 2: Edit in the MIDI Editor
Open the MIDI Editor (see [MIDI Editor Tutorial](tools/midi_editor/README.md)), load `vocal.mid`, and correct lyrics, pitches, or durations as needed. Save the result as e.g. `vocal_edited.mid`.
#### Step 3: MIDI → Metadata (for SoulX-Singer inference)
Convert the edited MIDI back into SoulX-Singer-style metadata (and cut wavs) for SVS:
```bash
python -m preprocess.tools.midi_parser \
--midi2meta \
--midi "${preprocess_root}/vocal_edited.mid" \
--meta "${preprocess_root}/edit_metadata.json" \
--vocal "${preprocess_root}/vocal.wav" \
```
Use `edit_metadata.json` (and the wavs under `edit_cut_wavs`) as the target metadata in your inference pipeline.
## 🔗 References & Dependencies
This project builds upon the following excellent open-source works:
### 🎧 Vocal Separation & Dereverberation
- [Music Source Separation Training](https://github.com/ZFTurbo/Music-Source-Separation-Training)
- [Lead Vocal Separation](https://huggingface.co/becruily/mel-band-roformer-karaoke)
- [Vocal Dereverberation](https://huggingface.co/anvuew/dereverb_mel_band_roformer)
### 🎼 F0 Extraction
- [RMVPE](https://github.com/Dream-High/RMVPE)
### 📝 Lyrics Transcription (ASR)
- [Paraformer](https://modelscope.cn/models/iic/speech_seaco_paraformer_large_asr_nat-zh-cn-16k-common-vocab8404-pytorch)
- [Parakeet-tdt-0.6b-v2](https://huggingface.co/nvidia/parakeet-tdt-0.6b-v2)
### 🎶 Note Transcription
- [ROSVOT](https://github.com/RickyL-2000/ROSVOT)
We sincerely thank the authors of these repositories for their exceptional open-source contributions, which have been fundamental to the development of this toolkit.
+146
View File
@@ -0,0 +1,146 @@
import json
import shutil
import soundfile as sf
from pathlib import Path
import librosa
from preprocess.utils import convert_metadata, merge_short_segments
from preprocess.tools import (
F0Extractor,
VocalDetector,
VocalSeparator,
NoteTranscriber,
LyricTranscriber,
)
class PreprocessPipeline:
def __init__(self, device: str, language: str, save_dir: str, vocal_sep: bool = True, max_merge_duration: int = 60000):
self.device = device
self.language = language
self.save_dir = save_dir
self.vocal_sep = vocal_sep
self.max_merge_duration = max_merge_duration
if vocal_sep:
self.vocal_separator = VocalSeparator(
sep_model_path="pretrained_models/SoulX-Singer-Preprocess/mel-band-roformer-karaoke/mel_band_roformer_karaoke_becruily.ckpt",
sep_config_path="pretrained_models/SoulX-Singer-Preprocess/mel-band-roformer-karaoke/config_karaoke_becruily.yaml",
der_model_path="pretrained_models/SoulX-Singer-Preprocess/dereverb_mel_band_roformer/dereverb_mel_band_roformer_anvuew_sdr_19.1729.ckpt",
der_config_path="pretrained_models/SoulX-Singer-Preprocess/dereverb_mel_band_roformer/dereverb_mel_band_roformer_anvuew.yaml",
device=device
)
else:
self.vocal_separator = None
self.f0_extractor = F0Extractor(
model_path="pretrained_models/SoulX-Singer-Preprocess/rmvpe/rmvpe.pt",
device=device,
)
self.vocal_detector = VocalDetector(
cut_wavs_output_dir= f"{save_dir}/cut_wavs",
)
self.lyric_transcriber = LyricTranscriber(
zh_model_path="pretrained_models/SoulX-Singer-Preprocess/speech_seaco_paraformer_large_asr_nat-zh-cn-16k-common-vocab8404-pytorch",
en_model_path="pretrained_models/SoulX-Singer-Preprocess/parakeet-tdt-0.6b-v2/parakeet-tdt-0.6b-v2.nemo",
device=device
)
self.note_transcriber = NoteTranscriber(
rosvot_model_path="pretrained_models/SoulX-Singer-Preprocess/rosvot/rosvot/model.pt",
rwbd_model_path="pretrained_models/SoulX-Singer-Preprocess/rosvot/rwbd/model.pt",
device=device
)
def run(
self,
audio_path: str,
vocal_sep: bool = True,
max_merge_duration: int = 60000,
language: str = "Mandarin"
) -> None:
vocal_sep = self.vocal_sep if vocal_sep is None else vocal_sep
max_merge_duration = self.max_merge_duration if max_merge_duration is None else max_merge_duration
language = self.language if language is None else language
output_dir = Path(self.save_dir)
output_dir.mkdir(parents=True, exist_ok=True)
if vocal_sep:
# Perform vocal/accompaniment separation
sep = self.vocal_separator.process(audio_path)
vocal = sep.vocals_dereverbed.T
acc = sep.accompaniment.T
sample_rate = sep.sample_rate
vocal_path = output_dir / "vocal.wav"
acc_path = output_dir / "acc.wav"
sf.write(vocal_path, vocal, sample_rate)
sf.write(acc_path, acc, sample_rate)
else:
# Use the original audio as vocal source (no separation)
vocal, sample_rate = librosa.load(audio_path, sr=None, mono=True)
vocal_path = output_dir / "vocal.wav"
sf.write(vocal_path, vocal, sample_rate)
vocal_f0 = self.f0_extractor.process(str(vocal_path))
segments = self.vocal_detector.process(str(vocal_path), f0=vocal_f0)
metadata = []
for seg in segments:
self.f0_extractor.process(seg["wav_fn"], f0_path=seg["wav_fn"].replace(".wav", "_f0.npy"))
words, durs = self.lyric_transcriber.process(
seg["wav_fn"], language
)
seg["words"] = words
seg["word_durs"] = durs
seg["language"] = language
metadata.append(
self.note_transcriber.process(seg, segment_info=seg)
)
merged = merge_short_segments(
vocal,
sample_rate,
metadata,
output_dir / "long_cut_wavs",
max_duration_ms=max_merge_duration,
)
final_metadata = []
for item in merged:
self.f0_extractor.process(item.wav_fn, f0_path=item.wav_fn.replace(".wav", "_f0.npy"))
final_metadata.append(convert_metadata(item))
with open(output_dir / "metadata.json", "w", encoding="utf-8") as f:
json.dump(final_metadata, f, ensure_ascii=False, indent=2)
shutil.copy(output_dir / "metadata.json", audio_path.replace(".wav", ".json").replace(".mp3", ".json").replace(".flac", ".json"))
def main(args):
pipeline = PreprocessPipeline(
device=args.device,
language=args.language,
save_dir=args.save_dir,
vocal_sep=args.vocal_sep,
max_merge_duration=args.max_merge_duration,
)
pipeline.run(
audio_path=args.audio_path,
language=args.language
)
if __name__ == "__main__":
import argparse
parser = argparse.ArgumentParser()
parser.add_argument("--audio_path", type=str, required=True, help="Path to the input audio file")
parser.add_argument("--save_dir", type=str, required=True, help="Directory to save the output files")
parser.add_argument("--language", type=str, default="Mandarin", help="Language of the audio")
parser.add_argument("--device", type=str, default="cuda:0", help="Device to run the models on")
parser.add_argument("--vocal_sep", type=bool, default=True, help="Whether to perform vocal separation")
parser.add_argument("--max_merge_duration", type=int, default=60000, help="Maximum merged segment duration in milliseconds")
args = parser.parse_args()
main(args)
+53
View File
@@ -0,0 +1,53 @@
"""Preprocess tools.
This package provides a thin, stable import surface for common preprocess components.
Examples:
from preprocess.tools import (
F0Extractor,
PitchExtractor,
VocalDetectionModel,
VocalSeparationModel,
VocalExtractionModel,
NoteTranscriptionModel,
LyricTranscriptionModel,
)
Note:
Keep these imports lightweight. If a tool pulls heavy dependencies at import time,
consider switching to lazy imports.
"""
from __future__ import annotations
# Core tools
from .f0_extraction import F0Extractor
from .vocal_detection import VocalDetector
# Some tools may live outside this package in different layouts across branches.
# Keep the public surface stable while avoiding hard import failures.
try:
from .vocal_separation.model import VocalSeparator # type: ignore
except Exception: # pragma: no cover
VocalSeparator = None # type: ignore
try:
from .note_transcription.model import NoteTranscriber # type: ignore
except Exception: # pragma: no cover
NoteTranscriber = None # type: ignore
try:
from .lyric_transcription import LyricTranscriber
except Exception: # pragma: no cover
LyricTranscriber = None # type: ignore
__all__ = [
"F0Extractor",
"VocalDetector",
]
if VocalSeparator is not None:
__all__.append("VocalSeparator")
if LyricTranscriber is not None:
__all__.append("LyricTranscriber")
if NoteTranscriber is not None:
__all__.append("NoteTranscriber")
+527
View File
@@ -0,0 +1,527 @@
# https://github.com/Dream-High/RMVPE
import math
import time
import librosa
import numpy as np
from librosa.filters import mel
from scipy.interpolate import interp1d
from typing import Optional
import torch
import torch.nn as nn
import torch.nn.functional as F
class BiGRU(nn.Module):
def __init__(self, input_features, hidden_features, num_layers):
super(BiGRU, self).__init__()
self.gru = nn.GRU(
input_features,
hidden_features,
num_layers=num_layers,
batch_first=True,
bidirectional=True,
)
def forward(self, x):
return self.gru(x)[0]
class ConvBlockRes(nn.Module):
def __init__(self, in_channels, out_channels, momentum=0.01):
super(ConvBlockRes, self).__init__()
self.conv = nn.Sequential(
nn.Conv2d(
in_channels=in_channels,
out_channels=out_channels,
kernel_size=(3, 3),
stride=(1, 1),
padding=(1, 1),
bias=False,
),
nn.BatchNorm2d(out_channels, momentum=momentum),
nn.ReLU(),
nn.Conv2d(
in_channels=out_channels,
out_channels=out_channels,
kernel_size=(3, 3),
stride=(1, 1),
padding=(1, 1),
bias=False,
),
nn.BatchNorm2d(out_channels, momentum=momentum),
nn.ReLU(),
)
if in_channels != out_channels:
self.shortcut = nn.Conv2d(in_channels, out_channels, (1, 1))
def forward(self, x):
if not hasattr(self, "shortcut"):
return self.conv(x) + x
else:
return self.conv(x) + self.shortcut(x)
class ResEncoderBlock(nn.Module):
def __init__(self, in_channels, out_channels, kernel_size, n_blocks=1, momentum=0.01):
super(ResEncoderBlock, self).__init__()
self.n_blocks = n_blocks
self.conv = nn.ModuleList()
self.conv.append(ConvBlockRes(in_channels, out_channels, momentum))
for i in range(n_blocks - 1):
self.conv.append(ConvBlockRes(out_channels, out_channels, momentum))
self.kernel_size = kernel_size
if self.kernel_size is not None:
self.pool = nn.AvgPool2d(kernel_size=kernel_size)
def forward(self, x):
for conv in self.conv:
x = conv(x)
if self.kernel_size is not None:
return x, self.pool(x)
else:
return x
class Encoder(nn.Module):
def __init__(self, in_channels, in_size, n_encoders, kernel_size, n_blocks, out_channels=16, momentum=0.01):
super(Encoder, self).__init__()
self.n_encoders = n_encoders
self.bn = nn.BatchNorm2d(in_channels, momentum=momentum)
self.layers = nn.ModuleList()
self.latent_channels = []
for i in range(self.n_encoders):
self.layers.append(
ResEncoderBlock(in_channels, out_channels, kernel_size, n_blocks, momentum=momentum)
)
self.latent_channels.append([out_channels, in_size])
in_channels = out_channels
out_channels *= 2
in_size //= 2
self.out_size = in_size
self.out_channel = out_channels
def forward(self, x):
concat_tensors = []
x = self.bn(x)
for layer in self.layers:
t, x = layer(x)
concat_tensors.append(t)
return x, concat_tensors
class Intermediate(nn.Module):
def __init__(self, in_channels, out_channels, n_inters, n_blocks, momentum=0.01):
super(Intermediate, self).__init__()
self.n_inters = n_inters
self.layers = nn.ModuleList()
self.layers.append(ResEncoderBlock(in_channels, out_channels, None, n_blocks, momentum))
for i in range(self.n_inters - 1):
self.layers.append(ResEncoderBlock(out_channels, out_channels, None, n_blocks, momentum))
def forward(self, x):
for layer in self.layers:
x = layer(x)
return x
class ResDecoderBlock(nn.Module):
def __init__(self, in_channels, out_channels, stride, n_blocks=1, momentum=0.01):
super(ResDecoderBlock, self).__init__()
out_padding = (0, 1) if stride == (1, 2) else (1, 1)
self.n_blocks = n_blocks
self.conv1 = nn.Sequential(
nn.ConvTranspose2d(
in_channels=in_channels,
out_channels=out_channels,
kernel_size=(3, 3),
stride=stride,
padding=(1, 1),
output_padding=out_padding,
bias=False,
),
nn.BatchNorm2d(out_channels, momentum=momentum),
nn.ReLU(),
)
self.conv2 = nn.ModuleList()
self.conv2.append(ConvBlockRes(out_channels * 2, out_channels, momentum))
for i in range(n_blocks - 1):
self.conv2.append(ConvBlockRes(out_channels, out_channels, momentum))
def forward(self, x, concat_tensor):
x = self.conv1(x)
x = torch.cat((x, concat_tensor), dim=1)
for conv2 in self.conv2:
x = conv2(x)
return x
class Decoder(nn.Module):
def __init__(self, in_channels, n_decoders, stride, n_blocks, momentum=0.01):
super(Decoder, self).__init__()
self.layers = nn.ModuleList()
self.n_decoders = n_decoders
for i in range(self.n_decoders):
out_channels = in_channels // 2
self.layers.append(
ResDecoderBlock(in_channels, out_channels, stride, n_blocks, momentum)
)
in_channels = out_channels
def forward(self, x, concat_tensors):
for i, layer in enumerate(self.layers):
x = layer(x, concat_tensors[-1 - i])
return x
class DeepUnet(nn.Module):
def __init__(self, kernel_size, n_blocks, en_de_layers=5, inter_layers=4, in_channels=1, en_out_channels=16):
super(DeepUnet, self).__init__()
self.encoder = Encoder(in_channels, 128, en_de_layers, kernel_size, n_blocks, en_out_channels)
self.intermediate = Intermediate(
self.encoder.out_channel // 2,
self.encoder.out_channel,
inter_layers,
n_blocks,
)
self.decoder = Decoder(self.encoder.out_channel, en_de_layers, kernel_size, n_blocks)
def forward(self, x):
x, concat_tensors = self.encoder(x)
x = self.intermediate(x)
x = self.decoder(x, concat_tensors)
return x
class E2E(nn.Module):
def __init__(self, n_blocks, n_gru, kernel_size, en_de_layers=5, inter_layers=4, in_channels=1, en_out_channels=16):
super(E2E, self).__init__()
self.unet = DeepUnet(kernel_size, n_blocks, en_de_layers, inter_layers, in_channels, en_out_channels)
self.cnn = nn.Conv2d(en_out_channels, 3, (3, 3), padding=(1, 1))
if n_gru:
self.fc = nn.Sequential(
BiGRU(3 * 128, 256, n_gru),
nn.Linear(512, 360),
nn.Dropout(0.25),
nn.Sigmoid(),
)
else:
self.fc = nn.Sequential(
nn.Linear(3 * 128, 360),
nn.Dropout(0.25),
nn.Sigmoid()
)
def forward(self, mel):
mel = mel.transpose(-1, -2).unsqueeze(1)
x = self.cnn(self.unet(mel)).transpose(1, 2).flatten(-2)
x = self.fc(x)
return x
class MelSpectrogram(torch.nn.Module):
def __init__(self, is_half, n_mel_channels, sampling_rate, win_length, hop_length,
n_fft=None, mel_fmin=0, mel_fmax=None, clamp=1e-5):
super().__init__()
n_fft = win_length if n_fft is None else n_fft
self.hann_window = {}
mel_basis = mel(
sr=sampling_rate,
n_fft=n_fft,
n_mels=n_mel_channels,
fmin=mel_fmin,
fmax=mel_fmax,
htk=True,
)
mel_basis = torch.from_numpy(mel_basis).float()
self.register_buffer("mel_basis", mel_basis)
self.n_fft = win_length if n_fft is None else n_fft
self.hop_length = hop_length
self.win_length = win_length
self.sampling_rate = sampling_rate
self.n_mel_channels = n_mel_channels
self.clamp = clamp
self.is_half = is_half
def forward(self, audio, keyshift=0, speed=1, center=True):
factor = 2 ** (keyshift / 12)
n_fft_new = int(np.round(self.n_fft * factor))
win_length_new = int(np.round(self.win_length * factor))
hop_length_new = int(np.round(self.hop_length * speed))
keyshift_key = str(keyshift) + "_" + str(audio.device)
if keyshift_key not in self.hann_window:
self.hann_window[keyshift_key] = torch.hann_window(win_length_new).to(audio.device)
fft = torch.stft(
audio,
n_fft=n_fft_new,
hop_length=hop_length_new,
win_length=win_length_new,
window=self.hann_window[keyshift_key],
center=center,
return_complex=True,
)
magnitude = torch.sqrt(fft.real.pow(2) + fft.imag.pow(2))
if keyshift != 0:
size = self.n_fft // 2 + 1
resize = magnitude.size(1)
if resize < size:
magnitude = F.pad(magnitude, (0, 0, 0, size - resize))
magnitude = magnitude[:, :size, :] * self.win_length / win_length_new
mel_output = torch.matmul(self.mel_basis, magnitude)
if self.is_half:
mel_output = mel_output.half()
log_mel_spec = torch.log(torch.clamp(mel_output, min=self.clamp))
return log_mel_spec
class RMVPE:
def __init__(self, model_path: str, is_half, device=None):
self.is_half = is_half
if device is None:
device = "cuda:0" if torch.cuda.is_available() else "cpu"
self.device = torch.device(device) if isinstance(device, str) else device
self.mel_extractor = MelSpectrogram(
is_half=is_half,
n_mel_channels=128,
sampling_rate=16000,
win_length=1024,
hop_length=160,
n_fft=None,
mel_fmin=30,
mel_fmax=8000
).to(self.device)
model = E2E(n_blocks=4, n_gru=1, kernel_size=(2, 2))
ckpt = torch.load(model_path, map_location=self.device)
model.load_state_dict(ckpt)
model.eval()
if is_half:
model = model.half()
else:
model = model.float()
self.model = model.to(self.device)
cents_mapping = 20 * np.arange(360) + 1997.3794084376191
self.cents_mapping = np.pad(cents_mapping, (4, 4)) # 368
def mel2hidden(self, mel):
with torch.no_grad():
n_frames = mel.shape[-1]
n_pad = 32 * ((n_frames - 1) // 32 + 1) - n_frames
if n_pad > 0:
mel = F.pad(mel, (0, n_pad), mode="constant")
mel = mel.half() if self.is_half else mel.float()
hidden = self.model(mel)
return hidden[:, :n_frames]
def decode(self, hidden, thred=0.03):
cents_pred = self.to_local_average_cents(hidden, thred=thred)
f0 = 10 * (2 ** (cents_pred / 1200))
f0[f0 == 10] = 0
return f0
def infer_from_audio(self, audio, thred=0.03):
if not torch.is_tensor(audio):
audio = torch.from_numpy(audio)
mel = self.mel_extractor(audio.float().to(self.device).unsqueeze(0), center=True)
hidden = self.mel2hidden(mel)
hidden = hidden.squeeze(0).cpu().numpy()
if self.is_half:
hidden = hidden.astype("float32")
f0 = self.decode(hidden, thred=thred)
return f0
def to_local_average_cents(self, salience, thred=0.05):
center = np.argmax(salience, axis=1)
salience = np.pad(salience, ((0, 0), (4, 4)))
center += 4
todo_salience = []
todo_cents_mapping = []
starts = center - 4
ends = center + 5
for idx in range(salience.shape[0]):
todo_salience.append(salience[:, starts[idx]:ends[idx]][idx])
todo_cents_mapping.append(self.cents_mapping[starts[idx]:ends[idx]])
todo_salience = np.array(todo_salience)
todo_cents_mapping = np.array(todo_cents_mapping)
product_sum = np.sum(todo_salience * todo_cents_mapping, 1)
weight_sum = np.sum(todo_salience, 1)
devided = product_sum / weight_sum
maxx = np.max(salience, axis=1)
devided[maxx <= thred] = 0
return devided
class F0Extractor:
"""Extract frame-level f0 from singing voice.
Wrapper around an RMVPE network that:
1) loads the checkpoint once in ``__init__``
2) exposes a simple :py:meth:`process` API and optionally saves ``*_f0.npy``.
"""
def __init__(
self,
model_path: str,
device: str = "cpu",
*,
is_half: bool = False,
input_sr: int = 16000,
target_sr: int = 24000,
hop_size: int = 480,
max_duration: float = 300,
thred: float = 0.03,
verbose: bool = True,
):
"""Initialize the f0 extractor.
Args:
model_path: Path to RMVPE checkpoint.
device: Torch device string, e.g. ``"cuda:0"`` / ``"cpu"``.
is_half: Whether to run the model in fp16.
input_sr: Input resample rate used by RMVPE frontend.
target_sr: Target sample rate for the output f0 grid.
hop_size: Target hop size for the output f0 grid.
max_duration: Max duration (seconds) for interpolation grid.
thred: Voicing threshold used when decoding salience.
verbose: Whether to print verbose logs.
"""
self.model_path = model_path
self.input_sr = input_sr
self.target_sr = target_sr
self.hop_size = hop_size
self.max_duration = max_duration
self.thred = thred
self.verbose = verbose
self.model = RMVPE(model_path, is_half=is_half, device=device)
if self.verbose:
print(
"[f0 extraction] init success:",
f"device={device}",
f"model_path={model_path}",
f"is_half={is_half}",
f"input_sr={input_sr}",
f"target_sr={target_sr}",
f"hop_size={hop_size}",
f"thred={thred}",
)
@staticmethod
def interpolate_f0(
f0_16k: np.ndarray,
original_length: int,
original_sr: int,
*,
target_sr: int = 48000,
hop_size: int = 256,
max_duration: float = 20.0,
) -> np.ndarray:
"""Interpolate f0 from RMVPE's 16k hop grid to target mel hop grid."""
mel_target_sr = target_sr
mel_hop_size = hop_size
mel_max_duration = max_duration
batch_max_length = int(mel_max_duration * mel_target_sr / mel_hop_size)
duration_in_seconds = original_length / original_sr
effective_target_length = int(duration_in_seconds * mel_target_sr)
original_frames = math.ceil(effective_target_length / mel_hop_size)
target_frames = min(original_frames, batch_max_length)
rmvpe_hop = 160
t_16k = np.arange(len(f0_16k)) * (rmvpe_hop / 16000.0)
t_target = np.arange(target_frames) * (mel_hop_size / float(mel_target_sr))
if len(f0_16k) > 0:
f_interp = interp1d(
t_16k,
f0_16k,
kind="linear",
bounds_error=False,
fill_value=0.0,
assume_sorted=True,
)
f0 = f_interp(t_target)
else:
f0 = np.zeros(target_frames)
if len(f0) != target_frames:
f0 = (
f0[:target_frames]
if len(f0) > target_frames
else np.pad(f0, (0, target_frames - len(f0)), "constant")
)
return f0
def process(self, audio_path: str, *, f0_path: str | None = None, verbose: Optional[bool] = None) -> np.ndarray:
"""Run f0 extraction for a single wav.
Args:
audio_path: Path to the input wav file.
f0_path: if is not None, save the f0 data to this path.
verbose: Override instance-level verbose flag for this call.
Returns:
np.ndarray: shape ``[T]``, f0 in Hz (0 for unvoiced).
"""
verbose = self.verbose if verbose is None else verbose
if verbose:
print(f"[f0 extraction] process: start: {audio_path}")
t0 = time.time()
audio, _ = librosa.load(audio_path, sr=self.input_sr)
f0_16k = self.model.infer_from_audio(audio, thred=self.thred)
f0 = self.interpolate_f0(
f0_16k,
original_length=audio.shape[-1],
original_sr=self.input_sr,
target_sr=self.target_sr,
hop_size=self.hop_size,
max_duration=self.max_duration,
)
if verbose:
dt = time.time() - t0
voiced_ratio = float(np.mean(f0 > 0)) if len(f0) else 0.0
print(
"[f0 extraction] process: done:",
f"frames={len(f0)}",
f"voiced_ratio={voiced_ratio:.3f}",
f"time={dt:.3f}s",
)
if f0_path is not None:
np.save(f0_path, f0)
return f0
if __name__ == "__main__":
model_path = (
"pretrained_models/rmvpe/rmvpe.pt"
)
audio_path = "./outputs/transcription/test.wav"
pe = F0Extractor(
model_path,
device="cuda",
)
f0 = pe.process(audio_path)
+72
View File
@@ -0,0 +1,72 @@
import re
import ToJyutping
from g2pM import G2pM
from g2p_en import G2p as G2pE
_EN_WORD_RE = re.compile(r"^[A-Za-z]+(?:'[A-Za-z]+)*$")
_ZH_WORD_RE = re.compile(r"[\u4e00-\u9fff]")
EN_FLAG = "en_"
YUE_FLAG = "yue_"
ZH_FLAG = "zh_"
g2p_zh = G2pM()
g2p_en = G2pE()
def is_chinese_char(word: str) -> bool:
if len(word) != 1:
return False
return bool(_ZH_WORD_RE.fullmatch(word))
def is_english_word(word: str) -> bool:
if not word:
return False
return bool(_EN_WORD_RE.fullmatch(word))
def g2p_cantonese(sent):
return ToJyutping.get_jyutping_list(sent) # with tone
def g2p_mandarin(sent):
return g2p_zh(sent, tone=True, char_split=False)
def g2p_english(word):
return g2p_en(word)
def g2p_transform(words, lang):
zh_words = []
transformed_words = [0] * len(words)
for idx, w in enumerate(words):
if w == "<SP>":
transformed_words[idx] = w
continue
w = w.replace("?", "").replace(".", "").replace("!", "").replace(",", "")
if is_chinese_char(w):
zh_words.append([idx, w])
else:
if is_english_word(w):
w = EN_FLAG + "-".join(g2p_english(w.lower()))
else:
w = "<SP>"
transformed_words[idx] = w
sent = "".join([k[1] for k in zh_words])
# zh (zh and yue) transformer to g2p
if len(sent) > 0:
if lang == "Cantonese":
g2pm_rst = g2p_cantonese(sent) # with tone
g2pm_rst = [YUE_FLAG + k[1] for k in g2pm_rst]
else:
g2pm_rst = g2p_mandarin(sent)
g2pm_rst = [ZH_FLAG + k for k in g2pm_rst]
for p, w in zip([k[0] for k in zh_words], g2pm_rst):
transformed_words[p] = w
return transformed_words
+279
View File
@@ -0,0 +1,279 @@
# https://modelscope.cn/models/iic/speech_seaco_paraformer_large_asr_nat-zh-cn-16k-common-vocab8404-pytorch/summary
# https://huggingface.co/nvidia/parakeet-tdt-0.6b-v2
import os
import re
import time
from typing import Any, Dict, List, Tuple
import librosa
import numpy as np
from funasr import AutoModel
def _build_words_with_gaps(raw_words, raw_timestamps, wav_fn: str):
words, word_durs = [], []
prev = 0.0
for w, t in zip(raw_words, raw_timestamps):
s, e = float(t[0]), float(t[1])
if s > prev:
words.append("<SP>")
word_durs.append(s - prev)
words.append(w)
word_durs.append(e - s)
prev = e
wav_len = librosa.get_duration(filename=wav_fn)
if wav_len > prev:
if len(words) == 0:
words.append("<SP>")
word_durs.append(wav_len)
return words, word_durs
if words[-1] != "<SP>":
words.append("<SP>")
word_durs.append(wav_len - prev)
else:
word_durs[-1] += wav_len - prev
return words, word_durs
def _word_dur_post_process(words, word_durs, f0):
"""Post-process word durations using f0 to better place silences.
"""
# f0 time grid parameters
sr = 24000 # f0 sample rate
hop_length = 480 # f0 hop length
# Convert word durations (seconds) to frame boundaries on the f0 grid.
boundaries = np.cumsum([
0,
*[
int(dur * sr / hop_length)
for dur in word_durs
],
]).tolist()
sil_tolerance = 5 # tolerance frames for silence detection
ext_tolerance = 5 # tolerance frames for vocal extension
new_words: list[str] = []
new_word_durs: list[float] = []
if words:
new_words.append(words[0])
new_word_durs.append(word_durs[0])
for i in range(1, len(words)):
word = words[i]
if word == "<SP>":
start_frame = boundaries[i]
end_frame = boundaries[i + 1]
num_frames = end_frame - start_frame
frame_idx = start_frame
# Find first region with at least 5 consecutive "unvoiced" frames.
unvoiced_count = 0
while frame_idx < end_frame:
if f0[frame_idx] <= 1: # unvoiced
unvoiced_count += 1
if unvoiced_count >= sil_tolerance:
frame_idx -= sil_tolerance - 1 # back to the last voiced frame
break
else:
unvoiced_count = 0
frame_idx += 1
voice_frames = frame_idx - start_frame
if voice_frames >= int(num_frames * 0.9): # over 90% voiced
# Treat the whole "<SP>" as silence and merge into previous word.
new_word_durs[-1] += word_durs[i]
elif voice_frames >= ext_tolerance: # over 5 frames voiced
# Split the "<SP>" into two parts: leading silence and tail kept as "<SP>".
dur = voice_frames * hop_length / sr
new_word_durs[-1] += dur
new_words.append("<SP>")
new_word_durs.append(word_durs[i] - dur)
else:
# Too short to adjust, keep as-is.
new_words.append(word)
new_word_durs.append(word_durs[i])
else:
new_words.append(word)
new_word_durs.append(word_durs[i])
return new_words, new_word_durs
class _ASRZhModel:
"""Mandarin/Cantonese ASR wrapper."""
def __init__(self, model_path: str, device: str):
self.model = AutoModel(
model=model_path,
disable_update=True,
device=device,
)
def process(self, wav_fn):
out = self.model.generate(wav_fn, output_timestamp=True)[0]
raw_words = out["text"].replace("@", "").split(" ")
raw_timestamps = [[t[0] / 1000, t[1] / 1000] for t in out["timestamp"]]
words, word_durs = _build_words_with_gaps(raw_words, raw_timestamps, wav_fn)
if os.path.exists(wav_fn.replace(".wav", "_f0.npy")):
words, word_durs = _word_dur_post_process(
words, word_durs, np.load(wav_fn.replace(".wav", "_f0.npy"))
)
return words, word_durs
class _ASREnModel:
"""English ASR wrapper for NeMo Parakeet-TDT."""
def __init__(self, model_path: str, device: str):
try:
import nemo.collections.asr as nemo_asr # type: ignore
except Exception as e: # pragma: no cover
raise ImportError(
"NeMo (nemo_toolkit) is required for ASR English but is not available in this Python env. "
"Install it in the active environment, then retry."
) from e
self.model = nemo_asr.models.ASRModel.restore_from(
restore_path=model_path,
map_location=device,
)
self.model.eval()
@staticmethod
def _clean_word(word: str) -> str:
return re.sub(r"[\?\.,:]", "", word).strip()
@staticmethod
def _extract_word_segments(output: Any) -> List[Dict[str, Any]]:
ts = getattr(output, "timestamp", None)
if not ts or not isinstance(ts, dict):
return []
word_ts = ts.get("word")
return word_ts if isinstance(word_ts, list) else []
def process(self, wav_fn: str) -> Tuple[List[str], List[float]]:
outputs = self.model.transcribe(
[wav_fn],
timestamps=True,
batch_size=1,
num_workers=0,
)
output = outputs[0] if outputs else None
raw_words: List[str] = []
raw_timestamps: List[List[float]] = []
if output is not None:
for w in self._extract_word_segments(output):
s, e = float(w.get("start", 0.0)), float(w.get("end", 0.0))
word = self._clean_word(str(w.get("word", "")))
if word:
raw_words.append(word)
raw_timestamps.append([s, e])
words, durs = _build_words_with_gaps(raw_words, raw_timestamps, wav_fn)
if os.path.exists(wav_fn.replace(".wav", "_f0.npy")):
words, durs = _word_dur_post_process(
words, durs, np.load(wav_fn.replace(".wav", "_f0.npy"))
)
return words, durs
class LyricTranscriber:
"""Transcribe lyrics from singing voice segment
"""
def __init__(
self,
zh_model_path: str,
en_model_path: str,
device: str = "cuda",
*,
verbose: bool = True,
):
"""Initialize lyric transcriber.
Args:
zh_model_path (str): Path to the Chinese model file.
en_model_path (str): Path to the English model file.
device (str): Device to use for tensor operations.
verbose (bool): Whether to print verbose logs.
"""
self.verbose = verbose
self.device = device
self.zh_model_path = zh_model_path
self.en_model_path = en_model_path
if self.verbose:
print(
"[lyric transcription] init: start:",
f"device={device}",
f"model_path={zh_model_path}",
)
# Always initialize Chinese ASR.
self.zh_model = _ASRZhModel(device=device, model_path=zh_model_path)
# English ASR will be lazily initialized on first English request to avoid long waiting cost when importing NeMo
self.en_model = None
if self.verbose:
print("[lyric transcription] init: success")
def process(self, wav_fn, language: str | None = "Mandarin", *, verbose: bool | None = None):
""" Lyric transcriber process
Args:
wav_fn (str): Path to the audio file.
language (str | None): Language of the audio. Defaults to "Mandarin". Supports "Mandarin", "Cantonese" and "English".
verbose (bool | None): Whether to print verbose logs. Defaults to None.
"""
v = self.verbose if verbose is None else verbose
if language not in {"Mandarin", "Cantonese", "English"}:
raise ValueError(f"Unsupported language: {language}, should be one of ['Mandarin', 'Cantonese', 'English']")
if v:
print(f"[lyric transcription] process: start: wav_fn={wav_fn} language={language}")
t0 = time.time()
lang = (language or "auto").lower()
if lang in {"english"}:
if self.en_model is None:
# Lazy-load NeMo model only when English is actually used.
if v:
print("[lyric transcription] init English ASR, please make sure NeMo is installed")
self.en_model = _ASREnModel(model_path=self.en_model_path, device=self.device)
out = self.en_model.process(wav_fn)
else:
out = self.zh_model.process(wav_fn)
if v:
words, durs = out
n_words = len(words) if isinstance(words, list) else 0
dur_sum = float(sum(durs)) if isinstance(durs, list) else 0.0
dt = time.time() - t0
print(
"[lyric transcription] process: done:",
f"n_words={n_words}",
f"dur_sum={dur_sum:.3f}s",
f"time={dt:.3f}s",
)
return out
if __name__ == "__main__":
m = LyricTranscriber(
zh_model_path="pretrained_models/speech_seaco_paraformer_large_asr_nat-zh-cn-16k-common-vocab8404-pytorch",
en_model_path="pretrained_models/parakeet-tdt-0.6b-v2/parakeet-tdt-0.6b-v2.nemo",
device="cuda"
)
print(m.process("example/test/asr_zh.wav", language="Mandarin"))
print(m.process("example/test/asr_en.wav", language="English"))
+173
View File
@@ -0,0 +1,173 @@
# 🎹 MIDI Editor - Web-based Singing MIDI Editor
[English](README.md) | [简体中文](README_CN.md)
A full-featured web MIDI editor for singing voice production, similar to ACE-Studio and VOCALOID. It supports real-time drag editing of MIDI notes, lyric editing, audio waveform alignment, and importing/exporting MIDI files with lyrics.
![MIDI Editor](https://img.shields.io/badge/React-19.2-blue) ![TypeScript](https://img.shields.io/badge/TypeScript-5.9-blue) ![Vite](https://img.shields.io/badge/Vite-7.2-purple)
## ✨ Features
### 🎼 Piano Roll Editing
- **Visual note editing**: Full range from C1 to C8 with intuitive piano keys
- **Drag operations**:
- Move notes: drag note blocks to adjust position and pitch
- Resize start: drag the left edge to adjust start time
- Resize end: drag the right edge to adjust end time
- **Double-click to add**: Add new notes quickly in empty areas
- **Piano key preview**: Click a key to audition the pitch
### 🔍 Zoom & Navigation
- **Horizontal zoom**
- **Vertical zoom**
- **Dynamic snapping**: finer snap granularity at higher zoom (min 0.01s)
- **Auto scroll**: keep the playhead visible during playback
### 📝 Lyric Editing
- **Inline editing**: edit lyrics for each note in the side list
- **Batch fill**: enter a string and auto-fill notes in order
- **Fill from selection**: start batch fill from the selected note
- **Precise fields**: edit PITCH, START, and END directly
- **Confirm edits**: press Enter or click ✓ to confirm
### 🎵 Audio Alignment
- **Waveform display**: sync waveform with the MIDI timeline
- **Formats**: MP3, WAV, OGG, FLAC, M4A, AAC
- **Sync playback**: play audio and MIDI together with independent mute
- **Click to seek**: click waveform or timeline to seek
### ⚠️ Overlap Detection
- **Visual highlight**: overlapping notes blink in red
- **Smart tolerance**: adjacent notes (end equals next start) are not overlaps
- **One-click fix**: remove all overlaps automatically
- **Export warning**: warn if overlaps exist during export
### 📥 Import & Export
- **MIDI import**: parse standard MIDI and lyric metadata
- **MIDI export**: export MIDI with lyrics
- **Chinese support**: full UTF-8 lyrics support
### 🎨 UI & UX
- **Theme toggle**: light and dark modes
- **Responsive layout**: adapts to window size
- **SVG grid**: cross-browser grid rendering
- **Status feedback**: real-time state and error tips
## 🚀 Quick Start
### Requirements
- Node.js 18+
- npm or yarn
### Install
```bash
# Clone
git clone <repository-url>
cd MIDI_Editor
# Install dependencies
npm install
# Start dev server
npm run dev
# Expose to LAN
npm run dev -- --host 0.0.0.0
```
### Build
```bash
# Build for production
npm run build
# Preview build
npm run preview
```
## 📖 Usage
### Basic Workflow
1. **Import MIDI**: click Import MIDI and select a .mid file
2. **Edit notes**: drag notes in the piano roll to adjust time and pitch
3. **Add lyrics**: edit lyrics in the right-side list
4. **Align audio** (optional): import reference audio
5. **Export**: click Export MIDI with lyrics
### Shortcuts
| Action | Description |
|------|------|
| Double-click piano roll | Add a new note |
| Double-click note | Edit lyric |
| Drag note | Move note and pitch |
| Drag note edges | Resize note |
| Backspace / Delete | Delete selected note |
| Enter | Confirm value edits |
| Escape | Cancel value edits |
| Ctrl(Command) + Wheel | Horizontal zoom |
| Ctrl(Command) + Shift(Option) + Wheel | Vertical zoom |
### Playback Controls
| Button | Description |
|------|------|
| ⏮ | Go to start |
| ⏪ 2s | Back 2 seconds |
| ▶ / ⏸ | Play / Pause |
| 2s ⏩ | Forward 2 seconds |
| ⏭ | Go to end |
## 🛠 Tech Stack
- **Frontend**: React 19 + TypeScript
- **Build**: Vite 7
- **State**: Zustand
- **Audio**: Tone.js
- **Waveform**: WaveSurfer.js
- **MIDI**: @tonejs/midi
- **Styles**: CSS with custom variables
## 📁 Project Structure
```
.
├── eslint.config.js
├── index.html
├── package.json
├── postcss.config.js
├── README.md
├── README_CN.md
├── tailwind.config.js
├── tsconfig.app.json
├── tsconfig.json
├── tsconfig.node.json
├── vite.config.ts
├── public/
└── src/
├── App.css
├── App.tsx
├── constants.ts
├── index.css
├── main.tsx
├── types.ts
├── assets/
├── components/
│ ├── AudioTrack.tsx
│ ├── LyricTable.tsx
│ └── PianoRoll.tsx
├── lib/
│ └── midi.ts
└── store/
└── useMidiStore.ts
```
+173
View File
@@ -0,0 +1,173 @@
# 🎹 MIDI Editor - 网页端歌声 MIDI 编辑器
[English](README.md) | [简体中文](README_CN.md)
一个功能完整的网页端歌声 MIDI 文件编辑器,类似 ACE-Studio 和 VOCALOID。支持实时拖拽调整 MIDI 音符、歌词编辑、音频波形对齐,以及导入导出含歌词的 MIDI 文件。
![MIDI Editor](https://img.shields.io/badge/React-19.2-blue) ![TypeScript](https://img.shields.io/badge/TypeScript-5.9-blue) ![Vite](https://img.shields.io/badge/Vite-7.2-purple)
## ✨ 功能特性
### 🎼 钢琴卷帘编辑
- **可视化音符编辑**:支持 C1-C8 全音域显示,直观的钢琴键布局
- **拖拽操作**:
- 移动音符:拖拽音符块调整位置和音高
- 调整音头:拖拽音符左边缘调整开始时间
- 调整音尾:拖拽音符右边缘调整结束时间
- **双击添加**:在钢琴卷帘空白处双击快速添加新音符
- **钢琴键试听**:点击左侧钢琴键可试听对应音高
### 🔍 缩放与导航
- **水平缩放**
- **垂直缩放**
- **动态精度**:缩放越大,音符调整的 snap 粒度越精细(最小 0.01 秒)
- **自动滚动**:播放时播放头自动保持可见
### 📝 歌词编辑
- **实时编辑**:右侧列表直接编辑每个音符的歌词
- **批量填充**:输入一段歌词,按字顺序自动填充到音符
- **从选中开始**:批量填充可从当前选中的音符开始
- **精确调整**:可直接编辑 PITCH(音高)、START(开始时间)、END(结束时间)
- **确认机制**:修改数值后按 Enter 或点击 ✓ 确认,避免误操作
### 🎵 音频对齐
- **波形显示**:导入音频后显示波形,与 MIDI 同步滚动
- **格式支持**:MP3、WAV、OGG、FLAC、M4A、AAC
- **同步播放**:音频与 MIDI 同步播放,可分别静音
- **点击定位**:点击波形或时间尺可快速定位播放位置
### ⚠️ 重叠检测
- **可视化标注**:时间重叠的音符显示为红色并闪烁
- **智能容差**:紧邻的音符(上一个结束 = 下一个开始)不视为重叠
- **一键修复**:点击消除重叠按钮自动修复所有重叠
- **导出提醒**:导出时如有重叠会弹出警告
### 📥 导入导出
- **MIDI 导入**:支持标准 MIDI 文件,自动解析歌词元数据
- **MIDI 导出**:导出包含歌词信息的 MIDI 文件
- **中文支持**:完整支持中文歌词的导入导出(UTF-8 编码)
### 🎨 界面特性
- **主题切换**:支持浅色/深色主题
- **响应式布局**:自适应窗口大小
- **SVG 网格**:跨浏览器兼容的网格渲染
- **状态提示**:实时显示操作状态和错误信息
## 🚀 快速开始
### 环境要求
- Node.js 18+
- npm 或 yarn
### 安装
```bash
# 克隆项目
git clone <repository-url>
cd MIDI_Editor
# 安装依赖
npm install
# 启动开发服务器
npm run dev
# 在局域网启动
npm run dev -- --host 0.0.0.0
```
### 构建
```bash
# 构建生产版本
npm run build
# 预览构建结果
npm run preview
```
## 📖 使用指南
### 基本工作流
1. **导入 MIDI**:点击导入 MIDI 按钮选择 .mid 文件
2. **编辑音符**:在钢琴卷帘中拖拽调整音符位置和时长
3. **添加歌词**:在右侧列表中输入每个音符的歌词
4. **对齐音频**(可选):导入参考音频进行对照编辑
5. **导出文件**:点击导出含歌词 MIDI 保存文件
### 快捷操作
| 操作 | 说明 |
|------|------|
| 双击钢琴卷帘 | 添加新音符 |
| 双击音符 | 修改歌词 |
| 拖拽音符 | 移动音符位置/音高 |
| 拖拽音符边缘 | 调整音符时长 |
| Backspace / Delete | 删除选中音符 |
| Enter | 确认数值修改 |
| Escape | 取消数值修改 |
| Ctrl(Command) + 滚轮 | 水平缩放 |
| Ctrl(Command) + Shift(Option) + 滚轮 | 垂直缩放 |
### 播放控制
| 按钮 | 功能 |
|------|------|
| ⏮ | 回到开头 |
| ⏪ 2s | 后退 2 秒 |
| ▶ / ⏸ | 播放 / 暂停 |
| 2s ⏩ | 前进 2 秒 |
| ⏭ | 跳到结尾 |
## 🛠 技术栈
- **前端框架**:React 19 + TypeScript
- **构建工具**:Vite 7
- **状态管理**:Zustand
- **音频引擎**:Tone.js
- **波形显示**:WaveSurfer.js
- **MIDI 解析**:@tonejs/midi
- **样式**:CSS(自定义变量主题)
## 📁 项目结构
```
.
├── eslint.config.js
├── index.html
├── package.json
├── postcss.config.js
├── README.md
├── README_CN.md
├── tailwind.config.js
├── tsconfig.app.json
├── tsconfig.json
├── tsconfig.node.json
├── vite.config.ts
├── public/
└── src/
├── App.css
├── App.tsx
├── constants.ts
├── index.css
├── main.tsx
├── types.ts
├── assets/
├── components/
│ ├── AudioTrack.tsx
│ ├── LyricTable.tsx
│ └── PianoRoll.tsx
├── lib/
│ └── midi.ts
└── store/
└── useMidiStore.ts
```
@@ -0,0 +1,23 @@
import js from '@eslint/js'
import globals from 'globals'
import reactHooks from 'eslint-plugin-react-hooks'
import reactRefresh from 'eslint-plugin-react-refresh'
import tseslint from 'typescript-eslint'
import { defineConfig, globalIgnores } from 'eslint/config'
export default defineConfig([
globalIgnores(['dist']),
{
files: ['**/*.{ts,tsx}'],
extends: [
js.configs.recommended,
tseslint.configs.recommended,
reactHooks.configs.flat.recommended,
reactRefresh.configs.vite,
],
languageOptions: {
ecmaVersion: 2020,
globals: globals.browser,
},
},
])
+13
View File
@@ -0,0 +1,13 @@
<!doctype html>
<html lang="en">
<head>
<meta charset="UTF-8" />
<link rel="icon" type="image/svg+xml" href="/vite.svg" />
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
<title>midi-editor</title>
</head>
<body>
<div id="root"></div>
<script type="module" src="/src/main.tsx"></script>
</body>
</html>
File diff suppressed because it is too large Load Diff
+39
View File
@@ -0,0 +1,39 @@
{
"name": "midi-editor",
"private": true,
"version": "0.0.0",
"type": "module",
"scripts": {
"dev": "vite",
"build": "tsc -b && vite build",
"lint": "eslint .",
"preview": "vite preview"
},
"dependencies": {
"@tonejs/midi": "^2.0.28",
"class-variance-authority": "^0.7.1",
"nanoid": "^5.1.6",
"react": "^19.2.0",
"react-dom": "^19.2.0",
"tone": "^15.1.22",
"wavesurfer.js": "^7.12.1",
"zustand": "^5.0.10"
},
"devDependencies": {
"@eslint/js": "^9.39.1",
"@types/node": "^24.10.1",
"@types/react": "^19.2.5",
"@types/react-dom": "^19.2.3",
"@vitejs/plugin-react": "^5.1.1",
"autoprefixer": "^10.4.20",
"eslint": "^9.39.1",
"eslint-plugin-react-hooks": "^7.0.1",
"eslint-plugin-react-refresh": "^0.4.24",
"globals": "^16.5.0",
"postcss": "^8.4.47",
"tailwindcss": "^3.4.15",
"typescript": "~5.9.3",
"typescript-eslint": "^8.46.4",
"vite": "^7.2.4"
}
}
@@ -0,0 +1,6 @@
export default {
plugins: {
tailwindcss: {},
autoprefixer: {},
},
}
@@ -0,0 +1 @@
<svg xmlns="http://www.w3.org/2000/svg" xmlns:xlink="http://www.w3.org/1999/xlink" aria-hidden="true" role="img" class="iconify iconify--logos" width="31.88" height="32" preserveAspectRatio="xMidYMid meet" viewBox="0 0 256 257"><defs><linearGradient id="IconifyId1813088fe1fbc01fb466" x1="-.828%" x2="57.636%" y1="7.652%" y2="78.411%"><stop offset="0%" stop-color="#41D1FF"></stop><stop offset="100%" stop-color="#BD34FE"></stop></linearGradient><linearGradient id="IconifyId1813088fe1fbc01fb467" x1="43.376%" x2="50.316%" y1="2.242%" y2="89.03%"><stop offset="0%" stop-color="#FFEA83"></stop><stop offset="8.333%" stop-color="#FFDD35"></stop><stop offset="100%" stop-color="#FFA800"></stop></linearGradient></defs><path fill="url(#IconifyId1813088fe1fbc01fb466)" d="M255.153 37.938L134.897 252.976c-2.483 4.44-8.862 4.466-11.382.048L.875 37.958c-2.746-4.814 1.371-10.646 6.827-9.67l120.385 21.517a6.537 6.537 0 0 0 2.322-.004l117.867-21.483c5.438-.991 9.574 4.796 6.877 9.62Z"></path><path fill="url(#IconifyId1813088fe1fbc01fb467)" d="M185.432.063L96.44 17.501a3.268 3.268 0 0 0-2.634 3.014l-5.474 92.456a3.268 3.268 0 0 0 3.997 3.378l24.777-5.718c2.318-.535 4.413 1.507 3.936 3.838l-7.361 36.047c-.495 2.426 1.782 4.5 4.151 3.78l15.304-4.649c2.372-.72 4.652 1.36 4.15 3.788l-11.698 56.621c-.732 3.542 3.979 5.473 5.943 2.437l1.313-2.028l72.516-144.72c1.215-2.423-.88-5.186-3.54-4.672l-25.505 4.922c-2.396.462-4.435-1.77-3.759-4.114l16.646-57.705c.677-2.35-1.37-4.583-3.769-4.113Z"></path></svg>

After

Width:  |  Height:  |  Size: 1.5 KiB

+785
View File
@@ -0,0 +1,785 @@
.app-shell {
padding: 24px;
color: var(--text-primary);
width: 100%;
max-width: 100%;
margin: 0;
height: 100vh;
max-height: 100vh;
display: flex;
flex-direction: column;
overflow: hidden;
box-sizing: border-box;
}
.topbar {
display: flex;
align-items: center;
justify-content: space-between;
gap: 24px;
background: var(--panel-strong);
border: 1px solid var(--border-subtle);
border-radius: 16px;
padding: 20px 24px;
box-shadow: var(--shadow-panel);
}
.topbar h1 {
margin: 4px 0 0 0;
font-size: 26px;
letter-spacing: -0.5px;
}
.eyebrow {
margin: 0;
text-transform: uppercase;
font-size: 12px;
letter-spacing: 2px;
color: var(--text-muted);
}
.muted {
margin: 6px 0 0 0;
color: var(--text-muted);
}
.actions {
display: flex;
gap: 10px;
}
.icon-toggle {
width: 40px;
height: 40px;
border-radius: 999px;
border: 1px solid var(--border-soft);
background: var(--button-ghost-bg);
color: var(--button-ghost-text);
display: inline-flex;
align-items: center;
justify-content: center;
font-size: 18px;
cursor: pointer;
}
.icon-toggle:hover {
transform: translateY(-1px);
}
.audio-bar {
margin-top: 14px;
padding: 12px 16px;
border-radius: 14px;
background: var(--panel-strong);
border: 1px solid var(--border-subtle);
display: flex;
align-items: center;
justify-content: space-between;
gap: 16px;
}
.audio-left {
display: flex;
align-items: center;
gap: 12px;
}
.audio-hint {
color: var(--text-muted);
font-size: 12px;
}
.audio-right {
display: flex;
align-items: center;
gap: 20px;
}
.volume-control {
display: flex;
align-items: center;
gap: 8px;
}
.volume-label {
font-size: 12px;
color: var(--text-muted);
min-width: 32px;
}
.volume-slider {
width: 80px;
height: 4px;
cursor: pointer;
accent-color: var(--accent);
}
.volume-value {
font-size: 11px;
color: var(--text-muted);
min-width: 36px;
text-align: right;
}
.toggle {
display: inline-flex;
align-items: center;
gap: 8px;
font-size: 13px;
color: var(--text-primary);
}
.panel {
margin-top: 18px;
background: var(--panel);
border: 1px solid var(--border-subtle);
border-radius: 16px;
padding: 18px;
box-shadow: var(--shadow-panel);
display: flex;
flex-direction: column;
flex: 1;
min-height: 0;
overflow: hidden;
}
.panel-split {
display: grid;
grid-template-columns: minmax(0, 1fr) 360px;
gap: 16px;
align-items: stretch;
flex: 1;
min-height: 0;
max-height: 100%;
overflow: hidden;
}
.panel-main {
min-width: 0;
display: flex;
flex-direction: column;
min-height: 0;
max-height: 100%;
overflow: hidden;
}
.panel-side {
display: flex;
flex-direction: column;
gap: 16px;
width: 360px;
max-width: 360px;
/* Use absolute positioning to enforce height */
position: relative;
overflow: hidden;
}
.controls {
display: grid;
grid-template-columns: repeat(auto-fit, minmax(180px, 1fr));
gap: 14px;
align-items: center;
background: var(--panel-soft);
padding: 12px 14px;
border-radius: 12px;
border: 1px solid var(--border-soft);
flex-shrink: 0;
}
.controls label {
display: block;
font-size: 12px;
text-transform: uppercase;
letter-spacing: 1px;
color: var(--text-muted);
margin-bottom: 4px;
}
.controls input[type='number'] {
width: 100%;
padding: 10px 12px;
border-radius: 10px;
border: 1px solid var(--border-soft);
background: var(--input-bg);
color: var(--text-primary);
}
.timesig {
display: flex;
align-items: center;
gap: 6px;
}
.timesig span {
font-weight: 700;
color: var(--text-muted);
}
.transport {
display: flex;
gap: 8px;
align-items: center;
}
.selection-controls {
display: flex;
gap: 8px;
align-items: center;
padding: 6px 0;
}
.selection-btn {
font-size: 12px !important;
padding: 6px 10px !important;
}
.selection-btn.active {
background: var(--accent) !important;
color: white !important;
}
.selection-info {
font-size: 12px;
color: var(--accent);
font-weight: 500;
padding: 4px 8px;
background: rgba(var(--accent-rgb), 0.1);
border-radius: 6px;
}
.status {
color: var(--text-muted);
font-size: 13px;
}
.button,
.actions button,
.transport button,
.ghost,
.primary,
.soft {
cursor: pointer;
border-radius: 12px;
border: 1px solid transparent;
padding: 10px 14px;
font-weight: 600;
transition: transform 140ms ease, box-shadow 140ms ease, background 140ms ease, border 140ms ease;
color: #0f1528;
}
.ghost {
background: var(--button-ghost-bg);
color: var(--button-ghost-text);
border-color: var(--border-soft);
}
.primary {
background: linear-gradient(135deg, var(--accent), var(--accent-strong));
color: var(--button-primary-text);
box-shadow: 0 8px 26px rgba(72, 228, 194, 0.2);
}
.soft {
background: var(--button-soft-bg);
color: var(--button-soft-text);
border: 1px solid var(--border-soft);
}
.ghost:disabled,
.primary:disabled,
.soft:disabled {
opacity: 0.6;
cursor: not-allowed;
}
.ghost:hover,
.primary:hover,
.soft:hover {
transform: translateY(-1px);
}
.piano-shell {
border-radius: 12px;
background: var(--panel-strong);
border: 1px solid var(--border-subtle);
overflow: hidden;
flex: 1;
min-height: 0;
max-height: 100%;
display: flex;
flex-direction: column;
}
.ruler {
position: relative;
height: 32px;
background: var(--panel-soft);
border-bottom: 1px solid var(--border-soft);
min-width: 100%;
}
.ruler-shell {
display: flex;
}
.ruler-spacer {
background: var(--panel-soft);
border-bottom: 1px solid var(--border-soft);
height: 32px;
}
.ruler-scroll {
overflow: hidden;
flex: 1;
height: 32px;
cursor: pointer;
}
.measure-mark {
position: absolute;
top: 0;
height: 100%;
display: flex;
flex-direction: column;
align-items: flex-start;
font-size: 10px;
color: var(--text-muted);
padding-left: 4px;
border-left: 1px solid var(--border-soft);
}
.measure-mark span {
margin-top: 2px;
}
.ruler-playhead {
position: absolute;
top: 0;
width: 2px;
height: 100%;
background: #ff7043;
pointer-events: none;
z-index: 10;
}
.ruler-scroll.selecting {
cursor: crosshair;
}
.selection-range {
position: absolute;
top: 0;
height: 100%;
background: rgba(66, 165, 245, 0.35);
border-left: 2px solid #42a5f5;
border-right: 2px solid #42a5f5;
pointer-events: none;
z-index: 5;
}
.grid-selection-range {
position: absolute;
top: 0;
background: rgba(66, 165, 245, 0.15);
border-left: 2px dashed #42a5f5;
border-right: 2px dashed #42a5f5;
pointer-events: none;
z-index: 1;
}
.roll-body {
display: flex;
flex: 1;
min-height: 0;
overflow: hidden;
}
.pitch-rail {
background: var(--panel-strong);
border-right: 1px solid var(--border-subtle);
color: var(--text-primary);
font-size: 12px;
text-align: right;
overflow: hidden;
flex-shrink: 0;
height: 100%;
}
.pitch-cell {
border-bottom: 1px solid var(--border-soft);
display: flex;
align-items: center;
justify-content: flex-end;
padding: 0 4px;
font-variant-numeric: tabular-nums;
box-sizing: border-box;
}
.pitch-white {
background: rgba(255, 255, 255, 0.06);
color: var(--text-primary);
}
.pitch-black {
background: rgba(0, 0, 0, 0.35);
color: rgba(233, 238, 247, 0.9);
}
.pitch-c {
background: rgba(100, 150, 255, 0.15);
font-weight: 600;
}
.pitch-label {
font-size: 10px;
}
.roll-grid {
position: relative;
overflow: auto;
flex: 1;
min-height: 0;
background-color: var(--grid-bg);
}
.grid-content {
background-color: var(--grid-bg);
}
.grid-svg {
shape-rendering: crispEdges;
}
.grid-overlay {
position: relative;
}
.note-chip {
position: absolute;
background: linear-gradient(135deg, var(--accent), var(--accent-strong));
border-radius: 6px;
border: 1px solid rgba(255, 255, 255, 0.16);
box-shadow: 0 10px 22px rgba(0, 0, 0, 0.25);
display: flex;
align-items: center;
justify-content: center;
color: var(--note-text);
font-weight: 700;
user-select: none;
box-sizing: border-box;
}
.note-active {
outline: 2px solid #ff7043;
z-index: 2;
}
.note-overlap {
background: linear-gradient(135deg, #ef5350 0%, #ff7043 100%) !important;
animation: pulse-overlap 1s ease-in-out infinite;
}
/* Selected overlapping note - more visible outline */
.note-overlap.note-active {
outline: 3px solid #1e40af;
outline-offset: 1px;
box-shadow: 0 0 12px rgba(30, 64, 175, 0.8);
animation: none;
}
@keyframes pulse-overlap {
0%, 100% { opacity: 1; }
50% { opacity: 0.7; }
}
.playhead {
position: absolute;
top: 0;
width: 2px;
background: #ff7043;
box-shadow: 0 0 12px rgba(255, 112, 67, 0.6);
pointer-events: none;
z-index: 20;
}
.pitch-rail-inner {
will-change: transform;
}
.note-label {
width: 100%;
text-align: center;
font-size: 12px;
padding: 0 12px;
overflow: hidden;
text-overflow: ellipsis;
white-space: nowrap;
}
.note-handle {
position: absolute;
top: 0;
width: 8px;
height: 100%;
background: rgba(255, 255, 255, 0.25);
cursor: ew-resize;
}
.note-handle.start {
left: 0;
border-radius: 6px 0 0 6px;
}
.note-handle.end {
right: 0;
border-radius: 0 6px 6px 0;
}
.lyric-container {
flex: 1;
min-height: 0;
position: relative;
}
.lyric-card {
border: 1px solid rgba(255, 255, 255, 0.06);
border-radius: 12px;
background: var(--panel-soft);
overflow: hidden;
display: flex;
flex-direction: column;
/* Force fixed height with absolute positioning */
position: absolute;
top: 0;
left: 0;
right: 0;
bottom: 0;
}
.lyric-bulk {
display: flex;
gap: 8px;
padding: 10px 12px;
border-bottom: 1px solid rgba(255, 255, 255, 0.06);
align-items: center;
}
.lyric-bulk-input {
flex: 1;
padding: 8px 10px;
border-radius: 10px;
border: 1px solid var(--border-soft);
background: var(--input-bg);
color: var(--text-primary);
resize: vertical;
}
.lyric-header,
.lyric-row {
display: grid;
grid-template-columns: 1.4fr 0.5fr 0.5fr 0.5fr;
gap: 8px;
padding: 10px 12px;
align-items: center;
}
.lyric-header {
font-size: 12px;
text-transform: uppercase;
letter-spacing: 1px;
color: var(--text-muted);
border-bottom: 1px solid rgba(255, 255, 255, 0.06);
}
.lyric-list {
overflow-y: auto;
overflow-x: hidden;
flex: 1;
min-height: 0;
}
.lyric-row {
border-bottom: 1px solid rgba(255, 255, 255, 0.04);
}
.lyric-row:hover {
background: rgba(255, 255, 255, 0.03);
}
.lyric-row-active {
background: rgba(72, 228, 194, 0.08);
border-left: 3px solid #48e4c2;
}
.lyric-input {
width: 100%;
padding: 8px 10px;
border-radius: 10px;
border: 1px solid var(--border-soft);
background: var(--input-bg);
color: var(--text-primary);
}
.lyric-meta {
color: var(--text-muted);
font-variant-numeric: tabular-nums;
}
.editable-cell {
position: relative;
display: flex;
align-items: center;
gap: 2px;
}
.lyric-meta-input {
width: 100%;
padding: 2px 4px;
border: 1px solid transparent;
border-radius: 4px;
background: transparent;
color: var(--text-muted);
font-size: 12px;
font-variant-numeric: tabular-nums;
text-align: center;
outline: none;
transition: border-color 0.15s, background-color 0.15s;
}
.lyric-meta-input:hover {
background: var(--surface-elevated);
}
.lyric-meta-input:focus {
border-color: var(--accent);
background: var(--surface-elevated);
color: var(--text-primary);
}
.lyric-meta-dirty {
border-color: #f59e0b !important;
background: rgba(245, 158, 11, 0.1) !important;
}
.confirm-btn {
flex-shrink: 0;
width: 18px;
height: 18px;
padding: 0;
border: none;
border-radius: 4px;
background: #22c55e;
color: white;
font-size: 12px;
font-weight: bold;
cursor: pointer;
display: flex;
align-items: center;
justify-content: center;
transition: background 0.15s;
}
.confirm-btn:hover {
background: #16a34a;
}
/* Hide number input spinners */
.lyric-meta-input::-webkit-outer-spin-button,
.lyric-meta-input::-webkit-inner-spin-button {
-webkit-appearance: none;
margin: 0;
}
.lyric-meta-input[type=number] {
-moz-appearance: textfield;
}
.lyric-empty {
padding: 16px;
color: var(--text-muted);
text-align: center;
}
.audio-track {
display: grid;
grid-template-columns: 80px 1fr;
gap: 12px;
align-items: center;
padding: 12px 14px;
border-radius: 12px;
border: 1px solid var(--border-soft);
background: var(--panel-soft);
margin-bottom: 12px;
flex-shrink: 0;
}
.audio-track-label {
font-size: 12px;
text-transform: uppercase;
letter-spacing: 1px;
color: var(--text-muted);
}
.audio-wave {
width: 100%;
height: 80px;
min-height: 80px;
}
:root {
--text-primary: #e9eef7;
--text-muted: rgba(233, 238, 247, 0.7);
--panel: rgba(13, 16, 28, 0.8);
--panel-strong: rgba(16, 21, 35, 0.95);
--panel-soft: rgba(255, 255, 255, 0.03);
--border-subtle: rgba(255, 255, 255, 0.08);
--border-soft: rgba(255, 255, 255, 0.12);
--input-bg: rgba(255, 255, 255, 0.06);
--grid-bg: rgba(14, 18, 30, 0.9);
--grid-line-minor: rgba(233, 238, 247, 0.08);
--grid-line-major: rgba(233, 238, 247, 0.16);
--accent: #48e4c2;
--accent-strong: #4b64bc;
--note-text: #0b1122;
--button-ghost-bg: rgba(233, 238, 247, 0.18);
--button-ghost-text: #ffffff;
--button-soft-bg: rgba(255, 255, 255, 0.14);
--button-soft-text: #ffffff;
--button-primary-text: #0b1122;
--shadow-panel: 0 18px 40px rgba(0, 0, 0, 0.32);
}
:root[data-theme='light'] {
--text-primary: #1b2238;
--text-muted: rgba(27, 34, 56, 0.7);
--panel: rgba(255, 255, 255, 0.9);
--panel-strong: rgba(250, 252, 255, 0.98);
--panel-soft: rgba(15, 23, 42, 0.04);
--border-subtle: rgba(15, 23, 42, 0.12);
--border-soft: rgba(15, 23, 42, 0.16);
--input-bg: rgba(15, 23, 42, 0.06);
--grid-bg: rgba(248, 250, 255, 0.95);
--grid-line-minor: rgba(15, 23, 42, 0.12);
--grid-line-major: rgba(15, 23, 42, 0.24);
--accent: #3f8cff;
--accent-strong: #4b64bc;
--note-text: #ffffff;
--button-ghost-bg: rgba(15, 23, 42, 0.06);
--button-ghost-text: #1b2238;
--button-soft-bg: rgba(15, 23, 42, 0.06);
--button-soft-text: #1b2238;
--button-primary-text: #0b1122;
--shadow-panel: 0 18px 40px rgba(15, 23, 42, 0.15);
}
.sr-only {
position: absolute;
width: 1px;
height: 1px;
padding: 0;
margin: -1px;
overflow: hidden;
clip: rect(0, 0, 0, 0);
white-space: nowrap;
border: 0;
}
+654
View File
@@ -0,0 +1,654 @@
import { useCallback, useEffect, useMemo, useRef, useState } from 'react'
import * as Tone from 'tone'
import { PianoRoll } from './components/PianoRoll'
import { LyricTable } from './components/LyricTable'
import { AudioTrack } from './components/AudioTrack'
import { useMidiStore } from './store/useMidiStore'
import { exportMidi, importMidiFile } from './lib/midi'
import type { TimeSignature } from './types'
import { BASE_GRID_SECOND_WIDTH, BASE_ROW_HEIGHT, LOW_NOTE, HIGH_NOTE } from './constants'
import './App.css'
type PlayEvent = {
time: number
midi: number
duration: number
velocity: number
}
function App() {
const {
notes,
tempo,
timeSignature,
selectedId,
playhead,
ppq,
addNote,
updateNote,
removeNote,
setNotes,
setTempo,
setTimeSignature,
setPpq,
select,
setPlayhead,
} = useMidiStore()
const [status, setStatus] = useState('准备就绪')
const [isPlaying, setIsPlaying] = useState(false)
const [theme, setTheme] = useState<'dark' | 'light'>('light')
const [audioUrl, setAudioUrl] = useState<string | null>(null)
const [audioDuration, setAudioDuration] = useState(0)
const [midiVolume, setMidiVolume] = useState(80) // 0-100
const [audioVolume, setAudioVolume] = useState(80) // 0-100
const [horizontalZoom, setHorizontalZoom] = useState(1)
const [verticalZoom, setVerticalZoom] = useState(1)
const [focusLyricId, setFocusLyricId] = useState<string | null>(null)
// Selection range for loop playback (in seconds)
const [selectionStart, setSelectionStart] = useState<number | null>(null)
const [selectionEnd, setSelectionEnd] = useState<number | null>(null)
const [isSelectingRange, setIsSelectingRange] = useState(false)
const fileInputRef = useRef<HTMLInputElement | null>(null)
const audioInputRef = useRef<HTMLInputElement | null>(null)
const audioRef = useRef<HTMLAudioElement | null>(null)
const partRef = useRef<Tone.Part<PlayEvent> | null>(null)
const synthRef = useRef<Tone.PolySynth | null>(null)
const rafRef = useRef<number | null>(null)
const audioScrollRef = useRef<HTMLDivElement | null>(null)
useEffect(() => {
return () => {
stopPlayback()
synthRef.current?.dispose()
}
}, [])
useEffect(() => {
document.documentElement.dataset.theme = theme
}, [theme])
// Sync audio volume - also trigger when audioUrl changes (new audio loaded)
useEffect(() => {
if (audioRef.current) {
audioRef.current.volume = audioVolume / 100
}
}, [audioVolume, audioUrl])
// Sync MIDI synth volume
useEffect(() => {
if (synthRef.current) {
// Convert 0-100 to dB scale (-60 to 0)
const dbValue = midiVolume === 0 ? -Infinity : (midiVolume / 100) * 60 - 60
synthRef.current.volume.value = dbValue
}
}, [midiVolume])
useEffect(() => {
if (!audioUrl) return
return () => {
URL.revokeObjectURL(audioUrl)
}
}, [audioUrl])
const ensureSynth = async () => {
await Tone.start()
if (!synthRef.current) {
synthRef.current = new Tone.PolySynth(Tone.Synth).toDestination()
// Apply current volume
const dbValue = midiVolume === 0 ? -Infinity : (midiVolume / 100) * 60 - 60
synthRef.current.volume.value = dbValue
}
}
const playPreviewNote = useCallback(async (midi: number) => {
await ensureSynth()
const frequency = Tone.Frequency(midi, 'midi').toFrequency()
synthRef.current?.triggerAttackRelease(frequency, '8n', Tone.now(), 0.7)
}, [midiVolume])
useEffect(() => {
const onKeyDown = (event: KeyboardEvent) => {
if (!selectedId) return
const target = event.target as HTMLElement | null
if (target && ['INPUT', 'TEXTAREA'].includes(target.tagName)) return
// Delete note
if (event.key === 'Backspace' || event.key === 'Delete') {
event.preventDefault()
removeNote(selectedId)
select(null)
return
}
// Cmd/Ctrl + Up/Down to adjust pitch
const isCmdOrCtrl = event.metaKey || event.ctrlKey
if (isCmdOrCtrl && (event.key === 'ArrowUp' || event.key === 'ArrowDown')) {
event.preventDefault()
const selectedNote = notes.find(n => n.id === selectedId)
if (!selectedNote) return
const delta = event.key === 'ArrowUp' ? 1 : -1
const newMidi = Math.max(LOW_NOTE, Math.min(HIGH_NOTE, selectedNote.midi + delta))
if (newMidi !== selectedNote.midi) {
updateNote(selectedId, { midi: newMidi })
playPreviewNote(newMidi)
}
}
}
window.addEventListener('keydown', onKeyDown)
return () => window.removeEventListener('keydown', onKeyDown)
}, [selectedId, notes, removeNote, select, updateNote, playPreviewNote])
const noteEvents = useMemo<PlayEvent[]>(
() =>
notes.map((note) => ({
time: (60 / tempo) * note.start,
duration: (60 / tempo) * note.duration,
midi: note.midi,
velocity: note.velocity,
})),
[notes, tempo],
)
const beatToSeconds = (beat: number) => beat * (60 / tempo)
const secondsToBeat = (seconds: number) => seconds / (60 / tempo)
const seekBySeconds = (deltaSeconds: number) => {
const maxNoteEnd = notes.reduce((acc, n) => Math.max(acc, n.start + n.duration), 0)
const maxBeat = Math.max(secondsToBeat(audioDuration), maxNoteEnd)
const nextSeconds = Math.max(0, Math.min(beatToSeconds(maxBeat), beatToSeconds(playhead) + deltaSeconds))
seekToBeat(secondsToBeat(nextSeconds))
}
const gridSecondWidth = BASE_GRID_SECOND_WIDTH * horizontalZoom
const rowHeight = BASE_ROW_HEIGHT * verticalZoom
// Calculate MIDI content width to sync with audio track
const midiContentWidth = useMemo(() => {
const noteEndSeconds = notes.reduce((acc, n) => {
const endBeat = n.start + n.duration
return Math.max(acc, beatToSeconds(endBeat))
}, 8)
const maxSeconds = Math.max(noteEndSeconds + 10, audioDuration + 10, 30)
return maxSeconds * gridSecondWidth
}, [notes, audioDuration, gridSecondWidth, beatToSeconds])
const seekToBeat = (beat: number) => {
setPlayhead(beat)
Tone.Transport.seconds = beatToSeconds(beat)
if (audioRef.current) {
audioRef.current.currentTime = beatToSeconds(beat)
}
}
const schedulePlayback = async () => {
if (!notes.length && !audioUrl) return
await ensureSynth()
partRef.current?.dispose()
Tone.Transport.cancel()
Tone.Transport.stop()
Tone.Transport.bpm.value = tempo
// Determine playback range
const hasSelection = selectionStart !== null && selectionEnd !== null && selectionEnd > selectionStart
const startSeconds = hasSelection ? selectionStart : beatToSeconds(playhead)
const endSeconds = hasSelection ? selectionEnd : null
Tone.Transport.seconds = startSeconds
// Filter notes within selection range if applicable
const filteredEvents = hasSelection
? noteEvents.filter(e => e.time >= startSeconds && e.time < endSeconds!)
: noteEvents
if (filteredEvents.length) {
partRef.current = new Tone.Part((time, event) => {
if (midiVolume === 0) return
const frequency = Tone.Frequency(event.midi, 'midi').toFrequency()
synthRef.current?.triggerAttackRelease(frequency, event.duration, time, event.velocity)
}, filteredEvents)
partRef.current.start(0)
}
Tone.Transport.start()
if (audioRef.current && audioUrl) {
audioRef.current.currentTime = startSeconds
if (audioVolume > 0) {
audioRef.current.play().catch(() => null)
}
}
setIsPlaying(true)
setStatus(hasSelection ? '选区回放中...' : '正在回放...')
const tick = () => {
const seconds =
audioRef.current && audioUrl && !audioRef.current.paused
? audioRef.current.currentTime
: Tone.Transport.seconds
// Stop at selection end
if (endSeconds !== null && seconds >= endSeconds) {
pausePlayback()
seekToBeat(secondsToBeat(selectionStart!))
setStatus('选区播放完毕')
return
}
const beat = seconds / (60 / tempo)
setPlayhead(beat)
rafRef.current = requestAnimationFrame(tick)
}
rafRef.current = requestAnimationFrame(tick)
}
const stopPlayback = () => {
Tone.Transport.stop()
Tone.Transport.cancel()
partRef.current?.dispose()
partRef.current = null
setIsPlaying(false)
setPlayhead(0)
if (audioRef.current) {
audioRef.current.pause()
audioRef.current.currentTime = 0
}
if (rafRef.current) {
cancelAnimationFrame(rafRef.current)
rafRef.current = null
}
}
const pausePlayback = () => {
Tone.Transport.stop()
partRef.current?.dispose()
partRef.current = null
setIsPlaying(false)
if (audioRef.current) {
audioRef.current.pause()
}
if (rafRef.current) {
cancelAnimationFrame(rafRef.current)
rafRef.current = null
}
}
const handlePlayToggle = async () => {
if (isPlaying) {
pausePlayback()
setStatus('已暂停')
} else {
await schedulePlayback()
}
}
const handleImportClick = () => fileInputRef.current?.click()
const handleAudioImportClick = () => audioInputRef.current?.click()
const handleFileChange = async (event: React.ChangeEvent<HTMLInputElement>) => {
const file = event.target.files?.[0]
if (!file) return
try {
const snapshot = await importMidiFile(file)
setNotes(snapshot.notes)
setTempo(snapshot.tempo)
setTimeSignature(snapshot.timeSignature as TimeSignature)
setPpq(snapshot.ppq) // Preserve original ppq for accurate export
setStatus(`已载入 ${file.name}`)
} catch (error) {
console.error(error)
setStatus('导入失败,请确认文件合法')
} finally {
event.target.value = ''
}
}
const handleAudioChange = (event: React.ChangeEvent<HTMLInputElement>) => {
const file = event.target.files?.[0]
if (!file) return
// Validate audio file type
const validAudioTypes = ['audio/mpeg', 'audio/wav', 'audio/ogg', 'audio/flac', 'audio/mp4', 'audio/aac', 'audio/x-m4a']
const validExtensions = ['.mp3', '.wav', '.ogg', '.flac', '.m4a', '.aac']
const fileName = file.name.toLowerCase()
const isValidType = validAudioTypes.includes(file.type) || file.type.startsWith('audio/')
const isValidExtension = validExtensions.some(ext => fileName.endsWith(ext))
if (!isValidType && !isValidExtension) {
setStatus(`不支持的文件格式,请选择音频文件(${validExtensions.join(', ')})`)
event.target.value = ''
return
}
const url = URL.createObjectURL(file)
setAudioUrl(url)
setStatus(`已载入音频 ${file.name}`)
event.target.value = ''
}
// Check for overlapping notes (any pitch)
const getOverlappingNotes = () => {
const overlapping: string[] = []
const sortedNotes = [...notes].sort((a, b) => a.start - b.start)
const EPSILON = 0.05 // Tolerance for floating point comparison
for (let i = 0; i < sortedNotes.length; i++) {
for (let j = i + 1; j < sortedNotes.length; j++) {
const noteA = sortedNotes[i]
const noteB = sortedNotes[j]
const noteAEnd = noteA.start + noteA.duration
// If noteB starts at or after noteA ends (with tolerance), no overlap
if (noteB.start >= noteAEnd - EPSILON) break
// True overlap: noteB starts before noteA ends
if (!overlapping.includes(noteA.id)) overlapping.push(noteA.id)
if (!overlapping.includes(noteB.id)) overlapping.push(noteB.id)
}
}
return overlapping
}
// Auto-fix overlapping notes by trimming the first note to end where the second begins
const handleFixOverlaps = () => {
const sortedNotes = [...notes].sort((a, b) => a.start - b.start)
let fixCount = 0
for (let i = 0; i < sortedNotes.length - 1; i++) {
const noteA = sortedNotes[i]
const noteB = sortedNotes[i + 1]
const noteAEnd = noteA.start + noteA.duration
// If noteA overlaps with noteB
if (noteAEnd > noteB.start) {
// Trim noteA to end at noteB's start
const newDuration = Math.max(0.01, noteB.start - noteA.start)
updateNote(noteA.id, { duration: newDuration })
fixCount++
}
}
if (fixCount > 0) {
setStatus(`已修复 ${fixCount} 个重叠音符`)
} else {
setStatus('没有检测到重叠音符')
}
}
const handleExport = () => {
const overlapping = getOverlappingNotes()
if (overlapping.length > 0) {
const confirm = window.confirm(
`检测到 ${overlapping.length} 个音符存在时间重叠(标红色的音符),这可能导致播放异常。\n\n是否仍要导出?`
)
if (!confirm) return
}
const blob = exportMidi({ notes, tempo, timeSignature, ppq })
const url = URL.createObjectURL(blob)
const anchor = document.createElement('a')
anchor.href = url
anchor.download = 'vocal-midi.mid'
anchor.click()
URL.revokeObjectURL(url)
setStatus('已导出包含歌词的 MIDI 文件')
}
return (
<div className="app-shell">
<header className="topbar">
<div>
<p className="eyebrow">歌声 MIDI 编辑器</p>
<h1>Lyric-ready Piano Roll</h1>
<p className="muted">导入、拖拽、实时修改歌词并导出标准 MIDI。</p>
</div>
<div className="actions">
<button className="icon-toggle" onClick={() => setTheme(theme === 'dark' ? 'light' : 'dark')}>
{theme === 'dark' ? (
<span className="icon" aria-label="切换到亮色">
☀️
</span>
) : (
<span className="icon" aria-label="切换到暗色">
🌙
</span>
)}
</button>
<button className="primary" onClick={handleImportClick}>
导入 MIDI
</button>
<button className="primary" onClick={handleExport}>
导出含歌词 MIDI
</button>
<button className="soft" onClick={handleFixOverlaps} title="自动消除重叠:将重叠音符的音尾提前到下一个音的音头">
消除重叠
</button>
<input ref={fileInputRef} type="file" accept=".mid,.midi" className="sr-only" onChange={handleFileChange} />
</div>
</header>
<section className="audio-bar">
<div className="audio-left">
<button className="ghost" onClick={handleAudioImportClick}>
对齐音频导入
</button>
<input
ref={audioInputRef}
type="file"
accept=".mp3,.wav,.ogg,.flac,.m4a,.aac"
className="sr-only"
onChange={handleAudioChange}
/>
<span className="audio-hint">导入后显示音频波形并与 MIDI 同步走带</span>
</div>
<div className="audio-right">
<div className="volume-control">
<span className="volume-label">MIDI</span>
<input
type="range"
min={0}
max={100}
value={midiVolume}
onChange={(e) => setMidiVolume(Number(e.target.value))}
className="volume-slider"
/>
<span className="volume-value">{midiVolume}%</span>
</div>
<div className="volume-control">
<span className="volume-label">音频</span>
<input
type="range"
min={0}
max={100}
value={audioVolume}
onChange={(e) => setAudioVolume(Number(e.target.value))}
className="volume-slider"
/>
<span className="volume-value">{audioVolume}%</span>
</div>
</div>
</section>
<section className="panel panel-split">
<div className="panel-main">
{audioUrl && (
<AudioTrack
key={audioUrl}
ref={audioScrollRef}
audioUrl={audioUrl}
muted={audioVolume === 0}
onSeek={(seconds) => seekToBeat(secondsToBeat(seconds))}
playheadSeconds={beatToSeconds(playhead)}
gridSecondWidth={gridSecondWidth}
minContentWidth={midiContentWidth}
/>
)}
<PianoRoll
notes={notes}
selectedId={selectedId}
timeSignature={timeSignature}
tempo={tempo}
playhead={playhead}
selectionStart={selectionStart}
selectionEnd={selectionEnd}
onAddNote={addNote}
onSelect={select}
onUpdateNote={updateNote}
onSeek={seekToBeat}
onScroll={(left) => {
if (audioScrollRef.current) {
audioScrollRef.current.scrollLeft = left
}
}}
onZoom={(deltaH, deltaV) => {
if (deltaH !== 0) {
setHorizontalZoom(prev => Math.max(0.5, prev + deltaH))
}
if (deltaV !== 0) {
setVerticalZoom(prev => Math.max(0.6, Math.min(2.5, prev + deltaV)))
}
}}
onPlayNote={playPreviewNote}
onFocusLyric={(noteId) => {
select(noteId)
setFocusLyricId(noteId)
}}
onSelectionChange={(start, end) => {
setSelectionStart(start)
setSelectionEnd(end)
}}
isSelectingRange={isSelectingRange}
audioDuration={audioDuration}
gridSecondWidth={gridSecondWidth}
rowHeight={rowHeight}
/>
</div>
<aside className="panel-side">
<div className="controls">
<div className="toggle" style={{ justifyContent: 'space-between' }}>
<span>水平缩放</span>
<input
type="range"
min={0.5}
max={10}
step={0.1}
value={Math.min(horizontalZoom, 10)}
onChange={(e) => setHorizontalZoom(Number(e.target.value))}
style={{ width: '140px' }}
/>
<span style={{ width: 42, textAlign: 'right' }}>{horizontalZoom.toFixed(1)}x</span>
</div>
<div className="toggle" style={{ justifyContent: 'space-between' }}>
<span>垂直缩放</span>
<input
type="range"
min={0.6}
max={2.5}
step={0.1}
value={verticalZoom}
onChange={(e) => setVerticalZoom(Number(e.target.value))}
style={{ width: '140px' }}
/>
<span style={{ width: 42, textAlign: 'right' }}>{verticalZoom.toFixed(1)}x</span>
</div>
<div className="transport">
<button
className="soft"
onClick={() => {
setPlayhead(0)
seekToBeat(0)
}}
title="回到开头"
>
⏮
</button>
<button
className="soft"
onClick={() => seekBySeconds(-2)}
title="后退 2 秒"
>
⏪ 2s
</button>
<button
className="primary"
onClick={handlePlayToggle}
disabled={!notes.length && !audioUrl}
title={isPlaying ? "暂停" : (selectionStart !== null && selectionEnd !== null ? "播放选区" : "播放")}
>
{isPlaying ? '⏸' : '▶'}
</button>
<button
className="soft"
onClick={() => seekBySeconds(2)}
title="前进 2 秒"
>
2s ⏩
</button>
<button
className="soft"
onClick={() => {
// Logic to find end of song? Max note end or audio duration
const maxNoteEnd = notes.reduce((acc, n) => Math.max(acc, n.start + n.duration), 0)
seekToBeat(Math.max(secondsToBeat(audioDuration), maxNoteEnd))
}}
title="回到结尾"
>
⏭
</button>
</div>
<div className="selection-controls">
<button
className={`soft selection-btn ${isSelectingRange ? 'active' : ''}`}
onClick={() => setIsSelectingRange(!isSelectingRange)}
title={isSelectingRange ? "退出选区模式" : "设置选区:在时间轴上拖拽选择播放范围"}
>
{isSelectingRange ? '📍 选区中' : '📍 设选区'}
</button>
{selectionStart !== null && selectionEnd !== null && (
<>
<span className="selection-info">
{selectionStart.toFixed(1)}s - {selectionEnd.toFixed(1)}s
</span>
<button
className="soft"
onClick={() => {
setSelectionStart(null)
setSelectionEnd(null)
}}
title="清除选区"
>
✕
</button>
</>
)}
</div>
<div className="status">{status}</div>
</div>
<div className="lyric-container">
<LyricTable
notes={notes}
selectedId={selectedId}
tempo={tempo}
focusLyricId={focusLyricId}
onSelect={select}
onUpdate={updateNote}
onFocusHandled={() => setFocusLyricId(null)}
/>
</div>
</aside>
</section>
<audio
ref={audioRef}
src={audioUrl ?? undefined}
preload="auto"
className="sr-only"
onLoadedMetadata={(e) => {
setAudioDuration(e.currentTarget.duration)
// Ensure volume is set when audio loads
e.currentTarget.volume = audioVolume / 100
}}
/>
</div>
)
}
export default App
@@ -0,0 +1,182 @@
import { useEffect, useRef, forwardRef, useState } from 'react'
import WaveSurfer from 'wavesurfer.js'
import { PITCH_WIDTH } from '../constants'
export type AudioTrackProps = {
audioUrl: string | null
muted: boolean
onSeek: (seconds: number) => void
mediaElement?: HTMLAudioElement | null
playheadSeconds: number
gridSecondWidth: number
minContentWidth?: number // Minimum width to match MIDI editor area
}
export const AudioTrack = forwardRef<HTMLDivElement, AudioTrackProps>(
({ audioUrl, muted, onSeek, playheadSeconds, gridSecondWidth, minContentWidth = 0 }, ref) => {
const containerRef = useRef<HTMLDivElement | null>(null)
const waveRef = useRef<WaveSurfer | null>(null)
const [waveWidth, setWaveWidth] = useState(0)
useEffect(() => {
if (!containerRef.current) return
if (!audioUrl) {
try {
waveRef.current?.destroy()
} catch {
// ignore teardown errors
}
waveRef.current = null
setWaveWidth(0)
return
}
let cancelled = false
// Clean up existing instance
if (waveRef.current) {
try {
waveRef.current.destroy()
} catch {
// ignore teardown errors
}
}
waveRef.current = WaveSurfer.create({
container: containerRef.current,
waveColor: '#4b64bc',
progressColor: '#4b64bc',
cursorColor: 'transparent',
barWidth: 2,
barGap: 2,
height: 60,
normalize: true,
minPxPerSec: gridSecondWidth,
interact: false,
hideScrollbar: true,
autoScroll: false,
})
waveRef.current.load(audioUrl).catch(() => null)
waveRef.current.on('error', () => null)
waveRef.current.on('ready', () => {
if (cancelled || !waveRef.current) return
const duration = waveRef.current.getDuration()
const requiredWidth = duration * gridSecondWidth
setWaveWidth(requiredWidth)
})
return () => {
cancelled = true
try {
waveRef.current?.destroy()
} catch {
// ignore teardown errors
}
waveRef.current = null
}
}, [audioUrl, gridSecondWidth])
useEffect(() => {
if (!waveRef.current) return
waveRef.current.setOptions({
waveColor: muted ? '#9aa6b2' : '#4b64bc',
progressColor: muted ? '#c0c9d4' : '#4b64bc',
})
}, [muted])
if (!audioUrl) return null
// Content width should be at least as wide as MIDI editor
const contentWidth = Math.max(waveWidth, minContentWidth)
return (
<div
className="audio-track-row"
style={{
display: 'flex',
borderBottom: '1px solid var(--border-soft)',
height: '70px',
flexShrink: 0
}}
>
<div
className="audio-gutter"
style={{
width: PITCH_WIDTH,
flexShrink: 0,
background: 'var(--panel-strong)',
borderRight: '1px solid var(--border-subtle)',
display: 'flex',
alignItems: 'center',
justifyContent: 'center',
fontSize: '11px',
color: 'var(--text-muted)',
fontWeight: 600,
}}
>
AUDIO
</div>
{/* Scroll Mask - Controlled by parent via ref */}
<div
ref={ref}
className="audio-scroll-mask"
style={{
flex: 1,
overflow: 'hidden',
position: 'relative',
background: 'var(--panel-soft)',
}}
onClick={(e) => {
const rect = e.currentTarget.getBoundingClientRect()
const scrollMask = e.currentTarget as HTMLDivElement
const x = e.clientX - rect.left + scrollMask.scrollLeft
const seconds = x / gridSecondWidth
onSeek(seconds)
}}
>
{/* Container that matches MIDI editor width */}
<div
className="audio-content"
style={{
width: contentWidth > 0 ? contentWidth : '100%',
height: '100%',
position: 'relative'
}}
>
{/* WaveSurfer container - only as wide as audio */}
<div
ref={containerRef}
className="wave-container"
style={{
width: waveWidth > 0 ? waveWidth : '100%',
height: '100%',
position: 'absolute',
left: 0,
top: 0
}}
/>
{/* Custom Playhead */}
<div
className="audio-playhead"
style={{
position: 'absolute',
top: 0,
bottom: 0,
width: '2px',
background: '#ff7043',
boxShadow: '0 0 12px rgba(255, 112, 67, 0.6)',
left: playheadSeconds * gridSecondWidth,
zIndex: 10,
pointerEvents: 'none',
}}
/>
</div>
</div>
</div>
)
}
)
@@ -0,0 +1,288 @@
import { useEffect, useMemo, useRef, useState } from 'react'
import type { NoteEvent } from '../types'
export type LyricTableProps = {
notes: NoteEvent[]
selectedId: string | null
tempo: number
focusLyricId: string | null
onSelect: (id: string | null) => void
onUpdate: (id: string, patch: Partial<NoteEvent>) => void
onScrollToNote?: (noteId: string) => void
onFocusHandled?: () => void
}
const formatSeconds = (beats: number, tempo: number) => {
const seconds = beats * (60 / tempo)
return Number.parseFloat(seconds.toFixed(2))
}
const secondsToBeats = (seconds: number, tempo: number) => {
return seconds * (tempo / 60)
}
// Editable cell with confirmation
function EditableCell({
value,
noteId,
field,
tempo,
onConfirm,
type = 'number',
min,
step
}: {
value: number
noteId: string
field: 'midi' | 'start' | 'end'
tempo: number
onConfirm: (noteId: string, field: string, value: number) => void
type?: string
min?: number
step?: number
}) {
const displayValue = field === 'midi' ? value : formatSeconds(value, tempo)
const [localValue, setLocalValue] = useState(String(displayValue))
const [isDirty, setIsDirty] = useState(false)
const inputRef = useRef<HTMLInputElement>(null)
// Sync with external value when it changes (and not dirty)
useEffect(() => {
if (!isDirty) {
setLocalValue(String(displayValue))
}
}, [displayValue, isDirty])
const handleChange = (e: React.ChangeEvent<HTMLInputElement>) => {
setLocalValue(e.target.value)
setIsDirty(true)
}
const handleConfirm = () => {
const parsed = parseFloat(localValue)
if (!isNaN(parsed)) {
if (field === 'midi') {
if (parsed >= 0 && parsed <= 127) {
onConfirm(noteId, field, Math.round(parsed))
}
} else {
if (parsed >= 0) {
onConfirm(noteId, field, secondsToBeats(parsed, tempo))
}
}
}
setIsDirty(false)
}
const handleKeyDown = (e: React.KeyboardEvent) => {
if (e.key === 'Enter') {
e.preventDefault()
handleConfirm()
inputRef.current?.blur()
} else if (e.key === 'Escape') {
setLocalValue(String(displayValue))
setIsDirty(false)
inputRef.current?.blur()
}
}
const handleBlur = () => {
if (isDirty) {
// Reset to original on blur without confirm
setLocalValue(String(displayValue))
setIsDirty(false)
}
}
return (
<div className="editable-cell">
<input
ref={inputRef}
className={`lyric-meta-input ${isDirty ? 'lyric-meta-dirty' : ''}`}
type={type}
min={min}
step={step}
value={localValue}
onChange={handleChange}
onKeyDown={handleKeyDown}
onBlur={handleBlur}
onClick={(e) => e.stopPropagation()}
/>
{isDirty && (
<button
className="confirm-btn"
onMouseDown={(e) => {
e.preventDefault() // Prevent input blur
e.stopPropagation()
}}
onClick={(e) => {
e.stopPropagation()
handleConfirm()
}}
title="确认修改 (Enter)"
>
✓
</button>
)}
</div>
)
}
export function LyricTable({ notes, selectedId, tempo, focusLyricId, onSelect, onUpdate, onScrollToNote, onFocusHandled }: LyricTableProps) {
const listRef = useRef<HTMLDivElement | null>(null)
const inputRefs = useRef<Map<string, HTMLInputElement>>(new Map())
const sorted = useMemo(() => [...notes].sort((a, b) => a.start - b.start), [notes])
// Scroll to selected note (no auto-focus on single click)
useEffect(() => {
if (!selectedId || !listRef.current) return
const target = listRef.current.querySelector<HTMLDivElement>(`[data-note-id="${selectedId}"]`)
if (target) {
target.scrollIntoView({ block: 'nearest', behavior: 'smooth' })
}
}, [selectedId])
// Focus lyric input when requested (double-click on note or click on list row)
useEffect(() => {
if (!focusLyricId) return
const input = inputRefs.current.get(focusLyricId)
if (input) {
setTimeout(() => {
input.focus()
input.select()
}, 50)
}
onFocusHandled?.()
}, [focusLyricId, onFocusHandled])
// Fill lyrics from selected note onwards
const handleBulkFill = (bulkText: string) => {
if (!sorted.length) return
const chars = Array.from(bulkText.replace(/\s+/g, ''))
let startIndex = 0
if (selectedId) {
const selectedIndex = sorted.findIndex(n => n.id === selectedId)
if (selectedIndex >= 0) {
startIndex = selectedIndex
}
}
let charIndex = 0
for (let i = startIndex; i < sorted.length && charIndex < chars.length; i++) {
onUpdate(sorted[i].id, { lyric: chars[charIndex] })
charIndex++
}
}
const handleRowClick = (noteId: string) => {
onSelect(noteId)
onScrollToNote?.(noteId)
}
const handleFieldConfirm = (noteId: string, field: string, value: number) => {
const note = notes.find(n => n.id === noteId)
if (!note) return
if (field === 'midi') {
onUpdate(noteId, { midi: value })
} else if (field === 'start') {
// Keep END the same, adjust duration accordingly
const currentEnd = note.start + note.duration
const newDuration = Math.max(0.01, currentEnd - value)
onUpdate(noteId, { start: value, duration: newDuration })
} else if (field === 'end') {
// End changed, update duration
const newDuration = Math.max(0.01, value - note.start)
onUpdate(noteId, { duration: newDuration })
}
}
return (
<div className="lyric-card">
<div className="lyric-bulk">
<textarea
className="lyric-bulk-input"
rows={2}
placeholder={selectedId ? "从选中音符开始按字填充" : "输入歌词,点击按字填充"}
onKeyDown={(e) => {
if (e.key === 'Enter' && !e.shiftKey) {
e.preventDefault()
handleBulkFill(e.currentTarget.value)
}
}}
/>
<button
className="soft"
type="button"
onClick={(e) => {
const textarea = e.currentTarget.previousElementSibling as HTMLTextAreaElement
handleBulkFill(textarea.value)
}}
>
按字<br/>填充
</button>
</div>
<div className="lyric-header" style={{ flexShrink: 0 }}>
<div>LYRIC</div>
<div>PITCH</div>
<div>START</div>
<div>END</div>
</div>
<div className="lyric-list" ref={listRef}>
{sorted.map((note) => (
<div
key={note.id}
className={`lyric-row ${selectedId === note.id ? 'lyric-row-active' : ''}`}
data-note-id={note.id}
onClick={() => handleRowClick(note.id)}
>
<input
ref={(el) => {
if (el) {
inputRefs.current.set(note.id, el)
} else {
inputRefs.current.delete(note.id)
}
}}
className="lyric-input"
value={note.lyric}
placeholder="Type lyric"
onChange={(event) => onUpdate(note.id, { lyric: event.target.value })}
onClick={(e) => e.stopPropagation()}
/>
<EditableCell
value={note.midi}
noteId={note.id}
field="midi"
tempo={tempo}
onConfirm={handleFieldConfirm}
min={0}
/>
<EditableCell
value={note.start}
noteId={note.id}
field="start"
tempo={tempo}
onConfirm={handleFieldConfirm}
min={0}
step={0.01}
/>
<EditableCell
value={note.start + note.duration}
noteId={note.id}
field="end"
tempo={tempo}
onConfirm={handleFieldConfirm}
min={0}
step={0.01}
/>
</div>
))}
{sorted.length === 0 && <div className="lyric-empty">Import或双击钢琴卷帘以添加音符</div>}
</div>
</div>
)
}
@@ -0,0 +1,704 @@
import { useEffect, useMemo, useRef, useState, useCallback, memo } from 'react'
import type React from 'react'
import type { NoteEvent, TimeSignature } from '../types'
import { PITCH_WIDTH, LOW_NOTE, HIGH_NOTE } from '../constants'
const midiToName = (midi: number) => {
const names = ['C', 'C#', 'D', 'D#', 'E', 'F', 'F#', 'G', 'G#', 'A', 'A#', 'B']
const octave = Math.floor(midi / 12) - 1
return `${names[midi % 12]}${octave}`
}
// Memoized note component to prevent unnecessary re-renders
const NoteChip = memo(function NoteChip({
note,
left,
top,
width,
height,
fontSize,
isSelected,
isOverlapping,
onPointerDown,
onDoubleClick,
}: {
note: NoteEvent
left: number
top: number
width: number
height: number
fontSize: number
isSelected: boolean
isOverlapping: boolean
onPointerDown: (event: React.PointerEvent<HTMLDivElement>, mode: 'move' | 'resize-start' | 'resize-end') => void
onDoubleClick: (event: React.MouseEvent<HTMLDivElement>) => void
}) {
return (
<div
className={`note-chip ${isSelected ? 'note-active' : ''} ${isOverlapping ? 'note-overlap' : ''}`}
style={{
left,
top: top + 1,
width,
height,
willChange: 'transform', // GPU acceleration hint
}}
onPointerDown={(e) => onPointerDown(e, 'move')}
onDoubleClick={onDoubleClick}
>
<div className="note-label" style={{ fontSize }}>
<span>{note.lyric || '\u00a0'}</span>
</div>
<div className="note-handle start" onPointerDown={(e) => { e.stopPropagation(); onPointerDown(e, 'resize-start') }} />
<div className="note-handle end" onPointerDown={(e) => { e.stopPropagation(); onPointerDown(e, 'resize-end') }} />
</div>
)
})
// Dynamic snap based on zoom level - higher zoom = finer snap
const getSnapSeconds = (gridSecondWidth: number) => {
// At base width (80px/s), snap is 0.1s
// At 2x zoom (160px/s), snap is 0.05s
// At 4x zoom (320px/s), snap is 0.025s
// At 8x zoom (640px/s), snap is 0.01s
const baseSnap = 0.1
const zoomFactor = gridSecondWidth / 80
return Math.max(0.01, baseSnap / zoomFactor)
}
const snapSeconds = (value: number, gridSecondWidth: number) => {
const snap = getSnapSeconds(gridSecondWidth)
return Math.max(0, Math.round(value / snap) * snap)
}
export type PianoRollProps = {
notes: NoteEvent[]
selectedId: string | null
timeSignature: TimeSignature
tempo: number
playhead: number // in beats
selectionStart: number | null // in seconds
selectionEnd: number | null // in seconds
onAddNote: (note: Partial<NoteEvent>) => NoteEvent
onUpdateNote: (id: string, patch: Partial<NoteEvent>) => void
onSelect: (id: string | null) => void
onSeek: (beat: number) => void
onScroll?: (left: number) => void
onZoom?: (deltaH: number, deltaV: number) => void
onPlayNote?: (midi: number) => void
onFocusLyric?: (noteId: string) => void
onSelectionChange?: (start: number | null, end: number | null) => void
isSelectingRange?: boolean
audioDuration?: number
gridSecondWidth: number
rowHeight: number
}
export function PianoRoll({
notes,
selectedId,
timeSignature: _timeSignature,
tempo,
playhead,
selectionStart,
selectionEnd,
onAddNote,
onSelect,
onUpdateNote,
onSeek,
onScroll,
onZoom,
onPlayNote,
onFocusLyric,
onSelectionChange,
isSelectingRange = false,
audioDuration = 0,
gridSecondWidth,
rowHeight
}: PianoRollProps) {
const scrollContainerRef = useRef<HTMLDivElement | null>(null)
const rulerScrollRef = useRef<HTMLDivElement | null>(null)
const [scrollTop, setScrollTop] = useState(0)
const [scrollLeft, setScrollLeft] = useState(0)
const [viewportWidth, setViewportWidth] = useState(800)
const [viewportHeight, setViewportHeight] = useState(400)
const dragRef = useRef<{
id: string
mode: 'move' | 'resize-start' | 'resize-end'
originX: number
originY: number
startSeconds: number
durationSeconds: number
midi: number
lastMidi?: number // Track last midi for pitch change sound
} | null>(null)
// Selection drag state
const selectionDragRef = useRef<{
startX: number
startSeconds: number
} | null>(null)
// Store callbacks in refs to avoid stale closures in event handlers
const onPlayNoteRef = useRef(onPlayNote)
const onUpdateNoteRef = useRef(onUpdateNote)
useEffect(() => {
onPlayNoteRef.current = onPlayNote
onUpdateNoteRef.current = onUpdateNote
}, [onPlayNote, onUpdateNote])
// Conversion helpers
const beatToSeconds = useCallback((beat: number) => beat * (60 / tempo), [tempo])
const secondsToBeat = useCallback((seconds: number) => seconds / (60 / tempo), [tempo])
// Calculate dimensions
const totalRows = HIGH_NOTE - LOW_NOTE + 1
const contentHeight = totalRows * rowHeight
const [containerWidth, setContainerWidth] = useState(1200)
// Track container size
useEffect(() => {
const container = scrollContainerRef.current
if (!container) return
const observer = new ResizeObserver((entries) => {
for (const entry of entries) {
setContainerWidth(entry.contentRect.width)
setViewportWidth(entry.contentRect.width)
setViewportHeight(entry.contentRect.height)
}
})
observer.observe(container)
return () => observer.disconnect()
}, [])
const maxSeconds = useMemo(() => {
const noteEndSeconds = notes.reduce((acc, n) => {
const endBeat = n.start + n.duration
return Math.max(acc, beatToSeconds(endBeat))
}, 8)
// Ensure grid extends at least 2x the visible area for smoother scrolling
const minSecondsForView = (containerWidth / gridSecondWidth) * 2
return Math.max(noteEndSeconds + 10, audioDuration + 10, minSecondsForView, 30)
}, [notes, audioDuration, beatToSeconds, containerWidth, gridSecondWidth])
const contentWidth = maxSeconds * gridSecondWidth
// Drag handlers - use refs to avoid stale closure issues
const handlePointerMove = useCallback((event: PointerEvent) => {
const drag = dragRef.current
if (!drag) return
const dxSeconds = (event.clientX - drag.originX) / gridSecondWidth
const dy = (event.clientY - drag.originY) / rowHeight
if (drag.mode === 'move') {
const nextSeconds = snapSeconds(drag.startSeconds + dxSeconds, gridSecondWidth)
const nextMidi = Math.min(HIGH_NOTE, Math.max(LOW_NOTE, Math.round(drag.midi - dy)))
// Play sound when pitch changes
if (nextMidi !== drag.lastMidi && onPlayNoteRef.current) {
onPlayNoteRef.current(nextMidi)
drag.lastMidi = nextMidi
}
onUpdateNoteRef.current(drag.id, {
start: secondsToBeat(nextSeconds),
midi: nextMidi
})
}
if (drag.mode === 'resize-start') {
const nextSeconds = snapSeconds(drag.startSeconds + dxSeconds, gridSecondWidth)
const delta = drag.startSeconds - nextSeconds
const nextDurationSeconds = Math.max(0.05, drag.durationSeconds + delta)
onUpdateNoteRef.current(drag.id, {
start: secondsToBeat(nextSeconds),
duration: secondsToBeat(nextDurationSeconds)
})
}
if (drag.mode === 'resize-end') {
const nextDurationSeconds = Math.max(0.05, snapSeconds(drag.durationSeconds + dxSeconds, gridSecondWidth))
onUpdateNoteRef.current(drag.id, { duration: secondsToBeat(nextDurationSeconds) })
}
}, [gridSecondWidth, rowHeight, secondsToBeat])
const handlePointerUp = useCallback(() => {
dragRef.current = null
window.removeEventListener('pointermove', handlePointerMove)
window.removeEventListener('pointerup', handlePointerUp)
}, [handlePointerMove])
useEffect(() => {
return () => {
window.removeEventListener('pointermove', handlePointerMove)
window.removeEventListener('pointerup', handlePointerUp)
}
}, [handlePointerMove, handlePointerUp])
// Scroll sync
useEffect(() => {
const container = scrollContainerRef.current
const ruler = rulerScrollRef.current
if (!container || !ruler) return
const handleScroll = () => {
ruler.scrollLeft = container.scrollLeft
setScrollTop(container.scrollTop)
setScrollLeft(container.scrollLeft)
if (onScroll) onScroll(container.scrollLeft)
}
container.addEventListener('scroll', handleScroll)
return () => container.removeEventListener('scroll', handleScroll)
}, [onScroll])
// Zoom support via wheel/trackpad
// Mac: Cmd+滚轮 (水平缩放), Cmd+Shift+滚轮 (垂直缩放), 或双指捏合
// Windows/Linux: Ctrl+滚轮 (水平缩放), Ctrl+Shift+滚轮 (垂直缩放)
useEffect(() => {
const container = scrollContainerRef.current
if (!container || !onZoom) return
const handleWheel = (e: WheelEvent) => {
// Ctrl (Windows/Linux/捏合) or Cmd (Mac) triggers zoom
const isZoomTrigger = e.ctrlKey || e.metaKey
if (isZoomTrigger) {
e.preventDefault()
e.stopPropagation()
// Use deltaY for zoom amount, normalize for different input methods
// Pinch gestures typically have smaller delta values
let delta = -e.deltaY
if (Math.abs(delta) > 10) {
// Likely a mouse wheel, scale down
delta = delta * 0.01
} else {
// Likely a trackpad pinch, scale appropriately
delta = delta * 0.05
}
// Shift or Alt/Option for vertical zoom, otherwise horizontal
if (e.shiftKey || e.altKey) {
onZoom(0, delta)
} else {
onZoom(delta, 0)
}
}
}
container.addEventListener('wheel', handleWheel, { passive: false })
return () => container.removeEventListener('wheel', handleWheel)
}, [onZoom])
// Playhead auto-scroll
useEffect(() => {
if (!scrollContainerRef.current) return
const container = scrollContainerRef.current
const playheadX = beatToSeconds(playhead) * gridSecondWidth
const viewStart = container.scrollLeft
const viewEnd = viewStart + container.clientWidth
if (playheadX > viewEnd) {
container.scrollLeft = playheadX
} else if (playheadX < viewStart) {
container.scrollLeft = playheadX
}
}, [playhead, gridSecondWidth, beatToSeconds])
// Selection auto-scroll
useEffect(() => {
if (!scrollContainerRef.current || !selectedId) return
const note = notes.find((n) => n.id === selectedId)
if (!note) return
const container = scrollContainerRef.current
const noteX = beatToSeconds(note.start) * gridSecondWidth
const noteY = (HIGH_NOTE - note.midi) * rowHeight
const viewStart = container.scrollLeft
const viewEnd = viewStart + container.clientWidth
if (noteX < viewStart + 50 || noteX > viewEnd - 50) {
container.scrollLeft = Math.max(0, noteX - container.clientWidth * 0.35)
}
const viewTop = container.scrollTop
const viewBottom = viewTop + container.clientHeight
if (noteY < viewTop || noteY > viewBottom - rowHeight) {
container.scrollTop = Math.max(0, noteY - container.clientHeight * 0.4)
}
}, [selectedId, notes, gridSecondWidth, rowHeight, beatToSeconds])
const handleGridDoubleClick = (event: React.MouseEvent<HTMLDivElement>) => {
// Only add note if clicking on empty space (not on a note)
const target = event.target as HTMLElement
if (target.closest('.note-chip')) return
if (!scrollContainerRef.current) return
const container = scrollContainerRef.current
const rect = container.getBoundingClientRect()
const x = event.clientX - rect.left + container.scrollLeft
const y = event.clientY - rect.top + container.scrollTop
const seconds = snapSeconds(x / gridSecondWidth, gridSecondWidth)
const pitch = Math.min(HIGH_NOTE, Math.max(LOW_NOTE, HIGH_NOTE - Math.floor(y / rowHeight)))
const created = onAddNote({
start: secondsToBeat(seconds),
midi: pitch,
duration: secondsToBeat(0.5),
lyric: ''
})
onSelect(created.id)
}
const startDrag = (
event: React.PointerEvent<HTMLDivElement>,
note: NoteEvent,
mode: 'move' | 'resize-start' | 'resize-end',
) => {
event.preventDefault()
event.stopPropagation()
dragRef.current = {
id: note.id,
mode,
originX: event.clientX,
originY: event.clientY,
startSeconds: beatToSeconds(note.start),
durationSeconds: beatToSeconds(note.duration),
midi: note.midi,
lastMidi: note.midi, // Initialize last midi
}
window.addEventListener('pointermove', handlePointerMove)
window.addEventListener('pointerup', handlePointerUp)
onSelect(note.id)
// Play sound when clicking/selecting note
if (onPlayNote) {
onPlayNote(note.midi)
}
}
// Second-based ruler labels
const secondLabels = useMemo(() => {
const labels = [] as Array<{ left: number; label: string }>
const totalSeconds = Math.ceil(maxSeconds)
for (let s = 0; s <= totalSeconds; s += 1) {
labels.push({ left: s * gridSecondWidth, label: `${s}s` })
}
return labels
}, [maxSeconds, gridSecondWidth])
// Piano keys
const pitchRows = useMemo(() => {
const rows = [] as Array<{ midi: number; isBlack: boolean; label: string; isC: boolean }>
const black = new Set([1, 3, 6, 8, 10])
for (let p = HIGH_NOTE; p >= LOW_NOTE; p -= 1) {
const name = midiToName(p)
const isC = p % 12 === 0
rows.push({ midi: p, isBlack: black.has(p % 12), label: name, isC })
}
return rows
}, [])
// Detect overlapping notes using optimized sweep line algorithm
const overlappingNoteIds = useMemo(() => {
if (notes.length < 2) return new Set<string>()
const overlapping = new Set<string>()
const sortedNotes = [...notes].sort((a, b) => a.start - b.start)
const EPSILON = 0.05 // Tolerance for floating point comparison
// Use a sliding window approach - more efficient for typical music data
// Active notes: notes that haven't ended yet
const activeNotes: typeof sortedNotes = []
for (const note of sortedNotes) {
// Remove notes that have ended before current note starts
while (activeNotes.length > 0) {
const firstActive = activeNotes[0]
const firstActiveEnd = firstActive.start + firstActive.duration
if (firstActiveEnd <= note.start + EPSILON) {
activeNotes.shift()
} else {
break
}
}
// Check overlap with remaining active notes
for (const activeNote of activeNotes) {
const activeEnd = activeNote.start + activeNote.duration
if (note.start < activeEnd - EPSILON) {
overlapping.add(activeNote.id)
overlapping.add(note.id)
}
}
// Add current note to active set (maintain sorted order by end time)
const noteEnd = note.start + note.duration
let insertIndex = activeNotes.length
for (let i = 0; i < activeNotes.length; i++) {
const aEnd = activeNotes[i].start + activeNotes[i].duration
if (noteEnd < aEnd) {
insertIndex = i
break
}
}
activeNotes.splice(insertIndex, 0, note)
}
return overlapping
}, [notes])
// Calculate visible area with buffer for smooth scrolling
const BUFFER_PX = 200 // Render notes slightly outside viewport for smooth scrolling
const visibleArea = useMemo(() => {
return {
left: Math.max(0, scrollLeft - BUFFER_PX),
right: scrollLeft + viewportWidth + BUFFER_PX,
top: Math.max(0, scrollTop - BUFFER_PX),
bottom: scrollTop + viewportHeight + BUFFER_PX,
}
}, [scrollLeft, scrollTop, viewportWidth, viewportHeight])
// Filter notes to only render visible ones (virtualization)
const visibleNotes = useMemo(() => {
return notes.filter(note => {
const noteSeconds = beatToSeconds(note.start)
const noteDurationSeconds = beatToSeconds(note.duration)
const noteLeft = noteSeconds * gridSecondWidth
const noteRight = noteLeft + noteDurationSeconds * gridSecondWidth
const noteTop = (HIGH_NOTE - note.midi) * rowHeight
const noteBottom = noteTop + rowHeight
// Check if note intersects with visible area
const horizontallyVisible = noteRight >= visibleArea.left && noteLeft <= visibleArea.right
const verticallyVisible = noteBottom >= visibleArea.top && noteTop <= visibleArea.bottom
return horizontallyVisible && verticallyVisible
})
}, [notes, visibleArea, gridSecondWidth, rowHeight, beatToSeconds])
// Calculate visible grid lines (virtualization)
const visibleGridLines = useMemo(() => {
const startSecond = Math.max(0, Math.floor(visibleArea.left / gridSecondWidth) - 1)
const endSecond = Math.ceil(visibleArea.right / gridSecondWidth) + 1
const startRow = Math.max(0, Math.floor(visibleArea.top / rowHeight) - 1)
const endRow = Math.min(totalRows, Math.ceil(visibleArea.bottom / rowHeight) + 1)
return {
horizontalLines: Array.from({ length: endRow - startRow + 1 }, (_, i) => startRow + i),
verticalLines: Array.from({ length: endSecond - startSecond + 1 }, (_, i) => startSecond + i),
}
}, [visibleArea, gridSecondWidth, rowHeight, totalRows])
const playheadSeconds = beatToSeconds(playhead)
// Selection drag handlers
const handleRulerPointerDown = (event: React.PointerEvent<HTMLDivElement>) => {
if (!isSelectingRange) {
// Normal click to seek
const rect = event.currentTarget.getBoundingClientRect()
const x = event.clientX - rect.left + (rulerScrollRef.current?.scrollLeft ?? 0)
const seconds = x / gridSecondWidth
onSeek(secondsToBeat(seconds))
return
}
// Start selection drag
const rect = event.currentTarget.getBoundingClientRect()
const x = event.clientX - rect.left + (rulerScrollRef.current?.scrollLeft ?? 0)
const seconds = Math.max(0, x / gridSecondWidth)
selectionDragRef.current = {
startX: event.clientX,
startSeconds: seconds,
}
onSelectionChange?.(seconds, seconds)
const handleSelectionMove = (e: PointerEvent) => {
if (!selectionDragRef.current) return
const currentX = e.clientX - rect.left + (rulerScrollRef.current?.scrollLeft ?? 0)
const currentSeconds = Math.max(0, currentX / gridSecondWidth)
const start = Math.min(selectionDragRef.current.startSeconds, currentSeconds)
const end = Math.max(selectionDragRef.current.startSeconds, currentSeconds)
onSelectionChange?.(start, end)
}
const handleSelectionUp = () => {
selectionDragRef.current = null
window.removeEventListener('pointermove', handleSelectionMove)
window.removeEventListener('pointerup', handleSelectionUp)
}
window.addEventListener('pointermove', handleSelectionMove)
window.addEventListener('pointerup', handleSelectionUp)
}
return (
<div className="piano-shell">
{/* Ruler */}
<div className="ruler-shell">
<div className="ruler-spacer" style={{ width: PITCH_WIDTH, flexShrink: 0 }} />
<div
ref={rulerScrollRef}
className={`ruler-scroll ${isSelectingRange ? 'selecting' : ''}`}
onPointerDown={handleRulerPointerDown}
>
<div className="ruler" style={{ width: contentWidth }}>
{secondLabels.map((mark) => (
<div key={mark.left} className="measure-mark" style={{ left: mark.left }}>
<span>{mark.label}</span>
</div>
))}
{/* Selection range indicator */}
{selectionStart !== null && selectionEnd !== null && selectionEnd > selectionStart && (
<div
className="selection-range"
style={{
left: selectionStart * gridSecondWidth,
width: (selectionEnd - selectionStart) * gridSecondWidth
}}
/>
)}
{/* Ruler playhead indicator */}
<div
className="ruler-playhead"
style={{ left: playheadSeconds * gridSecondWidth }}
/>
</div>
</div>
</div>
{/* Main content area */}
<div className="roll-body">
{/* Piano keys - synced with vertical scroll */}
<div className="pitch-rail" style={{ width: PITCH_WIDTH }}>
<div
className="pitch-rail-inner"
style={{
transform: `translateY(${-scrollTop}px)`,
height: contentHeight
}}
>
{pitchRows.map((pitch) => (
<div
key={pitch.midi}
className={`pitch-cell ${pitch.isBlack ? 'pitch-black' : 'pitch-white'} ${pitch.isC ? 'pitch-c' : ''}`}
style={{ height: rowHeight, cursor: 'pointer' }}
onClick={() => onPlayNote?.(pitch.midi)}
onMouseDown={(e) => e.preventDefault()}
>
<span className="pitch-label">{pitch.label}</span>
</div>
))}
</div>
</div>
{/* Scrollable grid area */}
<div
ref={scrollContainerRef}
className="roll-grid"
onDoubleClick={handleGridDoubleClick}
>
<div
className="grid-content"
style={{
width: contentWidth,
height: contentHeight,
position: 'relative'
}}
>
{/* SVG Grid - virtualized for performance */}
<svg
className="grid-svg"
width={contentWidth}
height={contentHeight}
style={{ position: 'absolute', top: 0, left: 0, pointerEvents: 'none' }}
>
{/* Horizontal lines (pitch rows) - only visible ones */}
{visibleGridLines.horizontalLines.map(i => (
<line
key={`h-${i}`}
x1={visibleArea.left}
y1={i * rowHeight}
x2={visibleArea.right}
y2={i * rowHeight}
stroke="var(--grid-line-minor)"
strokeWidth={1}
/>
))}
{/* Vertical lines (seconds) - only visible ones */}
{visibleGridLines.verticalLines.map(i => (
<line
key={`v-${i}`}
x1={i * gridSecondWidth}
y1={visibleArea.top}
x2={i * gridSecondWidth}
y2={visibleArea.bottom}
stroke="var(--grid-line-minor)"
strokeWidth={1}
/>
))}
</svg>
{/* Selection range in grid */}
{selectionStart !== null && selectionEnd !== null && selectionEnd > selectionStart && (
<div
className="grid-selection-range"
style={{
left: selectionStart * gridSecondWidth,
width: (selectionEnd - selectionStart) * gridSecondWidth,
height: contentHeight
}}
/>
)}
{/* Playhead */}
<div
className="playhead"
style={{
left: playheadSeconds * gridSecondWidth,
height: contentHeight
}}
/>
{/* Notes - virtualized: only render visible notes */}
{visibleNotes.map((note) => {
const noteSeconds = beatToSeconds(note.start)
const noteDurationSeconds = beatToSeconds(note.duration)
const left = noteSeconds * gridSecondWidth
const top = (HIGH_NOTE - note.midi) * rowHeight
const noteWidthPx = Math.max(noteDurationSeconds * gridSecondWidth, 4)
const noteHeight = rowHeight - 2
const isOverlapping = overlappingNoteIds.has(note.id)
// Dynamic font size based on row height (base: 12px at 20px row height)
const fontSize = Math.max(10, Math.min(24, rowHeight * 0.6))
return (
<NoteChip
key={note.id}
note={note}
left={left}
top={top}
width={noteWidthPx}
height={noteHeight}
fontSize={fontSize}
isSelected={selectedId === note.id}
isOverlapping={isOverlapping}
onPointerDown={(event, mode) => startDrag(event, note, mode)}
onDoubleClick={(event) => {
event.stopPropagation()
onFocusLyric?.(note.id)
}}
/>
)
})}
</div>
</div>
</div>
</div>
)
}
@@ -0,0 +1,8 @@
// Base values used for scaling; actual runtime values are derived in components
export const BASE_GRID_SECOND_WIDTH = 80
export const BASE_ROW_HEIGHT = 20
export const PITCH_WIDTH = 60
// C-1 to C8 range (MIDI note numbers)
// LOW_NOTE = 0 to support SP markers (pitch=0) in some MIDI files
export const LOW_NOTE = 0 // C-1 (also supports pitch=0 for SP markers)
export const HIGH_NOTE = 108 // C8
@@ -0,0 +1,37 @@
@tailwind base;
@tailwind components;
@tailwind utilities;
:root {
font-family: 'Space Grotesk', 'IBM Plex Sans', system-ui, sans-serif;
color: var(--text-primary);
background: radial-gradient(circle at 20% 20%, rgba(72, 228, 194, 0.08), transparent 35%),
radial-gradient(circle at 80% 0%, rgba(75, 100, 188, 0.24), transparent 40%),
#0f1528;
text-rendering: optimizeLegibility;
-webkit-font-smoothing: antialiased;
}
:root[data-theme='light'] {
background: radial-gradient(circle at 20% 20%, rgba(63, 140, 255, 0.08), transparent 35%),
radial-gradient(circle at 80% 0%, rgba(75, 100, 188, 0.14), transparent 40%),
#f5f7fb;
}
* {
box-sizing: border-box;
}
body {
margin: 0;
min-height: 100vh;
background: transparent;
}
#root {
min-height: 100vh;
}
a {
color: inherit;
}
@@ -0,0 +1,224 @@
import { Midi } from '@tonejs/midi'
import { writeMidi } from 'midi-file'
import type { MidiData, MidiEvent } from 'midi-file'
import type { NoteEvent, ProjectSnapshot, TimeSignature } from '../types'
const DEFAULT_SIGNATURE: TimeSignature = [4, 4]
// Decode UTF-8 byte string (latin1 encoded) to proper Unicode string
// This matches: text.encode("latin1").decode("utf-8") in Python
function decodeUtf8ByteString(byteString: string): string {
try {
const bytes = new Uint8Array(byteString.length)
for (let i = 0; i < byteString.length; i++) {
bytes[i] = byteString.charCodeAt(i)
}
return new TextDecoder('utf-8').decode(bytes)
} catch {
return byteString
}
}
// Encode Unicode string to UTF-8 byte string (latin1 encoding)
// This matches: text.encode("utf-8").decode("latin1") in Python
function encodeUtf8ByteString(text: string): string {
const bytes = new TextEncoder().encode(text)
let output = ''
bytes.forEach((b) => {
output += String.fromCharCode(b)
})
return output
}
export async function importMidiFile(file: File): Promise<ProjectSnapshot> {
const buffer = await file.arrayBuffer()
return parseMidiBuffer(buffer)
}
export async function parseMidiBuffer(buffer: ArrayBuffer): Promise<ProjectSnapshot> {
const midi = new Midi(buffer)
const tempo = midi.header.tempos[0]?.bpm ?? 120
const timeSignature = (midi.header.timeSignatures[0]?.timeSignature as TimeSignature | undefined) ?? DEFAULT_SIGNATURE
// Merge notes from all tracks and sort by ticks then by midi (for stable ordering)
const allNotes = midi.tracks
.flatMap(t => t.notes)
.sort((a, b) => a.ticks - b.ticks || a.midi - b.midi)
// Get lyrics from header.meta and sort by ticks
const lyricEvents = midi.header.meta
.filter((event) => event.type === 'lyrics')
.sort((a, b) => a.ticks - b.ticks)
// Match lyrics to notes by tick position
// Each lyric should be consumed by exactly one note at the same tick
const lyricsByTick = new Map<number, string[]>()
for (const event of lyricEvents) {
const existing = lyricsByTick.get(event.ticks) || []
existing.push(decodeUtf8ByteString(event.text))
lyricsByTick.set(event.ticks, existing)
}
// Track which lyrics have been used at each tick position
const usedLyricIndices = new Map<number, number>()
const notes: NoteEvent[] = allNotes.map((note, index) => {
const beat = note.ticks / midi.header.ppq
const durationBeats = note.durationTicks / midi.header.ppq
let lyric = ''
// First try exact tick match
const lyricsAtTick = lyricsByTick.get(note.ticks)
if (lyricsAtTick && lyricsAtTick.length > 0) {
const usedIndex = usedLyricIndices.get(note.ticks) || 0
if (usedIndex < lyricsAtTick.length) {
lyric = lyricsAtTick[usedIndex]
usedLyricIndices.set(note.ticks, usedIndex + 1)
}
}
// If no exact match, try nearby ticks (within small tolerance)
if (!lyric) {
const tolerance = midi.header.ppq / 100 // Very small tolerance
for (const [tick, lyrics] of lyricsByTick.entries()) {
if (Math.abs(tick - note.ticks) <= tolerance) {
const usedIndex = usedLyricIndices.get(tick) || 0
if (usedIndex < lyrics.length) {
lyric = lyrics[usedIndex]
usedLyricIndices.set(tick, usedIndex + 1)
break
}
}
}
}
return {
id: `${index}-${note.midi}-${Math.round(note.ticks)}`,
midi: note.midi,
start: beat,
duration: Math.max(durationBeats, 0.0625),
velocity: note.velocity,
lyric,
}
})
return { tempo, timeSignature, notes, ppq: midi.header.ppq }
}
// Used to add absoluteTime property for sorting
type WithAbsoluteTime<T> = T & { absoluteTime: number }
export function exportMidi(snapshot: ProjectSnapshot): Blob {
const ppq = snapshot.ppq ?? 480 // Use original ppq if available, otherwise default to 480
const microsecondsPerBeat = Math.round(60000000 / snapshot.tempo) // Convert BPM to microseconds per beat
// Sort notes by start time, then by midi for stable ordering
const sortedNotes = [...snapshot.notes].sort((a, b) => a.start - b.start || a.midi - b.midi)
// Build events for a single track containing both lyrics and notes
// Event order at same tick: note_off (0) < lyrics (1) < note_on (2)
// This matches meta.py's tg2midi implementation
const events: Array<WithAbsoluteTime<MidiEvent>> = []
// Add all note events and their corresponding lyrics
sortedNotes.forEach((note) => {
const startTicks = Math.round(note.start * ppq)
const endTicks = Math.round((note.start + note.duration) * ppq)
const velocity = Math.round(note.velocity * 127)
// Add lyric event at the same tick as note_on (but will be sorted before it)
const lyricText = note.lyric ?? ''
const encodedLyric = encodeUtf8ByteString(lyricText)
// Lyric event - sort key 1 (after note_off, before note_on)
events.push({
absoluteTime: startTicks,
deltaTime: 0,
meta: true,
type: 'lyrics',
text: encodedLyric,
_sortKey: 1,
} as WithAbsoluteTime<MidiEvent> & { _sortKey: number })
// Note on event - sort key 2 (after lyrics)
events.push({
absoluteTime: startTicks,
deltaTime: 0,
type: 'noteOn',
channel: 0,
noteNumber: note.midi,
velocity: velocity,
_sortKey: 2,
} as WithAbsoluteTime<MidiEvent> & { _sortKey: number })
// Note off event - sort key 0 (before everything at same tick)
events.push({
absoluteTime: endTicks,
deltaTime: 0,
type: 'noteOff',
channel: 0,
noteNumber: note.midi,
velocity: 0,
_sortKey: 0,
} as WithAbsoluteTime<MidiEvent> & { _sortKey: number })
})
// Sort events by absoluteTime, then by _sortKey
events.sort((a, b) => {
const aKey = (a as { _sortKey?: number })._sortKey ?? 1
const bKey = (b as { _sortKey?: number })._sortKey ?? 1
return a.absoluteTime - b.absoluteTime || aKey - bKey
})
// Convert absolute time to delta time
let lastTick = 0
events.forEach(event => {
event.deltaTime = event.absoluteTime - lastTick
lastTick = event.absoluteTime
delete (event as { absoluteTime?: number }).absoluteTime
delete (event as { _sortKey?: number })._sortKey
})
// Build the MIDI track with header events
const track: MidiEvent[] = [
// Set tempo
{
deltaTime: 0,
meta: true,
type: 'setTempo',
microsecondsPerBeat: microsecondsPerBeat,
},
// Time signature
{
deltaTime: 0,
meta: true,
type: 'timeSignature',
numerator: snapshot.timeSignature[0],
denominator: snapshot.timeSignature[1],
metronome: 24,
thirtyseconds: 8,
},
// All note and lyric events
...events,
// End of track
{
deltaTime: 0,
meta: true,
type: 'endOfTrack',
},
]
// Build MIDI data structure
const midiData: MidiData = {
header: {
format: 0, // Single track format (type 0)
numTracks: 1,
ticksPerBeat: ppq,
},
tracks: [track],
}
const bytes = writeMidi(midiData)
return new Blob([new Uint8Array(bytes)], { type: 'audio/midi' })
}
+10
View File
@@ -0,0 +1,10 @@
import { StrictMode } from 'react'
import { createRoot } from 'react-dom/client'
import './index.css'
import App from './App.tsx'
createRoot(document.getElementById('root')!).render(
<StrictMode>
<App />
</StrictMode>,
)
@@ -0,0 +1,78 @@
import { nanoid } from 'nanoid'
import { create } from 'zustand'
import type { NoteEvent, TimeSignature } from '../types'
const clamp = (value: number, min: number, max: number) =>
Math.min(Math.max(value, min), max)
export type MidiStore = {
tempo: number
timeSignature: TimeSignature
notes: NoteEvent[]
selectedId: string | null
playhead: number
ppq: number | undefined // Ticks per quarter note (for preserving original MIDI timing)
addNote: (partial?: Partial<NoteEvent>) => NoteEvent
updateNote: (id: string, partial: Partial<NoteEvent>) => void
removeNote: (id: string) => void
setNotes: (notes: NoteEvent[]) => void
setTempo: (tempo: number) => void
setTimeSignature: (sig: TimeSignature) => void
setPpq: (ppq: number | undefined) => void
select: (id: string | null) => void
setLyric: (id: string, lyric: string) => void
setPlayhead: (beat: number) => void
clear: () => void
}
const defaultNotes: NoteEvent[] = [
{ id: nanoid(), midi: 64, start: 0, duration: 1.5, velocity: 0.9, lyric: 'la' },
{ id: nanoid(), midi: 67, start: 1.5, duration: 1.5, velocity: 0.85, lyric: 'na' },
{ id: nanoid(), midi: 69, start: 3, duration: 2, velocity: 0.8, lyric: 'ah' },
]
export const useMidiStore = create<MidiStore>((set) => ({
tempo: 110,
timeSignature: [4, 4],
notes: defaultNotes,
selectedId: null,
playhead: 0,
ppq: undefined,
addNote: (partial = {}) => {
const note: NoteEvent = {
id: nanoid(),
midi: partial.midi ?? 64,
start: partial.start ?? 0,
duration: partial.duration ?? 1,
velocity: clamp(partial.velocity ?? 0.85, 0, 1),
lyric: partial.lyric ?? '',
}
set((state) => ({ notes: [...state.notes, note] }))
return note
},
updateNote: (id, partial) => {
set((state) => ({
notes: state.notes.map((note) =>
note.id === id
? {
...note,
...partial,
duration: Math.max(partial.duration ?? note.duration, 0.0625),
}
: note,
),
}))
},
removeNote: (id) => set((state) => ({ notes: state.notes.filter((n) => n.id !== id) })),
setNotes: (notes) => set(() => ({ notes })),
setTempo: (tempo) => set(() => ({ tempo: clamp(tempo, 30, 240) })),
setTimeSignature: (sig) => set(() => ({ timeSignature: sig })),
setPpq: (ppq) => set(() => ({ ppq })),
select: (id) => set(() => ({ selectedId: id })),
setLyric: (id, lyric) =>
set((state) => ({
notes: state.notes.map((note) => (note.id === id ? { ...note, lyric } : note)),
})),
setPlayhead: (beat) => set(() => ({ playhead: Math.max(beat, 0) })),
clear: () => set(() => ({ notes: [], selectedId: null })),
}))
+17
View File
@@ -0,0 +1,17 @@
export type NoteEvent = {
id: string
midi: number
start: number // in beats
duration: number // in beats
velocity: number
lyric: string
}
export type TimeSignature = [number, number]
export type ProjectSnapshot = {
tempo: number
timeSignature: TimeSignature
notes: NoteEvent[]
ppq?: number // Ticks per quarter note (for preserving original MIDI timing)
}
@@ -0,0 +1,33 @@
/** @type {import('tailwindcss').Config} */
export default {
content: ['./index.html', './src/**/*.{ts,tsx,js,jsx}'],
theme: {
extend: {
fontFamily: {
display: ['"Space Grotesk"', '"IBM Plex Sans"', 'system-ui', 'sans-serif'],
mono: ['"JetBrains Mono"', 'ui-monospace', 'SFMono-Regular', 'monospace'],
},
colors: {
ink: {
50: '#f4f7fb',
100: '#dfe7f5',
200: '#beceec',
300: '#95addf',
400: '#6a87ce',
500: '#4b64bc',
600: '#3b4ea7',
700: '#32418a',
800: '#2c376f',
900: '#262f5c',
},
ember: '#ff7043',
mint: '#48e4c2',
},
boxShadow: {
panel: '0 14px 35px rgba(0, 0, 0, 0.25)',
},
},
},
plugins: [],
}
@@ -0,0 +1,28 @@
{
"compilerOptions": {
"tsBuildInfoFile": "./node_modules/.tmp/tsconfig.app.tsbuildinfo",
"target": "ES2022",
"useDefineForClassFields": true,
"lib": ["ES2022", "DOM", "DOM.Iterable"],
"module": "ESNext",
"types": ["vite/client"],
"skipLibCheck": true,
/* Bundler mode */
"moduleResolution": "bundler",
"allowImportingTsExtensions": true,
"verbatimModuleSyntax": true,
"moduleDetection": "force",
"noEmit": true,
"jsx": "react-jsx",
/* Linting */
"strict": true,
"noUnusedLocals": true,
"noUnusedParameters": true,
"erasableSyntaxOnly": true,
"noFallthroughCasesInSwitch": true,
"noUncheckedSideEffectImports": true
},
"include": ["src"]
}
@@ -0,0 +1,7 @@
{
"files": [],
"references": [
{ "path": "./tsconfig.app.json" },
{ "path": "./tsconfig.node.json" }
]
}
@@ -0,0 +1,26 @@
{
"compilerOptions": {
"tsBuildInfoFile": "./node_modules/.tmp/tsconfig.node.tsbuildinfo",
"target": "ES2023",
"lib": ["ES2023"],
"module": "ESNext",
"types": ["node"],
"skipLibCheck": true,
/* Bundler mode */
"moduleResolution": "bundler",
"allowImportingTsExtensions": true,
"verbatimModuleSyntax": true,
"moduleDetection": "force",
"noEmit": true,
/* Linting */
"strict": true,
"noUnusedLocals": true,
"noUnusedParameters": true,
"erasableSyntaxOnly": true,
"noFallthroughCasesInSwitch": true,
"noUncheckedSideEffectImports": true
},
"include": ["vite.config.ts"]
}
@@ -0,0 +1,7 @@
import { defineConfig } from 'vite'
import react from '@vitejs/plugin-react'
// https://vite.dev/config/
export default defineConfig({
plugins: [react()],
})
+669
View File
@@ -0,0 +1,669 @@
"""
SoulX-Singer MIDI <-> metadata converter.
Converts between SoulX-Singer-style metadata JSON (with note_text, note_dur,
note_pitch, note_type per segment) and standard MIDI files. Uses an internal
Note dataclass (start_s, note_dur, note_text, note_pitch, note_type) as the
intermediate representation.
"""
import os
import json
import shutil
from dataclasses import dataclass
from typing import Any, List, Tuple, Union
import librosa
import mido
from soundfile import write
from .f0_extraction import F0Extractor
from .g2p import g2p_transform
# Audio and segmenting constants (used by _edit_data_to_meta)
SAMPLE_RATE = 44100
DEFAULT_LANGUAGE = "Mandarin"
MAX_GAP_SEC = 5.0 # gap (sec) above which we start a new segment
MAX_SEGMENT_DUR_SUM_SEC = 60.0 # max cumulative note duration per segment (sec)
MIN_GAP_THRESHOLD_SEC = 0.001 # ignore gaps smaller than this
LONG_SILENCE_THRESHOLD_SEC = 0.05 # treat as separate <SP> if gap larger
MAX_LEADING_SP_DUR_SEC = 2.0 # cap leading silence in a segment to this (sec)
DEFAULT_RMVPE_MODEL_PATH = "pretrained_models/SoulX-Singer-Preprocess/rmvpe/rmvpe.pt"
@dataclass
class Note:
"""Single note: text, duration (seconds), pitch (MIDI), type. start_s is absolute start time in seconds (for ordering / MIDI)."""
start_s: float
note_dur: float
note_text: str
note_pitch: int
note_type: int
@property
def end_s(self) -> float:
return self.start_s + self.note_dur
def remove_duplicate_segments(meta_data: List[dict]) -> None:
"""Merge consecutive identical notes (same text, pitch, type) within each segment. Mutates meta_data in place."""
for idx, segment in enumerate(meta_data):
texts = segment["note_text"]
durs = segment["note_dur"]
pitches = segment["note_pitch"]
types = segment["note_type"]
new_texts = []
new_durs = []
new_pitches = []
new_types = []
for i in range(len(texts)):
if i == 0:
new_texts.append(texts[i])
new_durs.append(durs[i])
new_pitches.append(pitches[i])
new_types.append(types[i])
continue
t, d, p, ty = texts[i], durs[i], pitches[i], types[i]
if t == "<SP>" and texts[i - 1] == "<SP>":
new_durs[-1] += d
continue
if t == texts[i - 1] and p == pitches[i - 1] and ty == types[i - 1]:
new_durs[-1] += d
else:
new_texts.append(t)
new_durs.append(d)
new_pitches.append(p)
new_types.append(ty)
meta_data[idx]["note_text"] = new_texts
meta_data[idx]["note_dur"] = new_durs
meta_data[idx]["note_pitch"] = new_pitches
meta_data[idx]["note_type"] = new_types
def meta2notes(meta_path: str) -> List[Note]:
"""Parse SoulX-Singer metadata JSON into a flat list of Note (absolute start_s)."""
with open(meta_path, "r", encoding="utf-8") as f:
segments = json.load(f)
if not isinstance(segments, list):
raise ValueError(f"Metadata must be a list of segments, got {type(segments).__name__}")
if not segments:
raise ValueError("Metadata has no segments.")
notes: List[Note] = []
for seg in segments:
offset_s = seg["time"][0] / 1000
words = [str(x).replace("<AP>", "<SP>") for i, x in enumerate(seg["text"].split())]
word_durs = [float(x) for x in seg["duration"].split()]
pitches = [int(x) for x in seg["note_pitch"].split()]
types = [int(x) if words[i] != "<SP>" else 1 for i, x in enumerate(seg["note_type"].split())]
if len(words) != len(word_durs) or len(word_durs) != len(pitches) or len(pitches) != len(types):
raise ValueError(
f"Length mismatch in segment {seg.get('item_name', '?')}: "
"note_text, note_dur, note_pitch, note_type must have same length"
)
current_s = offset_s
for text, dur, pitch, type_ in zip(words, word_durs, pitches, types):
notes.append(
Note(
start_s=current_s,
note_dur=float(dur),
note_text=str(text),
note_pitch=int(pitch),
note_type=int(type_),
)
)
current_s += float(dur)
return notes
def _append_segment_to_meta(
meta_path_str: str,
cut_wavs_output_dir: str,
vocal_file: str,
audio_data: Any,
meta_data: List[dict],
note_start: List[float],
note_end: List[float],
note_text: List[Any],
note_pitch: List[Any],
note_type: List[Any],
note_dur: List[float],
end_time_ms_override: float | None = None,
) -> None:
"""Write one segment wav and append one segment dict to meta_data. Caller clears note_* lists after."""
base_name = os.path.splitext(os.path.basename(meta_path_str))[0]
item_name = f"{base_name}_{len(meta_data)}"
wav_fn = os.path.join(cut_wavs_output_dir, f"{item_name}.wav")
start_ms = int(note_start[0] * 1000)
end_ms = (
int(end_time_ms_override)
if end_time_ms_override is not None
else int(note_end[-1] * 1000)
)
start_sample = int(note_start[0] * SAMPLE_RATE)
end_sample = int(note_end[-1] * SAMPLE_RATE)
write(wav_fn, audio_data[start_sample:end_sample], SAMPLE_RATE)
meta_data.append({
"item_name": item_name,
"wav_fn": wav_fn,
"origin_wav_fn": vocal_file,
"start_time_ms": start_ms,
"end_time_ms": end_ms,
"language": DEFAULT_LANGUAGE,
"note_text": list(note_text),
"note_pitch": list(note_pitch),
"note_type": list(note_type),
"note_dur": list(note_dur),
})
def convert_meta(meta_data: List[dict], rmvpe_model_path, device="cuda"):
pitch_extractor = F0Extractor(rmvpe_model_path, device=device, verbose=False)
converted_data = []
for item in meta_data:
wav_fn = item.get("wav_fn")
if not wav_fn or not os.path.isfile(wav_fn):
raise FileNotFoundError(f"Segment wav file not found: {wav_fn}")
f0 = pitch_extractor.process(wav_fn)
converted_item = {
"index": item.get("item_name"),
"language": item.get("language"),
"time": [item.get("start_time_ms", 0), item.get("end_time_ms", sum(item["note_dur"]) * 1000)],
"duration": " ".join(str(round(x, 2)) for x in item.get("note_dur", [])),
"text": " ".join(item.get("note_text", [])),
"phoneme": " ".join(g2p_transform(item.get("note_text", []), DEFAULT_LANGUAGE)),
"note_pitch": " ".join(str(x) for x in item.get("note_pitch", [])),
"note_type": " ".join(str(x) for x in item.get("note_type", [])),
"f0": " ".join(str(round(float(x), 1)) for x in f0),
}
converted_data.append(converted_item)
return converted_data
def _edit_data_to_meta(
meta_path_str: str,
edit_data: List[dict],
vocal_file: str,
rmvpe_model_path: str | None = None,
device: str = "cuda",
) -> None:
"""Write SoulX-Singer metadata JSON from edit_data (list of {start, end, note_text, note_pitch, note_type})."""
# Use a fixed temporary directory for cut wavs
cut_wavs_output_dir = os.path.join(os.path.dirname(vocal_file), "cut_wavs_tmp")
os.makedirs(cut_wavs_output_dir, exist_ok=True)
note_text: List[Any] = []
note_pitch: List[Any] = []
note_type: List[Any] = []
note_dur: List[float] = []
note_start: List[float] = []
note_end: List[float] = []
prev_end = 0.0
meta_data: List[dict] = []
audio_data, _ = librosa.load(vocal_file, sr=SAMPLE_RATE, mono=True)
dur_sum = 0.0
for entry in edit_data:
start = float(entry["start"])
end = float(entry["end"])
text = entry["note_text"]
pitch = entry["note_pitch"]
type_ = entry["note_type"]
if text == "" or pitch == "" or type_ == "":
note_text.append("<SP>")
note_pitch.append(0)
note_type.append(1)
note_dur.append(end - start)
note_start.append(start)
note_end.append(end)
prev_end = end
dur_sum += end - start
continue
if (
len(note_text) > 0
and note_text[-1] == "<SP>"
and note_dur[-1] > MAX_LEADING_SP_DUR_SEC
):
cut_time = note_dur[-1] - MAX_LEADING_SP_DUR_SEC
note_dur[-1] = MAX_LEADING_SP_DUR_SEC
end_ms_override = note_end[-1] * 1000 - cut_time * 1000
_append_segment_to_meta(
meta_path_str,
cut_wavs_output_dir,
vocal_file,
audio_data,
meta_data,
note_start,
note_end,
note_text,
note_pitch,
note_type,
note_dur,
end_time_ms_override=end_ms_override,
)
note_text = []
note_pitch = []
note_type = []
note_dur = []
note_start = []
note_end = []
prev_end = start
dur_sum = 0.0
gap_from_prev = start - prev_end
gap_from_last_note = (start - note_end[-1]) if note_end else 0.0
if (
gap_from_prev >= MAX_GAP_SEC
or gap_from_last_note >= MAX_GAP_SEC
or dur_sum >= MAX_SEGMENT_DUR_SUM_SEC
):
if len(note_text) > 0:
_append_segment_to_meta(
meta_path_str,
cut_wavs_output_dir,
vocal_file,
audio_data,
meta_data,
note_start,
note_end,
note_text,
note_pitch,
note_type,
note_dur,
)
note_text = []
note_pitch = []
note_type = []
note_dur = []
note_start = []
note_end = []
prev_end = start
dur_sum = 0.0
if start - prev_end > MIN_GAP_THRESHOLD_SEC:
if start - prev_end > LONG_SILENCE_THRESHOLD_SEC or len(note_text) == 0:
note_text.append("<SP>")
note_pitch.append(0)
note_type.append(1)
note_dur.append(start - prev_end)
note_start.append(prev_end)
note_end.append(start)
else:
if len(note_dur) > 0:
note_dur[-1] += start - prev_end
note_end[-1] = start
prev_end = end
note_text.append(text)
note_pitch.append(int(pitch))
note_type.append(int(type_))
note_dur.append(end - start)
note_start.append(start)
note_end.append(end)
dur_sum += end - start
if len(note_text) > 0:
_append_segment_to_meta(
meta_path_str,
cut_wavs_output_dir,
vocal_file,
audio_data,
meta_data,
note_start,
note_end,
note_text,
note_pitch,
note_type,
note_dur,
)
remove_duplicate_segments(meta_data)
_rmvpe_path = rmvpe_model_path or DEFAULT_RMVPE_MODEL_PATH
converted_data = convert_meta(meta_data, _rmvpe_path, device)
with open(meta_path_str, "w", encoding="utf-8") as f:
json.dump(converted_data, f, ensure_ascii=False, indent=2)
# Clean up temporary cut wavs directory
try:
shutil.rmtree(cut_wavs_output_dir, ignore_errors=True)
except Exception:
pass
def notes2meta(
notes: List[Note],
meta_path: str,
vocal_file: str,
rmvpe_model_path: str | None = None,
device: str = "cuda",
) -> None:
"""Write SoulX-Singer metadata JSON from a list of Note (segmenting + wav cuts)."""
edit_data = [
{
"start": n.start_s,
"end": n.end_s,
"note_text": n.note_text,
"note_pitch": str(n.note_pitch),
"note_type": str(n.note_type),
}
for n in notes
]
_edit_data_to_meta(
str(meta_path),
edit_data,
vocal_file,
rmvpe_model_path=rmvpe_model_path,
device=device,
)
@dataclass(frozen=True)
class MidiDefaults:
ticks_per_beat: int = 500
tempo: int = 500000 # microseconds per beat (120 BPM)
time_signature: Tuple[int, int] = (4, 4)
velocity: int = 64
def _seconds_to_ticks(seconds: float, ticks_per_beat: int, tempo: int) -> int:
return int(round(seconds * ticks_per_beat * 1_000_000 / tempo))
def notes2midi(
notes: List[Note],
midi_path: str,
defaults: MidiDefaults | None = None,
) -> None:
"""Write MIDI file from a list of Note."""
defaults = defaults or MidiDefaults()
if not notes:
raise ValueError("Empty note list.")
events: List[Tuple[int, int, Union[mido.Message, mido.MetaMessage]]] = []
for n in notes:
start_s = n.start_s
end_s = n.end_s
if end_s <= start_s:
continue
start_ticks = _seconds_to_ticks(
start_s, defaults.ticks_per_beat, defaults.tempo
)
end_ticks = _seconds_to_ticks(
end_s, defaults.ticks_per_beat, defaults.tempo
)
if end_ticks <= start_ticks:
end_ticks = start_ticks + 1
lyric = n.note_text
try:
lyric = lyric.encode("utf-8").decode("latin1")
except (UnicodeEncodeError, UnicodeDecodeError):
pass
if n.note_type == 3:
lyric = "-"
events.append(
(start_ticks, 1, mido.MetaMessage("lyrics", text=lyric, time=0))
)
events.append(
(
start_ticks,
2,
mido.Message(
"note_on",
note=n.note_pitch,
velocity=defaults.velocity,
time=0,
),
)
)
events.append(
(
end_ticks,
0,
mido.Message("note_off", note=n.note_pitch, velocity=0, time=0),
)
)
events.sort(key=lambda x: (x[0], x[1]))
mid = mido.MidiFile(ticks_per_beat=defaults.ticks_per_beat)
track = mido.MidiTrack()
mid.tracks.append(track)
track.append(mido.MetaMessage("set_tempo", tempo=defaults.tempo, time=0))
track.append(
mido.MetaMessage(
"time_signature",
numerator=defaults.time_signature[0],
denominator=defaults.time_signature[1],
time=0,
)
)
last_tick = 0
for tick, _, msg in events:
msg.time = max(0, tick - last_tick)
track.append(msg)
last_tick = tick
track.append(mido.MetaMessage("end_of_track", time=0))
mid.save(midi_path)
def midi2notes(midi_path: str) -> List[Note]:
"""Parse MIDI file into a list of Note. Merges all tracks; tempo from last set_tempo event."""
mid = mido.MidiFile(midi_path)
ticks_per_beat = mid.ticks_per_beat
tempo = 500000
raw_notes: List[dict] = []
lyrics: List[Tuple[int, str]] = []
for track in mid.tracks:
abs_ticks = 0
active = {}
for msg in track:
abs_ticks += msg.time
if msg.type == "set_tempo":
tempo = msg.tempo
elif msg.type == "lyrics":
text = msg.text
try:
text = text.encode("latin1").decode("utf-8")
except Exception:
pass
lyrics.append((abs_ticks, text))
elif msg.type == "note_on":
key = (msg.channel, msg.note)
if msg.velocity > 0:
active[key] = (abs_ticks, msg.velocity)
else:
if key in active:
start_ticks, vel = active.pop(key)
raw_notes.append(
{
"midi": msg.note,
"start_ticks": start_ticks,
"duration_ticks": abs_ticks - start_ticks,
"velocity": vel,
"lyric": "",
}
)
elif msg.type == "note_off":
key = (msg.channel, msg.note)
if key in active:
start_ticks, vel = active.pop(key)
raw_notes.append(
{
"midi": msg.note,
"start_ticks": start_ticks,
"duration_ticks": abs_ticks - start_ticks,
"velocity": vel,
"lyric": "",
}
)
if not raw_notes:
raise ValueError("No notes found in MIDI file")
for n in raw_notes:
n["end_ticks"] = n["start_ticks"] + n["duration_ticks"]
raw_notes.sort(key=lambda n: n["start_ticks"])
lyrics.sort(key=lambda x: x[0])
trimmed = []
for note in raw_notes:
while trimmed:
prev = trimmed[-1]
if note["start_ticks"] < prev["end_ticks"]:
prev["end_ticks"] = note["start_ticks"]
prev["duration_ticks"] = prev["end_ticks"] - prev["start_ticks"]
if prev["duration_ticks"] <= 0:
trimmed.pop()
continue
break
trimmed.append(note)
raw_notes = trimmed
tolerance = ticks_per_beat // 100
lyric_idx = 0
for note in raw_notes:
while lyric_idx < len(lyrics) and lyrics[lyric_idx][0] < note["start_ticks"] - tolerance:
lyric_idx += 1
if lyric_idx < len(lyrics):
lyric_ticks, lyric_text = lyrics[lyric_idx]
if abs(lyric_ticks - note["start_ticks"]) <= tolerance:
note["lyric"] = lyric_text
lyric_idx += 1
def ticks_to_seconds(ticks: int) -> float:
return (ticks / ticks_per_beat) * (tempo / 1_000_000)
result: List[Note] = []
prev_end_s = 0.0
for idx, n in enumerate(raw_notes):
start_s = ticks_to_seconds(n["start_ticks"])
end_s = ticks_to_seconds(n["end_ticks"])
if prev_end_s > start_s:
start_s = prev_end_s
dur_s = end_s - start_s
if dur_s <= 0:
continue
lyric = n.get("lyric", "")
if not lyric:
tp = 2
text = "啦"
elif lyric == "<SP>":
tp = 1
text = "<SP>"
elif lyric == "-":
tp = 3
text = raw_notes[idx - 1].get("lyric", "-") if idx > 0 else "-"
else:
tp = 2
text = lyric
result.append(
Note(
start_s=start_s,
note_dur=dur_s,
note_text=text,
note_pitch=n["midi"],
note_type=tp,
)
)
prev_end_s = end_s
return result
def meta2midi(meta_path: str, midi_path: str, defaults: MidiDefaults | None = None) -> None:
"""Convert SoulX-Singer metadata JSON to MIDI file (meta -> List[Note] -> midi)."""
notes = meta2notes(meta_path)
notes2midi(notes, midi_path, defaults)
print(f"Saved MIDI to {midi_path}")
def midi2meta(
midi_path: str,
meta_path: str,
vocal_file: str,
rmvpe_model_path: str | None = None,
device: str = "cuda",
) -> None:
"""Convert MIDI file to SoulX-Singer metadata JSON (midi -> List[Note] -> meta)."""
meta_dir = os.path.dirname(meta_path)
if meta_dir:
os.makedirs(meta_dir, exist_ok=True)
# cut_wavs will be written to a fixed temporary directory inside _edit_data_to_meta
notes = midi2notes(midi_path)
notes2meta(
notes,
meta_path,
vocal_file,
rmvpe_model_path=rmvpe_model_path,
device=device,
)
print(f"Saved Meta to {meta_path}")
if __name__ == "__main__":
import argparse
parser = argparse.ArgumentParser(
description="Convert SoulX-Singer metadata JSON <-> MIDI."
)
parser.add_argument("--meta", type=str, help="Path to metadata JSON")
parser.add_argument("--midi", type=str, help="Path to MIDI file")
parser.add_argument("--vocal", type=str, help="Path to vocal wav (for midi2meta)")
parser.add_argument(
"--meta2midi",
action="store_true",
help="Convert meta -> midi (requires --meta and --midi)",
)
parser.add_argument(
"--midi2meta",
action="store_true",
help="Convert midi -> meta (requires --midi, --meta, --vocal, --cut_wavs_dir)",
)
parser.add_argument(
"--rmvpe_model_path",
type=str,
help="Path to RMVPE model",
default="pretrained_models/SoulX-Singer-Preprocess/rmvpe/rmvpe.pt",
)
parser.add_argument(
"--device",
type=str,
help="Device to use for RMVPE",
default="cuda",
)
args = parser.parse_args()
if args.meta2midi:
if not args.meta or not args.midi:
parser.error("--meta2midi requires --meta and --midi")
meta2midi(args.meta, args.midi)
elif args.midi2meta:
if not args.midi or not args.meta or not args.vocal:
parser.error(
"--midi2meta requires --midi, --meta, --vocal"
)
midi2meta(
args.midi,
args.meta,
args.vocal,
rmvpe_model_path=args.rmvpe_model_path,
device=args.device,
)
else:
parser.print_help()
@@ -0,0 +1,522 @@
# https://github.com/RickyL-2000/ROSVOT
import math
import sys
import traceback
import json
import time
from pathlib import Path
from typing import Any, Dict, Optional
import librosa
import numpy as np
import torch
import matplotlib.pyplot as plt
from .utils.os_utils import safe_path
from .utils.commons.hparams import set_hparams
from .utils.commons.ckpt_utils import load_ckpt
from .utils.commons.dataset_utils import pad_or_cut_xd
from .utils.audio.mel import MelNet
from .utils.audio.pitch_utils import (
norm_interp_f0,
denorm_f0,
f0_to_coarse,
boundary2Interval,
save_midi,
midi_to_hz,
)
from .utils.rosvot_utils import (
get_mel_len,
align_word,
regulate_real_note_itv,
regulate_ill_slur,
bd_to_durs,
)
from .modules.pe.rmvpe import RMVPE
from .modules.rosvot.rosvot import MidiExtractor, WordbdExtractor
@torch.no_grad()
def infer_sample(
item: Dict[str, Any],
hparams: Dict[str, Any],
models: Dict[str, Any],
device: torch.device,
*,
save_dir: Optional[str] = None,
apply_rwbd: Optional[bool] = None,
# outputs
save_plot: bool = False,
no_save_midi: bool = True,
no_save_npy: bool = True,
verbose: bool = False,
) -> Dict[str, Any]:
if "item_name" not in item or "wav_fn" not in item:
raise ValueError('item must contain keys: "item_name" and "wav_fn"')
item_name = item["item_name"]
wav_src = item["wav_fn"]
# Decide RWBD usage
if apply_rwbd is None:
apply_rwbd_ = ("word_durs" not in item)
else:
apply_rwbd_ = bool(apply_rwbd)
# Models
model = models["model"]
mel_net = models["mel_net"]
pe = models.get("pe")
wbd_predictor = models.get("wbd_predictor")
if wbd_predictor is None and apply_rwbd_:
raise ValueError("apply_rwbd is True but wbd_predictor model is not provided in models")
# ---- Prepare Data ----
if isinstance(wav_src, str):
wav, _ = librosa.core.load(wav_src, sr=hparams["audio_sample_rate"])
else:
wav = wav_src
if not isinstance(wav, np.ndarray):
wav = np.asarray(wav)
wav = wav.astype(np.float32)
# Calculate timestamps and alignment lengths
wav_len_samples = wav.shape[-1]
mel_len = get_mel_len(wav_len_samples, hparams["hop_size"])
# Word boundary preparation
mel2word = None
word_durs_filtered = None
if not apply_rwbd_:
if "word_durs" not in item:
raise ValueError('apply_rwbd=False but item has no "word_durs"')
wd_raw = list(item["word_durs"])
min_word_dur = hparams.get("min_word_dur", 20) / 1000
word_durs_filtered = []
for i, wd in enumerate(wd_raw):
if wd < min_word_dur:
if i == 0 and len(wd_raw) > 1:
wd_raw[i + 1] += wd
elif len(word_durs_filtered) > 0:
word_durs_filtered[-1] += wd
else:
word_durs_filtered.append(wd)
mel2word, _ = align_word(word_durs_filtered, mel_len, hparams["hop_size"], hparams["audio_sample_rate"])
mel2word = np.asarray(mel2word)
if mel2word.size > 0 and mel2word[0] == 0:
mel2word = mel2word + 1
mel2word_len = int(np.sum(mel2word > 0))
real_len = min(mel_len, mel2word_len)
else:
real_len = min(mel_len, hparams["max_frames"])
T = math.ceil(min(real_len, hparams["max_frames"]) / hparams["frames_multiple"]) * hparams["frames_multiple"]
# ---- Input Tensors & Padding ----
target_samples = T * hparams["hop_size"]
wav_t = torch.from_numpy(wav).float().to(device).unsqueeze(0) # [1, L]
if wav_t.shape[-1] < target_samples:
wav_t = pad_or_cut_xd(wav_t, target_samples, 1)
# ---- Pitch Extraction ----
if pe is not None:
f0s, uvs = pe.get_pitch_batch(
wav_t,
sample_rate=hparams["audio_sample_rate"],
hop_size=hparams["hop_size"],
lengths=[real_len],
fmax=hparams["f0_max"],
fmin=hparams["f0_min"],
)
f0_1d, uv_1d = norm_interp_f0(f0s[0][:T])
f0_t = pad_or_cut_xd(torch.FloatTensor(f0_1d).to(device), T, 0).unsqueeze(0)
uv_t = pad_or_cut_xd(torch.FloatTensor(uv_1d).to(device), T, 0).long().unsqueeze(0)
pitch_coarse = f0_to_coarse(denorm_f0(f0_t, uv_t)).to(device)
f0_np = denorm_f0(f0_t, uv_t)[0].detach().cpu().numpy()[:real_len]
else:
f0_t = uv_t = pitch_coarse = None
f0_np = None
# ---- Mel Extraction ----
mel = mel_net(wav_t) # [1, T_padded, C]
mel = pad_or_cut_xd(mel, T, 1)
# Construct non-padding mask
mel_nonpadding_mask = torch.zeros(1, T, device=device)
mel_nonpadding_mask[:, :real_len] = 1.0
# Apply mask to mel (zero out padding)
mel = (mel.transpose(1, 2) * mel_nonpadding_mask.unsqueeze(1)).transpose(1, 2)
# Re-calculate non_padding bool mask
mel_nonpadding = mel.abs().sum(-1) > 0
# ---- Word Boundary ----
word_durs_used = None
if apply_rwbd_:
mel_input = mel[:, :, : hparams.get("wbd_use_mel_bins", 80)]
wbd_outputs = wbd_predictor(
mel=mel_input,
pitch=pitch_coarse,
uv=uv_t,
non_padding=mel_nonpadding,
train=False,
)
word_bd = wbd_outputs["word_bd_pred"] # [1, T]
else:
# Construct word_bd from provided durs
mel2word_t = pad_or_cut_xd(torch.LongTensor(mel2word).to(device), T, 0)
word_bd = torch.zeros_like(mel2word_t)
# Vectorized check
word_bd[1:] = (mel2word_t[1:] != mel2word_t[:-1]).long()
word_bd[real_len:] = 0
word_bd = word_bd.unsqueeze(0) # [1, T]
word_durs_used = np.array(word_durs_filtered)
# ---- Main Inference ----
mel_input = mel[:, :, : hparams.get("use_mel_bins", 80)]
outputs = model(
mel=mel_input,
word_bd=word_bd,
pitch=pitch_coarse,
uv=uv_t,
non_padding=mel_nonpadding,
train=False,
)
note_lengths = outputs["note_lengths"].detach().cpu().numpy()
note_bd_pred = outputs["note_bd_pred"][0].detach().cpu().numpy()[:real_len]
note_pred = outputs["note_pred"][0].detach().cpu().numpy()[: note_lengths[0]]
note_bd_logits = torch.sigmoid(outputs["note_bd_logits"])[0].detach().cpu().numpy()[:real_len]
if note_pred.shape == (0,):
if verbose:
print(f"skip {item_name}: no notes detected")
return {
"item_name": item_name,
"pitches": [],
"note_durs": [],
"note2words": None,
}
# ---- Post-Processing & Regulation ----
note_itv_pred = boundary2Interval(note_bd_pred)
note2words = None
if apply_rwbd_:
word_bd_np = outputs['word_bd_pred'][0].detach().cpu().numpy()[:real_len]
word_durs_derived = np.array(bd_to_durs(word_bd_np)) * hparams['hop_size'] / hparams['audio_sample_rate']
word_durs_for_reg = word_durs_derived
word_bd_for_reg = word_bd_np
else:
word_bd_for_reg = word_bd[0].detach().cpu().numpy()[:real_len]
word_durs_for_reg = word_durs_used
should_regulate = hparams.get("infer_regulate_real_note_itv", True) and (not apply_rwbd_)
if should_regulate and (word_durs_for_reg is not None):
try:
note_itv_pred_secs, note2words = regulate_real_note_itv(
note_itv_pred,
note_bd_pred,
word_bd_for_reg,
word_durs_for_reg,
hparams["hop_size"],
hparams["audio_sample_rate"],
)
note_pred, note_itv_pred_secs, note2words = regulate_ill_slur(note_pred, note_itv_pred_secs, note2words)
except Exception as err:
if verbose:
_, exc_value, exc_tb = sys.exc_info()
tb = traceback.extract_tb(exc_tb)[-1]
print(f"postprocess failed: {err}: {exc_value} in {tb[0]}:{tb[1]} '{tb[2]}' in {tb[3]}")
# Fallback
note_itv_pred_secs = note_itv_pred * hparams["hop_size"] / hparams["audio_sample_rate"]
note2words = None
else:
note_itv_pred_secs = note_itv_pred * hparams["hop_size"] / hparams["audio_sample_rate"]
# ---- Output ----
note_durs = [float((itv[1] - itv[0])) for itv in note_itv_pred_secs]
out = {
"item_name": item_name,
"pitches": note_pred.tolist(),
"note_durs": note_durs,
"note2words": note2words.tolist() if note2words is not None else None,
}
# ---- Saving ----
if save_dir is not None:
save_dir_path = Path(save_dir)
save_dir_path.mkdir(parents=True, exist_ok=True)
fn = str(item_name)
if not no_save_midi:
save_midi(note_pred, note_itv_pred_secs, safe_path(save_dir_path / "midi" / f"{fn}.mid"))
if not no_save_npy:
np.save(safe_path(save_dir_path / "npy" / f"[note]{fn}.npy"), out, allow_pickle=True)
if save_plot:
fig = plt.figure()
if f0_np is not None:
plt.plot(f0_np, color="red", label="f0")
midi_pred = np.zeros(note_bd_pred.shape[0], dtype=np.float32)
itvs = np.round(note_itv_pred_secs * hparams["audio_sample_rate"] / hparams["hop_size"]).astype(int)
for i, itv in enumerate(itvs):
midi_pred[itv[0] : itv[1]] = note_pred[i]
plt.plot(midi_to_hz(midi_pred), color="blue", label="pred midi")
plt.plot(note_bd_logits * 100, color="green", label="note bd logits x100")
plt.legend()
plt.tight_layout()
plt.savefig(safe_path(save_dir_path / "plot" / f"[MIDI]{fn}.png"), format="png")
plt.close(fig)
return out
def load_rosvot_models(ckpt, config="", wbd_ckpt="", wbd_config="", device="cuda:0", verbose=False, thr=0.85):
"""
Load models once to reuse across multiple items.
"""
dev = torch.device(device)
# 1. Hparams
config_path = Path(ckpt).with_name("config.yaml") if config == "" else config
pe_ckpt = Path(ckpt).parent.parent / "rmvpe/model.pt"
hparams = set_hparams(
config=config_path,
print_hparams=verbose,
hparams_str=f"note_bd_threshold={thr}",
)
# 2. Main Model
model = MidiExtractor(hparams)
load_ckpt(model, ckpt, verbose=verbose)
model.eval().to(dev)
# 3. MelNet
mel_net = MelNet(hparams)
mel_net.to(dev)
# 4. Pitch Extractor
pe = None
if hparams.get("use_pitch_embed", False):
pe = RMVPE(pe_ckpt, device=dev)
# 5. Word Boundary Predictor (optional but we load if ckpt provided or needed)
wbd_predictor = None
if wbd_ckpt:
wbd_config_path = Path(wbd_ckpt).with_name("config.yaml") if wbd_config == "" else wbd_config
wbd_hparams = set_hparams(
config=wbd_config_path,
print_hparams=False,
hparams_str="",
)
hparams.update({
"wbd_use_mel_bins": wbd_hparams["use_mel_bins"],
"min_word_dur": wbd_hparams["min_word_dur"],
})
wbd_predictor = WordbdExtractor(wbd_hparams)
load_ckpt(wbd_predictor, wbd_ckpt, verbose=verbose)
wbd_predictor.eval().to(dev)
models = {
"model": model,
"mel_net": mel_net,
"pe": pe,
"wbd_predictor": wbd_predictor
}
return hparams, models
class NoteTranscriber:
"""Note transcription wrapper based on ROSVOT.
Loads ROSVOT and optional RWBD models once in ``__init__`` and
exposes a :py:meth:`process` API that turns an item dict into
aligned note metadata for downstream SVS.
"""
def __init__(
self,
rosvot_model_path: str,
rwbd_model_path: str,
*,
rosvot_config_path: str = "",
rwbd_config_path: str = "",
device: str = "cuda:0",
thr: float = 0.85,
verbose: bool = True,
):
"""Initialize the note transcriber.
Args:
ckpt: Path to the main ROSVOT checkpoint.
config: Optional config YAML path for ROSVOT.
wbd_ckpt: Optional word-boundary checkpoint path.
wbd_config: Optional config YAML path for RWBD.
device: Torch device string, e.g. ``"cuda:0"`` / ``"cpu"``.
thr: Note boundary threshold.
verbose: Whether to print verbose logs.
"""
self.verbose = verbose
self.device = torch.device(device)
self.hparams, self.models = load_rosvot_models(
ckpt=rosvot_model_path,
config=rosvot_config_path,
wbd_ckpt=rwbd_model_path,
wbd_config=rwbd_config_path,
device=device,
verbose=verbose,
thr=thr,
)
if self.verbose:
print(
"[note transcription] init success:",
f"device={self.device}",
f"rosvot_model_path={rosvot_model_path}",
f"rwbd_model_path={rwbd_model_path if rwbd_model_path else 'None'}",
f"thr={thr}",
)
def process(
self,
item: Dict[str, Any],
*,
segment_info: Optional[Dict[str, Any]] = None,
save_dir: Optional[str] = None,
apply_rwbd: Optional[bool] = None,
save_plot: bool = False,
no_save_midi: bool = True,
no_save_npy: bool = True,
verbose: Optional[bool] = None,
) -> Dict[str, Any]:
"""Run ROSVOT on a single item and post-process outputs.
Args:
item: Input metadata dict with at least ``item_name`` and ``wav_fn``.
segment_info: Optional segment metadata for sliced audio.
save_dir: Optional directory for debug artifacts (plots, midis).
apply_rwbd: Whether to run RWBD-based word boundary refinement.
save_plot: Whether to save diagnostic plots.
no_save_midi: If True, skip saving midi.
no_save_npy: If True, skip saving numpy intermediates.
verbose: Override instance-level verbose flag for this call.
Returns:
Dict with aligned note information for downstream SVS.
"""
v = self.verbose if verbose is None else verbose
if v:
item_name = item.get("item_name", "")
wav_fn = item.get("wav_fn", "")
print(f"[note transcription] process: start: item_name={item_name} wav_fn={wav_fn}")
t0 = time.time()
rosvot_out = infer_sample(
item,
self.hparams,
self.models,
device=self.device,
save_dir=save_dir,
apply_rwbd=apply_rwbd,
save_plot=save_plot,
no_save_midi=no_save_midi,
no_save_npy=no_save_npy,
verbose=v,
)
out = self.post_process(
metadata=item,
segment_info=segment_info,
rosvot_out=rosvot_out,
)
if v:
dt = time.time() - t0
print(
"[note transcription] process: done:",
f"item_name={out.get('item_name','')}",
f"n_notes={len(out.get('note_pitch', []) or [])}",
f"time={dt:.3f}s",
)
return out
@staticmethod
def _normalize_note2words(note2words: list[int]) -> list[int]:
if not note2words:
return []
normalized = [note2words[0]]
for idx in range(1, len(note2words)):
if note2words[idx] < normalized[-1]:
normalized.append(normalized[-1])
else:
normalized.append(note2words[idx])
return normalized
@staticmethod
def _build_ep_types(note2words: list[int], align_words: list[str]) -> list[int]:
ep_types: list[int] = []
prev = -1
for i, w in zip(note2words, align_words):
if w == "<SP>":
ep_types.append(1)
else:
ep_types.append(2 if i != prev else 3)
prev = i
return ep_types
def post_process(
self,
*,
metadata: Dict[str, Any],
segment_info: Dict[str, Any],
rosvot_out: Dict[str, Any],
) -> Dict[str, Any]:
"""Build aligned note metadata using ROSVOT outputs."""
note2words_raw = rosvot_out.get("note2words") or []
note2words = self._normalize_note2words(note2words_raw)
align_words = [
metadata["words"][idx - 1]
for idx in note2words_raw
if 0 < idx <= len(metadata["words"])
]
ep_types = self._build_ep_types(note2words, align_words) if align_words else []
return {
"item_name": rosvot_out.get("item_name", "") if not segment_info else segment_info["item_name"],
"wav_fn": metadata.get("wav_fn", "") if not segment_info else segment_info["wav_fn"],
"origin_wav_fn": metadata.get("origin_wav_fn", "") if not segment_info else segment_info["origin_wav_fn"],
"start_time_ms": "" if not segment_info else segment_info["start_time_ms"],
"end_time_ms": "" if not segment_info else segment_info["end_time_ms"],
"language": metadata.get("language", ""),
"note_text": align_words,
"note_dur": rosvot_out.get("note_durs", []),
"note_type": ep_types,
"note_pitch": rosvot_out.get("pitches", []),
}
if __name__ == "__main__":
items = json.load(open("example/test/rosvot_input.json", "r"))
item = items[0]
m = NoteTranscriber(
rosvot_model_path="pretrained_models/rosvot/rosvot/model.pt",
rwbd_model_path="pretrained_models/rosvot/rwbd/model.pt",
device="cuda"
)
out = m.process(item)
print(out)
@@ -0,0 +1 @@
"""ROSVOT model submodules."""
@@ -0,0 +1 @@
"""Common ROSVOT layers and utilities."""
@@ -0,0 +1 @@
"""Conformer layers for ROSVOT."""
@@ -0,0 +1,96 @@
from torch import nn
from .espnet_positional_embedding import RelPositionalEncoding, ScaledPositionalEncoding, PositionalEncoding
from .espnet_transformer_attn import RelPositionMultiHeadedAttention, MultiHeadedAttention
from .layers import Swish, ConvolutionModule, EncoderLayer, MultiLayeredConv1d
from ..layers import Embedding
class ConformerLayers(nn.Module):
def __init__(self, hidden_size, num_layers, kernel_size=9, dropout=0.0, num_heads=4,
use_last_norm=True, save_hidden=False):
super().__init__()
self.use_last_norm = use_last_norm
self.layers = nn.ModuleList()
positionwise_layer = MultiLayeredConv1d
positionwise_layer_args = (hidden_size, hidden_size * 4, 1, dropout)
self.pos_embed = RelPositionalEncoding(hidden_size, dropout)
self.encoder_layers = nn.ModuleList([EncoderLayer(
hidden_size,
RelPositionMultiHeadedAttention(num_heads, hidden_size, 0.0),
positionwise_layer(*positionwise_layer_args),
positionwise_layer(*positionwise_layer_args),
ConvolutionModule(hidden_size, kernel_size, Swish()),
dropout,
) for _ in range(num_layers)])
if self.use_last_norm:
self.layer_norm = nn.LayerNorm(hidden_size)
else:
self.layer_norm = nn.Linear(hidden_size, hidden_size)
self.save_hidden = save_hidden
if save_hidden:
self.hiddens = []
def forward(self, x, padding_mask=None):
"""
:param x: [B, T, H]
:param padding_mask: [B, T]
:return: [B, T, H]
"""
self.hiddens = []
nonpadding_mask = x.abs().sum(-1) > 0
x = self.pos_embed(x)
for l in self.encoder_layers:
x, mask = l(x, nonpadding_mask[:, None, :])
if self.save_hidden:
self.hiddens.append(x[0])
x = x[0]
x = self.layer_norm(x) * nonpadding_mask.float()[:, :, None]
return x
class FastConformerLayers(ConformerLayers):
def __init__(self, hidden_size, num_layers, kernel_size=9, dropout=0.0, num_heads=4,
use_last_norm=True, save_hidden=False):
super(ConformerLayers, self).__init__()
self.use_last_norm = use_last_norm
self.layers = nn.ModuleList()
positionwise_layer = MultiLayeredConv1d
positionwise_layer_args = (hidden_size, hidden_size * 4, 1, dropout)
self.pos_embed = PositionalEncoding(hidden_size, dropout)
self.encoder_layers = nn.ModuleList([EncoderLayer(
hidden_size,
MultiHeadedAttention(num_heads, hidden_size, 0.0, flash=True),
positionwise_layer(*positionwise_layer_args),
positionwise_layer(*positionwise_layer_args),
ConvolutionModule(hidden_size, kernel_size, Swish()),
dropout,
) for _ in range(num_layers)])
if self.use_last_norm:
self.layer_norm = nn.LayerNorm(hidden_size)
else:
self.layer_norm = nn.Linear(hidden_size, hidden_size)
self.save_hidden = save_hidden
if save_hidden:
self.hiddens = []
class ConformerEncoder(ConformerLayers):
def __init__(self, hidden_size, dict_size, num_layers=None):
conformer_enc_kernel_size = 9
super().__init__(hidden_size, num_layers, conformer_enc_kernel_size)
self.embed = Embedding(dict_size, hidden_size, padding_idx=0)
def forward(self, x):
"""
:param src_tokens: [B, T]
:return: [B x T x C]
"""
x = self.embed(x) # [B, T, H]
x = super(ConformerEncoder, self).forward(x)
return x
class ConformerDecoder(ConformerLayers):
def __init__(self, hidden_size, num_layers):
conformer_dec_kernel_size = 9
super().__init__(hidden_size, num_layers, conformer_dec_kernel_size)
@@ -0,0 +1,113 @@
import math
import torch
class PositionalEncoding(torch.nn.Module):
"""Positional encoding.
Args:
d_model (int): Embedding dimension.
dropout_rate (float): Dropout rate.
max_len (int): Maximum input length.
reverse (bool): Whether to reverse the input position.
"""
def __init__(self, d_model, dropout_rate, max_len=5000, reverse=False):
"""Construct an PositionalEncoding object."""
super(PositionalEncoding, self).__init__()
self.d_model = d_model
self.reverse = reverse
self.xscale = math.sqrt(self.d_model)
self.dropout = torch.nn.Dropout(p=dropout_rate)
self.pe = None
self.extend_pe(torch.tensor(0.0).expand(1, max_len))
def extend_pe(self, x):
"""Reset the positional encodings."""
if self.pe is not None:
if self.pe.size(1) >= x.size(1):
if self.pe.dtype != x.dtype or self.pe.device != x.device:
self.pe = self.pe.to(dtype=x.dtype, device=x.device)
return
pe = torch.zeros(x.size(1), self.d_model)
if self.reverse:
position = torch.arange(
x.size(1) - 1, -1, -1.0, dtype=torch.float32
).unsqueeze(1)
else:
position = torch.arange(0, x.size(1), dtype=torch.float32).unsqueeze(1)
div_term = torch.exp(
torch.arange(0, self.d_model, 2, dtype=torch.float32)
* -(math.log(10000.0) / self.d_model)
)
pe[:, 0::2] = torch.sin(position * div_term)
pe[:, 1::2] = torch.cos(position * div_term)
pe = pe.unsqueeze(0)
self.pe = pe.to(device=x.device, dtype=x.dtype)
def forward(self, x: torch.Tensor):
"""Add positional encoding.
Args:
x (torch.Tensor): Input tensor (batch, time, `*`).
Returns:
torch.Tensor: Encoded tensor (batch, time, `*`).
"""
self.extend_pe(x)
x = x * self.xscale + self.pe[:, : x.size(1)]
return self.dropout(x)
class ScaledPositionalEncoding(PositionalEncoding):
"""Scaled positional encoding module.
See Sec. 3.2 https://arxiv.org/abs/1809.08895
Args:
d_model (int): Embedding dimension.
dropout_rate (float): Dropout rate.
max_len (int): Maximum input length.
"""
def __init__(self, d_model, dropout_rate, max_len=5000):
"""Initialize class."""
super().__init__(d_model=d_model, dropout_rate=dropout_rate, max_len=max_len)
self.alpha = torch.nn.Parameter(torch.tensor(1.0))
def reset_parameters(self):
"""Reset parameters."""
self.alpha.data = torch.tensor(1.0)
def forward(self, x):
"""Add positional encoding.
Args:
x (torch.Tensor): Input tensor (batch, time, `*`).
Returns:
torch.Tensor: Encoded tensor (batch, time, `*`).
"""
self.extend_pe(x)
x = x + self.alpha * self.pe[:, : x.size(1)]
return self.dropout(x)
class RelPositionalEncoding(PositionalEncoding):
"""Relative positional encoding module.
See : Appendix B in https://arxiv.org/abs/1901.02860
Args:
d_model (int): Embedding dimension.
dropout_rate (float): Dropout rate.
max_len (int): Maximum input length.
"""
def __init__(self, d_model, dropout_rate, max_len=5000):
"""Initialize class."""
super().__init__(d_model, dropout_rate, max_len, reverse=True)
def forward(self, x):
"""Compute positional encoding.
Args:
x (torch.Tensor): Input tensor (batch, time, `*`).
Returns:
torch.Tensor: Encoded tensor (batch, time, `*`).
torch.Tensor: Positional embedding tensor (1, time, `*`).
"""
self.extend_pe(x)
x = x * self.xscale
pos_emb = self.pe[:, : x.size(1)]
return self.dropout(x), self.dropout(pos_emb)
@@ -0,0 +1,198 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
# Copyright 2019 Shigeki Karita
# Apache 2.0 (http://www.apache.org/licenses/LICENSE-2.0)
"""Multi-Head Attention layer definition."""
from packaging import version
import math
import numpy
import torch
from torch import nn
class MultiHeadedAttention(nn.Module):
"""Multi-Head Attention layer.
Args:
n_head (int): The number of heads.
n_feat (int): The number of features.
dropout_rate (float): Dropout rate.
"""
def __init__(self, n_head, n_feat, dropout_rate, flash=False):
"""Construct an MultiHeadedAttention object."""
super(MultiHeadedAttention, self).__init__()
assert n_feat % n_head == 0
# We assume d_v always equals d_k
self.d_k = n_feat // n_head
self.h = n_head
self.linear_q = nn.Linear(n_feat, n_feat)
self.linear_k = nn.Linear(n_feat, n_feat)
self.linear_v = nn.Linear(n_feat, n_feat)
self.linear_out = nn.Linear(n_feat, n_feat)
self.attn = None
self.dropout = nn.Dropout(p=dropout_rate)
self.dropout_rate = dropout_rate
self.flash = flash
def forward_qkv(self, query, key, value):
"""Transform query, key and value.
Args:
query (torch.Tensor): Query tensor (#batch, time1, size).
key (torch.Tensor): Key tensor (#batch, time2, size).
value (torch.Tensor): Value tensor (#batch, time2, size).
Returns:
torch.Tensor: Transformed query tensor (#batch, n_head, time1, d_k).
torch.Tensor: Transformed key tensor (#batch, n_head, time2, d_k).
torch.Tensor: Transformed value tensor (#batch, n_head, time2, d_k).
"""
n_batch = query.size(0)
q = self.linear_q(query).view(n_batch, -1, self.h, self.d_k)
k = self.linear_k(key).view(n_batch, -1, self.h, self.d_k)
v = self.linear_v(value).view(n_batch, -1, self.h, self.d_k)
q = q.transpose(1, 2) # (batch, head, time1, d_k)
k = k.transpose(1, 2) # (batch, head, time2, d_k)
v = v.transpose(1, 2) # (batch, head, time2, d_k)
return q, k, v
def forward_attention(self, value, scores, mask):
"""Compute attention context vector.
Args:
value (torch.Tensor): Transformed value (#batch, n_head, time2, d_k).
scores (torch.Tensor): Attention score (#batch, n_head, time1, time2).
mask (torch.Tensor): Mask (#batch, 1, time2) or (#batch, time1, time2).
Returns:
torch.Tensor: Transformed value (#batch, time1, d_model)
weighted by the attention score (#batch, time1, time2).
"""
n_batch = value.size(0)
if mask is not None:
mask = mask.unsqueeze(1).eq(0) # (batch, 1, *, time2)
min_value = float(
numpy.finfo(torch.tensor(0, dtype=scores.dtype).numpy().dtype).min
)
scores = scores.masked_fill(mask, min_value)
self.attn = torch.softmax(scores, dim=-1).masked_fill(
mask, 0.0
) # (batch, head, time1, time2)
else:
self.attn = torch.softmax(scores, dim=-1) # (batch, head, time1, time2)
p_attn = self.dropout(self.attn)
x = torch.matmul(p_attn, value) # (batch, head, time1, d_k)
x = (
x.transpose(1, 2).contiguous().view(n_batch, -1, self.h * self.d_k)
) # (batch, time1, d_model)
return self.linear_out(x) # (batch, time1, d_model)
def forward(self, query, key, value, mask):
"""Compute scaled dot product attention.
Args:
query (torch.Tensor): Query tensor (#batch, time1, size).
key (torch.Tensor): Key tensor (#batch, time2, size).
value (torch.Tensor): Value tensor (#batch, time2, size).
mask (torch.Tensor): Mask tensor (#batch, 1, time2) or
(#batch, time1, time2).
Returns:
torch.Tensor: Output tensor (#batch, time1, d_model).
"""
q, k, v = self.forward_qkv(query, key, value)
if version.parse(torch.__version__) >= version.parse("2.0") and self.flash:
n_batch = value.size(0)
x = torch.nn.functional.scaled_dot_product_attention(
q, k, v, attn_mask=mask.unsqueeze(1) if mask is not None else None, dropout_p=self.dropout_rate)
x = (
x.transpose(1, 2).contiguous().view(n_batch, -1, self.h * self.d_k)
) # (batch, time1, d_model)
return self.linear_out(x)
else:
scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.d_k)
return self.forward_attention(v, scores, mask)
class RelPositionMultiHeadedAttention(MultiHeadedAttention):
"""Multi-Head Attention layer with relative position encoding.
Paper: https://arxiv.org/abs/1901.02860
Args:
n_head (int): The number of heads.
n_feat (int): The number of features.
dropout_rate (float): Dropout rate.
"""
def __init__(self, n_head, n_feat, dropout_rate):
"""Construct an RelPositionMultiHeadedAttention object."""
super().__init__(n_head, n_feat, dropout_rate)
# linear transformation for positional ecoding
self.linear_pos = nn.Linear(n_feat, n_feat, bias=False)
# these two learnable bias are used in matrix c and matrix d
# as described in https://arxiv.org/abs/1901.02860 Section 3.3
self.pos_bias_u = nn.Parameter(torch.Tensor(self.h, self.d_k))
self.pos_bias_v = nn.Parameter(torch.Tensor(self.h, self.d_k))
torch.nn.init.xavier_uniform_(self.pos_bias_u)
torch.nn.init.xavier_uniform_(self.pos_bias_v)
def rel_shift(self, x, zero_triu=False):
"""Compute relative positinal encoding.
Args:
x (torch.Tensor): Input tensor (batch, time, size).
zero_triu (bool): If true, return the lower triangular part of the matrix.
Returns:
torch.Tensor: Output tensor.
"""
zero_pad = torch.zeros((*x.size()[:3], 1), device=x.device, dtype=x.dtype)
x_padded = torch.cat([zero_pad, x], dim=-1)
x_padded = x_padded.view(*x.size()[:2], x.size(3) + 1, x.size(2))
x = x_padded[:, :, 1:].view_as(x)
if zero_triu:
ones = torch.ones((x.size(2), x.size(3)))
x = x * torch.tril(ones, x.size(3) - x.size(2))[None, None, :, :]
return x
def forward(self, query, key, value, pos_emb, mask):
"""Compute 'Scaled Dot Product Attention' with rel. positional encoding.
Args:
query (torch.Tensor): Query tensor (#batch, time1, size).
key (torch.Tensor): Key tensor (#batch, time2, size).
value (torch.Tensor): Value tensor (#batch, time2, size).
pos_emb (torch.Tensor): Positional embedding tensor (#batch, time2, size).
mask (torch.Tensor): Mask tensor (#batch, 1, time2) or
(#batch, time1, time2).
Returns:
torch.Tensor: Output tensor (#batch, time1, d_model).
"""
q, k, v = self.forward_qkv(query, key, value)
q = q.transpose(1, 2) # (batch, time1, head, d_k)
n_batch_pos = pos_emb.size(0)
p = self.linear_pos(pos_emb).view(n_batch_pos, -1, self.h, self.d_k)
p = p.transpose(1, 2) # (batch, head, time1, d_k)
# (batch, head, time1, d_k)
q_with_bias_u = (q + self.pos_bias_u).transpose(1, 2)
# (batch, head, time1, d_k)
q_with_bias_v = (q + self.pos_bias_v).transpose(1, 2)
# compute attention score
# first compute matrix a and matrix c
# as described in https://arxiv.org/abs/1901.02860 Section 3.3
# (batch, head, time1, time2)
matrix_ac = torch.matmul(q_with_bias_u, k.transpose(-2, -1))
# compute matrix b and matrix d
# (batch, head, time1, time2)
matrix_bd = torch.matmul(q_with_bias_v, p.transpose(-2, -1))
matrix_bd = self.rel_shift(matrix_bd)
scores = (matrix_ac + matrix_bd) / math.sqrt(
self.d_k
) # (batch, head, time1, time2)
return self.forward_attention(v, scores, mask)
@@ -0,0 +1,260 @@
from torch import nn
import torch
from ..layers import LayerNorm
class ConvolutionModule(nn.Module):
"""ConvolutionModule in Conformer model.
Args:
channels (int): The number of channels of conv layers.
kernel_size (int): Kernerl size of conv layers.
"""
def __init__(self, channels, kernel_size, activation=nn.ReLU(), bias=True):
"""Construct an ConvolutionModule object."""
super(ConvolutionModule, self).__init__()
# kernerl_size should be a odd number for 'SAME' padding
assert (kernel_size - 1) % 2 == 0
self.pointwise_conv1 = nn.Conv1d(
channels,
2 * channels,
kernel_size=1,
stride=1,
padding=0,
bias=bias,
)
self.depthwise_conv = nn.Conv1d(
channels,
channels,
kernel_size,
stride=1,
padding=(kernel_size - 1) // 2,
groups=channels,
bias=bias,
)
self.norm = nn.BatchNorm1d(channels)
self.pointwise_conv2 = nn.Conv1d(
channels,
channels,
kernel_size=1,
stride=1,
padding=0,
bias=bias,
)
self.activation = activation
def forward(self, x):
"""Compute convolution module.
Args:
x (torch.Tensor): Input tensor (#batch, time, channels).
Returns:
torch.Tensor: Output tensor (#batch, time, channels).
"""
# exchange the temporal dimension and the feature dimension
x = x.transpose(1, 2)
# GLU mechanism
x = self.pointwise_conv1(x) # (batch, 2*channel, dim)
x = nn.functional.glu(x, dim=1) # (batch, channel, dim)
# 1D Depthwise Conv
x = self.depthwise_conv(x)
x = self.activation(self.norm(x))
x = self.pointwise_conv2(x)
return x.transpose(1, 2)
class MultiLayeredConv1d(torch.nn.Module):
"""Multi-layered conv1d for Transformer block.
This is a module of multi-leyered conv1d designed
to replace positionwise feed-forward network
in Transforner block, which is introduced in
`FastSpeech: Fast, Robust and Controllable Text to Speech`_.
.. _`FastSpeech: Fast, Robust and Controllable Text to Speech`:
https://arxiv.org/pdf/1905.09263.pdf
"""
def __init__(self, in_chans, hidden_chans, kernel_size, dropout_rate):
"""Initialize MultiLayeredConv1d module.
Args:
in_chans (int): Number of input channels.
hidden_chans (int): Number of hidden channels.
kernel_size (int): Kernel size of conv1d.
dropout_rate (float): Dropout rate.
"""
super(MultiLayeredConv1d, self).__init__()
self.w_1 = torch.nn.Conv1d(
in_chans,
hidden_chans,
kernel_size,
stride=1,
padding=(kernel_size - 1) // 2,
)
self.w_2 = torch.nn.Conv1d(
hidden_chans,
in_chans,
kernel_size,
stride=1,
padding=(kernel_size - 1) // 2,
)
self.dropout = torch.nn.Dropout(dropout_rate)
def forward(self, x):
"""Calculate forward propagation.
Args:
x (torch.Tensor): Batch of input tensors (B, T, in_chans).
Returns:
torch.Tensor: Batch of output tensors (B, T, hidden_chans).
"""
x = torch.relu(self.w_1(x.transpose(-1, 1))).transpose(-1, 1)
return self.w_2(self.dropout(x).transpose(-1, 1)).transpose(-1, 1)
class Swish(torch.nn.Module):
"""Construct an Swish object."""
def forward(self, x):
"""Return Swich activation function."""
return x * torch.sigmoid(x)
class EncoderLayer(nn.Module):
"""Encoder layer module.
Args:
size (int): Input dimension.
self_attn (torch.nn.Module): Self-attention module instance.
`MultiHeadedAttention` or `RelPositionMultiHeadedAttention` instance
can be used as the argument.
feed_forward (torch.nn.Module): Feed-forward module instance.
`PositionwiseFeedForward`, `MultiLayeredConv1d`, or `Conv1dLinear` instance
can be used as the argument.
feed_forward_macaron (torch.nn.Module): Additional feed-forward module instance.
`PositionwiseFeedForward`, `MultiLayeredConv1d`, or `Conv1dLinear` instance
can be used as the argument.
conv_module (torch.nn.Module): Convolution module instance.
`ConvlutionModule` instance can be used as the argument.
dropout_rate (float): Dropout rate.
normalize_before (bool): Whether to use layer_norm before the first block.
concat_after (bool): Whether to concat attention layer's input and output.
if True, additional linear will be applied.
i.e. x -> x + linear(concat(x, att(x)))
if False, no additional linear will be applied. i.e. x -> x + att(x)
"""
def __init__(
self,
size,
self_attn,
feed_forward,
feed_forward_macaron,
conv_module,
dropout_rate,
normalize_before=True,
concat_after=False,
):
"""Construct an EncoderLayer object."""
super(EncoderLayer, self).__init__()
self.self_attn = self_attn
self.feed_forward = feed_forward
self.feed_forward_macaron = feed_forward_macaron
self.conv_module = conv_module
self.norm_ff = LayerNorm(size) # for the FNN module
self.norm_mha = LayerNorm(size) # for the MHA module
if feed_forward_macaron is not None:
self.norm_ff_macaron = LayerNorm(size)
self.ff_scale = 0.5
else:
self.ff_scale = 1.0
if self.conv_module is not None:
self.norm_conv = LayerNorm(size) # for the CNN module
self.norm_final = LayerNorm(size) # for the final output of the block
self.dropout = nn.Dropout(dropout_rate)
self.size = size
self.normalize_before = normalize_before
self.concat_after = concat_after
if self.concat_after:
self.concat_linear = nn.Linear(size + size, size)
def forward(self, x_input, mask, cache=None):
"""Compute encoded features.
Args:
x_input (Union[Tuple, torch.Tensor]): Input tensor w/ or w/o pos emb.
- w/ pos emb: Tuple of tensors [(#batch, time, size), (1, time, size)].
- w/o pos emb: Tensor (#batch, time, size).
mask (torch.Tensor): Mask tensor for the input (#batch, time).
cache (torch.Tensor): Cache tensor of the input (#batch, time - 1, size).
Returns:
torch.Tensor: Output tensor (#batch, time, size).
torch.Tensor: Mask tensor (#batch, time).
"""
if isinstance(x_input, tuple):
x, pos_emb = x_input[0], x_input[1]
else:
x, pos_emb = x_input, None
# whether to use macaron style
if self.feed_forward_macaron is not None:
residual = x
if self.normalize_before:
x = self.norm_ff_macaron(x)
x = residual + self.ff_scale * self.dropout(self.feed_forward_macaron(x))
if not self.normalize_before:
x = self.norm_ff_macaron(x)
# multi-headed self-attention module
residual = x
if self.normalize_before:
x = self.norm_mha(x)
if cache is None:
x_q = x
else:
assert cache.shape == (x.shape[0], x.shape[1] - 1, self.size)
x_q = x[:, -1:, :]
residual = residual[:, -1:, :]
mask = None if mask is None else mask[:, -1:, :]
if pos_emb is not None:
x_att = self.self_attn(x_q, x, x, pos_emb, mask)
else:
x_att = self.self_attn(x_q, x, x, mask)
if self.concat_after:
x_concat = torch.cat((x, x_att), dim=-1)
x = residual + self.concat_linear(x_concat)
else:
x = residual + self.dropout(x_att)
if not self.normalize_before:
x = self.norm_mha(x)
# convolution module
if self.conv_module is not None:
residual = x
if self.normalize_before:
x = self.norm_conv(x)
x = residual + self.dropout(self.conv_module(x))
if not self.normalize_before:
x = self.norm_conv(x)
# feed forward module
residual = x
if self.normalize_before:
x = self.norm_ff(x)
x = residual + self.ff_scale * self.dropout(self.feed_forward(x))
if not self.normalize_before:
x = self.norm_ff(x)
if self.conv_module is not None:
x = self.norm_final(x)
if cache is not None:
x = torch.cat([cache, x], dim=1)
if pos_emb is not None:
return (x, pos_emb), mask
return x, mask
@@ -0,0 +1,175 @@
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from .layers import LayerNorm, Embedding
class LambdaLayer(nn.Module):
def __init__(self, lambd):
super(LambdaLayer, self).__init__()
self.lambd = lambd
def forward(self, x):
return self.lambd(x)
def init_weights_func(m):
classname = m.__class__.__name__
if classname.find("Conv1d") != -1:
torch.nn.init.xavier_uniform_(m.weight)
def get_norm_builder(norm_type, channels, ln_eps=1e-6):
if norm_type == 'bn':
norm_builder = lambda: nn.BatchNorm1d(channels)
elif norm_type == 'in':
norm_builder = lambda: nn.InstanceNorm1d(channels, affine=True)
elif norm_type == 'gn':
norm_builder = lambda: nn.GroupNorm(8, channels)
elif norm_type == 'ln':
norm_builder = lambda: LayerNorm(channels, dim=1, eps=ln_eps)
else:
norm_builder = lambda: nn.Identity()
return norm_builder
def get_act_builder(act_type):
if act_type == 'gelu':
act_builder = lambda: nn.GELU()
elif act_type == 'relu':
act_builder = lambda: nn.ReLU(inplace=True)
elif act_type == 'leakyrelu':
act_builder = lambda: nn.LeakyReLU(negative_slope=0.01, inplace=True)
elif act_type == 'swish':
act_builder = lambda: nn.SiLU(inplace=True)
else:
act_builder = lambda: nn.Identity()
return act_builder
class ResidualBlock(nn.Module):
"""Implements conv->PReLU->norm n-times"""
def __init__(self, channels, kernel_size, dilation, n=2, norm_type='bn', dropout=0.0,
c_multiple=2, ln_eps=1e-12, act_type='gelu'):
super(ResidualBlock, self).__init__()
norm_builder = get_norm_builder(norm_type, channels, ln_eps)
act_builder = get_act_builder(act_type)
self.blocks = [
nn.Sequential(
norm_builder(),
nn.Conv1d(channels, c_multiple * channels, kernel_size, dilation=dilation,
padding=(dilation * (kernel_size - 1)) // 2),
LambdaLayer(lambda x: x * kernel_size ** -0.5),
act_builder(),
nn.Conv1d(c_multiple * channels, channels, 1, dilation=dilation),
)
for i in range(n)
]
self.blocks = nn.ModuleList(self.blocks)
self.dropout = dropout
def forward(self, x):
nonpadding = (x.abs().sum(1) > 0).float()[:, None, :]
for b in self.blocks:
x_ = b(x)
if self.dropout > 0 and self.training:
x_ = F.dropout(x_, self.dropout, training=self.training)
x = x + x_
x = x * nonpadding
return x
class ConvBlocks(nn.Module):
"""Decodes the expanded phoneme encoding into spectrograms"""
def __init__(self, hidden_size, out_dims, dilations, kernel_size,
norm_type='ln', layers_in_block=2, c_multiple=2,
dropout=0.0, ln_eps=1e-5,
init_weights=True, is_BTC=True, num_layers=None, post_net_kernel=3, act_type='gelu'):
super(ConvBlocks, self).__init__()
self.is_BTC = is_BTC
if num_layers is not None:
dilations = [1] * num_layers
self.res_blocks = nn.Sequential(
*[ResidualBlock(hidden_size, kernel_size, d,
n=layers_in_block, norm_type=norm_type, c_multiple=c_multiple,
dropout=dropout, ln_eps=ln_eps, act_type=act_type)
for d in dilations],
)
norm = get_norm_builder(norm_type, hidden_size, ln_eps)()
self.last_norm = norm
self.post_net1 = nn.Conv1d(hidden_size, out_dims, kernel_size=post_net_kernel,
padding=post_net_kernel // 2)
if init_weights:
self.apply(init_weights_func)
def forward(self, x, nonpadding=None):
"""
:param x: [B, T, H]
:return: [B, T, H]
"""
if self.is_BTC:
x = x.transpose(1, 2)
if nonpadding is None:
nonpadding = (x.abs().sum(1) > 0).float()[:, None, :]
elif self.is_BTC:
nonpadding = nonpadding.transpose(1, 2)
x = self.res_blocks(x) * nonpadding
x = self.last_norm(x) * nonpadding
x = self.post_net1(x) * nonpadding
if self.is_BTC:
x = x.transpose(1, 2)
return x
class TextConvEncoder(ConvBlocks):
def __init__(self, dict_size, hidden_size, out_dims, dilations, kernel_size,
norm_type='ln', layers_in_block=2, c_multiple=2,
dropout=0.0, ln_eps=1e-5, init_weights=True, num_layers=None, post_net_kernel=3):
super().__init__(hidden_size, out_dims, dilations, kernel_size,
norm_type, layers_in_block, c_multiple,
dropout, ln_eps, init_weights, num_layers=num_layers,
post_net_kernel=post_net_kernel)
self.embed_tokens = Embedding(dict_size, hidden_size, 0)
self.embed_scale = math.sqrt(hidden_size)
def forward(self, txt_tokens):
"""
:param txt_tokens: [B, T]
:return: {
'encoder_out': [B x T x C]
}
"""
x = self.embed_scale * self.embed_tokens(txt_tokens)
return super().forward(x)
class ConditionalConvBlocks(ConvBlocks):
def __init__(self, hidden_size, c_cond, c_out, dilations, kernel_size,
norm_type='ln', layers_in_block=2, c_multiple=2,
dropout=0.0, ln_eps=1e-5, init_weights=True, is_BTC=True, num_layers=None):
super().__init__(hidden_size, c_out, dilations, kernel_size,
norm_type, layers_in_block, c_multiple,
dropout, ln_eps, init_weights, is_BTC=False, num_layers=num_layers)
self.g_prenet = nn.Conv1d(c_cond, hidden_size, 3, padding=1)
self.is_BTC_ = is_BTC
if init_weights:
self.g_prenet.apply(init_weights_func)
def forward(self, x, cond, nonpadding=None):
if self.is_BTC_:
x = x.transpose(1, 2)
cond = cond.transpose(1, 2)
if nonpadding is not None:
nonpadding = nonpadding.transpose(1, 2)
if nonpadding is None:
nonpadding = x.abs().sum(1)[:, None]
x = x + self.g_prenet(cond)
x = x * nonpadding
x = super(ConditionalConvBlocks, self).forward(x) # input needs to be BTC
if self.is_BTC_:
x = x.transpose(1, 2)
return x
@@ -0,0 +1,85 @@
import torch
from torch import nn
from torch.autograd import Function
class LayerNorm(torch.nn.LayerNorm):
"""Layer normalization module.
:param int nout: output dim size
:param int dim: dimension to be normalized
"""
def __init__(self, nout, dim=-1, eps=1e-5):
"""Construct an LayerNorm object."""
super(LayerNorm, self).__init__(nout, eps=eps)
self.dim = dim
def forward(self, x):
"""Apply layer normalization.
:param torch.Tensor x: input tensor
:return: layer normalized tensor
:rtype torch.Tensor
"""
if self.dim == -1:
return super(LayerNorm, self).forward(x)
return super(LayerNorm, self).forward(x.transpose(1, -1)).transpose(1, -1)
class Reshape(nn.Module):
def __init__(self, *args):
super(Reshape, self).__init__()
self.shape = args
def forward(self, x):
return x.view(self.shape)
class Permute(nn.Module):
def __init__(self, *args):
super(Permute, self).__init__()
self.args = args
def forward(self, x):
return x.permute(self.args)
def Linear(in_features, out_features, bias=True, init_type='xavier'):
m = nn.Linear(in_features, out_features, bias)
if init_type == 'xavier':
nn.init.xavier_uniform_(m.weight)
elif init_type == 'kaiming':
nn.init.kaiming_normal_(m.weight, mode='fan_in')
if bias:
nn.init.constant_(m.bias, 0.)
return m
def Embedding(num_embeddings, embedding_dim, padding_idx=None, init_type='normal'):
m = nn.Embedding(num_embeddings, embedding_dim, padding_idx=padding_idx)
if init_type == 'normal':
nn.init.normal_(m.weight, mean=0, std=embedding_dim ** -0.5)
elif init_type == 'kaiming':
nn.init.kaiming_normal_(m.weight, mode='fan_in')
if padding_idx is not None:
nn.init.constant_(m.weight[padding_idx], 0)
return m
class GradientReverseFunction(Function):
@staticmethod
def forward(ctx, input, coeff=1.):
ctx.coeff = coeff
output = input * 1.0
return output
@staticmethod
def backward(ctx, grad_output):
return grad_output.neg() * ctx.coeff, None
class GRL(nn.Module):
def __init__(self):
super(GRL, self).__init__()
def forward(self, *input):
return GradientReverseFunction.apply(*input)
@@ -0,0 +1,378 @@
import math
import torch
from torch import nn
from torch.nn import functional as F
from .layers import Embedding
def convert_pad_shape(pad_shape):
l = pad_shape[::-1]
pad_shape = [item for sublist in l for item in sublist]
return pad_shape
def shift_1d(x):
x = F.pad(x, convert_pad_shape([[0, 0], [0, 0], [1, 0]]))[:, :, :-1]
return x
def sequence_mask(length, max_length=None):
if max_length is None:
max_length = length.max()
x = torch.arange(max_length, dtype=length.dtype, device=length.device)
return x.unsqueeze(0) < length.unsqueeze(1)
class Encoder(nn.Module):
def __init__(self, hidden_channels, filter_channels, n_heads, n_layers, kernel_size=1, p_dropout=0.,
window_size=None, block_length=None, pre_ln=False, **kwargs):
super().__init__()
self.hidden_channels = hidden_channels
self.filter_channels = filter_channels
self.n_heads = n_heads
self.n_layers = n_layers
self.kernel_size = kernel_size
self.p_dropout = p_dropout
self.window_size = window_size
self.block_length = block_length
self.pre_ln = pre_ln
self.drop = nn.Dropout(p_dropout)
self.attn_layers = nn.ModuleList()
self.norm_layers_1 = nn.ModuleList()
self.ffn_layers = nn.ModuleList()
self.norm_layers_2 = nn.ModuleList()
for i in range(self.n_layers):
self.attn_layers.append(
MultiHeadAttention(hidden_channels, hidden_channels, n_heads, window_size=window_size,
p_dropout=p_dropout, block_length=block_length))
self.norm_layers_1.append(LayerNorm(hidden_channels))
self.ffn_layers.append(
FFN(hidden_channels, hidden_channels, filter_channels, kernel_size, p_dropout=p_dropout))
self.norm_layers_2.append(LayerNorm(hidden_channels))
if pre_ln:
self.last_ln = LayerNorm(hidden_channels)
def forward(self, x, x_mask):
attn_mask = x_mask.unsqueeze(2) * x_mask.unsqueeze(-1)
for i in range(self.n_layers):
x = x * x_mask
x_ = x
if self.pre_ln:
x = self.norm_layers_1[i](x)
y = self.attn_layers[i](x, x, attn_mask)
y = self.drop(y)
x = x_ + y
if not self.pre_ln:
x = self.norm_layers_1[i](x)
x_ = x
if self.pre_ln:
x = self.norm_layers_2[i](x)
y = self.ffn_layers[i](x, x_mask)
y = self.drop(y)
x = x_ + y
if not self.pre_ln:
x = self.norm_layers_2[i](x)
if self.pre_ln:
x = self.last_ln(x)
x = x * x_mask
return x
class MultiHeadAttention(nn.Module):
def __init__(self, channels, out_channels, n_heads, window_size=None, heads_share=True, p_dropout=0.,
block_length=None, proximal_bias=False, proximal_init=False):
super().__init__()
assert channels % n_heads == 0
self.channels = channels
self.out_channels = out_channels
self.n_heads = n_heads
self.window_size = window_size
self.heads_share = heads_share
self.block_length = block_length
self.proximal_bias = proximal_bias
self.p_dropout = p_dropout
self.attn = None
self.k_channels = channels // n_heads
self.conv_q = nn.Conv1d(channels, channels, 1)
self.conv_k = nn.Conv1d(channels, channels, 1)
self.conv_v = nn.Conv1d(channels, channels, 1)
if window_size is not None:
n_heads_rel = 1 if heads_share else n_heads
rel_stddev = self.k_channels ** -0.5
self.emb_rel_k = nn.Parameter(torch.randn(n_heads_rel, window_size * 2 + 1, self.k_channels) * rel_stddev)
self.emb_rel_v = nn.Parameter(torch.randn(n_heads_rel, window_size * 2 + 1, self.k_channels) * rel_stddev)
self.conv_o = nn.Conv1d(channels, out_channels, 1)
self.drop = nn.Dropout(p_dropout)
nn.init.xavier_uniform_(self.conv_q.weight)
nn.init.xavier_uniform_(self.conv_k.weight)
if proximal_init:
self.conv_k.weight.data.copy_(self.conv_q.weight.data)
self.conv_k.bias.data.copy_(self.conv_q.bias.data)
nn.init.xavier_uniform_(self.conv_v.weight)
def forward(self, x, c, attn_mask=None):
q = self.conv_q(x)
k = self.conv_k(c)
v = self.conv_v(c)
x, self.attn = self.attention(q, k, v, mask=attn_mask)
x = self.conv_o(x)
return x
def attention(self, query, key, value, mask=None):
# reshape [b, d, t] -> [b, n_h, t, d_k]
b, d, t_s, t_t = (*key.size(), query.size(2))
query = query.view(b, self.n_heads, self.k_channels, t_t).transpose(2, 3)
key = key.view(b, self.n_heads, self.k_channels, t_s).transpose(2, 3)
value = value.view(b, self.n_heads, self.k_channels, t_s).transpose(2, 3)
scores = torch.matmul(query, key.transpose(-2, -1)) / math.sqrt(self.k_channels)
if self.window_size is not None:
assert t_s == t_t, "Relative attention is only available for self-attention."
key_relative_embeddings = self._get_relative_embeddings(self.emb_rel_k, t_s)
rel_logits = self._matmul_with_relative_keys(query, key_relative_embeddings)
rel_logits = self._relative_position_to_absolute_position(rel_logits)
scores_local = rel_logits / math.sqrt(self.k_channels)
scores = scores + scores_local
if self.proximal_bias:
assert t_s == t_t, "Proximal bias is only available for self-attention."
scores = scores + self._attention_bias_proximal(t_s).to(device=scores.device, dtype=scores.dtype)
if mask is not None:
scores = scores.masked_fill(mask == 0, -1e4)
if self.block_length is not None:
block_mask = torch.ones_like(scores).triu(-self.block_length).tril(self.block_length)
scores = scores * block_mask + -1e4 * (1 - block_mask)
p_attn = F.softmax(scores, dim=-1) # [b, n_h, t_t, t_s]
p_attn = self.drop(p_attn)
output = torch.matmul(p_attn, value)
if self.window_size is not None:
relative_weights = self._absolute_position_to_relative_position(p_attn)
value_relative_embeddings = self._get_relative_embeddings(self.emb_rel_v, t_s)
output = output + self._matmul_with_relative_values(relative_weights, value_relative_embeddings)
output = output.transpose(2, 3).contiguous().view(b, d, t_t) # [b, n_h, t_t, d_k] -> [b, d, t_t]
return output, p_attn
def _matmul_with_relative_values(self, x, y):
"""
x: [b, h, l, m]
y: [h or 1, m, d]
ret: [b, h, l, d]
"""
ret = torch.matmul(x, y.unsqueeze(0))
return ret
def _matmul_with_relative_keys(self, x, y):
"""
x: [b, h, l, d]
y: [h or 1, m, d]
ret: [b, h, l, m]
"""
ret = torch.matmul(x, y.unsqueeze(0).transpose(-2, -1))
return ret
def _get_relative_embeddings(self, relative_embeddings, length):
max_relative_position = 2 * self.window_size + 1
# Pad first before slice to avoid using cond ops.
pad_length = max(length - (self.window_size + 1), 0)
slice_start_position = max((self.window_size + 1) - length, 0)
slice_end_position = slice_start_position + 2 * length - 1
if pad_length > 0:
padded_relative_embeddings = F.pad(
relative_embeddings,
convert_pad_shape([[0, 0], [pad_length, pad_length], [0, 0]]))
else:
padded_relative_embeddings = relative_embeddings
used_relative_embeddings = padded_relative_embeddings[:, slice_start_position:slice_end_position]
return used_relative_embeddings
def _relative_position_to_absolute_position(self, x):
"""
x: [b, h, l, 2*l-1]
ret: [b, h, l, l]
"""
batch, heads, length, _ = x.size()
# Concat columns of pad to shift from relative to absolute indexing.
x = F.pad(x, convert_pad_shape([[0, 0], [0, 0], [0, 0], [0, 1]]))
# Concat extra elements so to add up to shape (len+1, 2*len-1).
x_flat = x.view([batch, heads, length * 2 * length])
x_flat = F.pad(x_flat, convert_pad_shape([[0, 0], [0, 0], [0, length - 1]]))
# Reshape and slice out the padded elements.
x_final = x_flat.view([batch, heads, length + 1, 2 * length - 1])[:, :, :length, length - 1:]
return x_final
def _absolute_position_to_relative_position(self, x):
"""
x: [b, h, l, l]
ret: [b, h, l, 2*l-1]
"""
batch, heads, length, _ = x.size()
# padd along column
x = F.pad(x, convert_pad_shape([[0, 0], [0, 0], [0, 0], [0, length - 1]]))
x_flat = x.view([batch, heads, length ** 2 + length * (length - 1)])
# add 0's in the beginning that will skew the elements after reshape
x_flat = F.pad(x_flat, convert_pad_shape([[0, 0], [0, 0], [length, 0]]))
x_final = x_flat.view([batch, heads, length, 2 * length])[:, :, :, 1:]
return x_final
def _attention_bias_proximal(self, length):
"""Bias for self-attention to encourage attention to close positions.
Args:
length: an integer scalar.
Returns:
a Tensor with shape [1, 1, length, length]
"""
r = torch.arange(length, dtype=torch.float32)
diff = torch.unsqueeze(r, 0) - torch.unsqueeze(r, 1)
return torch.unsqueeze(torch.unsqueeze(-torch.log1p(torch.abs(diff)), 0), 0)
class FFN(nn.Module):
def __init__(self, in_channels, out_channels, filter_channels, kernel_size, p_dropout=0., activation=None):
super().__init__()
self.in_channels = in_channels
self.out_channels = out_channels
self.filter_channels = filter_channels
self.kernel_size = kernel_size
self.p_dropout = p_dropout
self.activation = activation
self.conv_1 = nn.Conv1d(in_channels, filter_channels, kernel_size, padding=kernel_size // 2)
self.conv_2 = nn.Conv1d(filter_channels, out_channels, 1)
self.drop = nn.Dropout(p_dropout)
def forward(self, x, x_mask):
x = self.conv_1(x * x_mask)
if self.activation == "gelu":
x = x * torch.sigmoid(1.702 * x)
else:
x = torch.relu(x)
x = self.drop(x)
x = self.conv_2(x * x_mask)
return x * x_mask
class LayerNorm(nn.Module):
def __init__(self, channels, eps=1e-4):
super().__init__()
self.channels = channels
self.eps = eps
self.gamma = nn.Parameter(torch.ones(channels))
self.beta = nn.Parameter(torch.zeros(channels))
def forward(self, x):
n_dims = len(x.shape)
mean = torch.mean(x, 1, keepdim=True)
variance = torch.mean((x - mean) ** 2, 1, keepdim=True)
x = (x - mean) * torch.rsqrt(variance + self.eps)
shape = [1, -1] + [1] * (n_dims - 2)
x = x * self.gamma.view(*shape) + self.beta.view(*shape)
return x
class ConvReluNorm(nn.Module):
def __init__(self, in_channels, hidden_channels, out_channels, kernel_size, n_layers, p_dropout):
super().__init__()
self.in_channels = in_channels
self.hidden_channels = hidden_channels
self.out_channels = out_channels
self.kernel_size = kernel_size
self.n_layers = n_layers
self.p_dropout = p_dropout
assert n_layers > 1, "Number of layers should be larger than 0."
self.conv_layers = nn.ModuleList()
self.norm_layers = nn.ModuleList()
self.conv_layers.append(nn.Conv1d(in_channels, hidden_channels, kernel_size, padding=kernel_size // 2))
self.norm_layers.append(LayerNorm(hidden_channels))
self.relu_drop = nn.Sequential(
nn.ReLU(),
nn.Dropout(p_dropout))
for _ in range(n_layers - 1):
self.conv_layers.append(nn.Conv1d(hidden_channels, hidden_channels, kernel_size, padding=kernel_size // 2))
self.norm_layers.append(LayerNorm(hidden_channels))
self.proj = nn.Conv1d(hidden_channels, out_channels, 1)
self.proj.weight.data.zero_()
self.proj.bias.data.zero_()
def forward(self, x, x_mask):
x_org = x
for i in range(self.n_layers):
x = self.conv_layers[i](x * x_mask)
x = self.norm_layers[i](x)
x = self.relu_drop(x)
x = x_org + self.proj(x)
return x * x_mask
class RelTransformerEncoder(nn.Module):
def __init__(self,
n_vocab,
out_channels,
hidden_channels,
filter_channels,
n_heads,
n_layers,
kernel_size,
p_dropout=0.0,
window_size=4,
block_length=None,
prenet=True,
pre_ln=True,
):
super().__init__()
self.n_vocab = n_vocab
self.out_channels = out_channels
self.hidden_channels = hidden_channels
self.filter_channels = filter_channels
self.n_heads = n_heads
self.n_layers = n_layers
self.kernel_size = kernel_size
self.p_dropout = p_dropout
self.window_size = window_size
self.block_length = block_length
self.prenet = prenet
if n_vocab > 0:
self.emb = Embedding(n_vocab, hidden_channels, padding_idx=0)
if prenet:
self.pre = ConvReluNorm(hidden_channels, hidden_channels, hidden_channels,
kernel_size=5, n_layers=3, p_dropout=0)
self.encoder = Encoder(
hidden_channels,
filter_channels,
n_heads,
n_layers,
kernel_size,
p_dropout,
window_size=window_size,
block_length=block_length,
pre_ln=pre_ln,
)
def forward(self, x, x_mask=None):
if self.n_vocab > 0:
x_lengths = (x > 0).long().sum(-1)
x = self.emb(x) * math.sqrt(self.hidden_channels) # [b, t, h]
else:
x_lengths = (x.abs().sum(-1) > 0).long().sum(-1)
x = torch.transpose(x, 1, -1) # [b, h, t]
x_mask = torch.unsqueeze(sequence_mask(x_lengths, x.size(2)), 1).to(x.dtype)
if self.prenet:
x = self.pre(x, x_mask)
x = self.encoder(x, x_mask)
return x.transpose(1, 2)
@@ -0,0 +1,261 @@
import torch
from torch import nn
import torch.nn.functional as F
class PreNet(nn.Module):
def __init__(self, in_dims, fc1_dims=256, fc2_dims=128, dropout=0.5):
super().__init__()
self.fc1 = nn.Linear(in_dims, fc1_dims)
self.fc2 = nn.Linear(fc1_dims, fc2_dims)
self.p = dropout
def forward(self, x):
x = self.fc1(x)
x = F.relu(x)
x = F.dropout(x, self.p, training=self.training)
x = self.fc2(x)
x = F.relu(x)
x = F.dropout(x, self.p, training=self.training)
return x
class HighwayNetwork(nn.Module):
def __init__(self, size):
super().__init__()
self.W1 = nn.Linear(size, size)
self.W2 = nn.Linear(size, size)
self.W1.bias.data.fill_(0.)
def forward(self, x):
x1 = self.W1(x)
x2 = self.W2(x)
g = torch.sigmoid(x2)
y = g * F.relu(x1) + (1. - g) * x
return y
class BatchNormConv(nn.Module):
def __init__(self, in_channels, out_channels, kernel, relu=True):
super().__init__()
self.conv = nn.Conv1d(in_channels, out_channels, kernel, stride=1, padding=kernel // 2, bias=False)
self.bnorm = nn.BatchNorm1d(out_channels)
self.relu = relu
def forward(self, x):
x = self.conv(x)
x = F.relu(x) if self.relu is True else x
return self.bnorm(x)
class ConvNorm(torch.nn.Module):
def __init__(self, in_channels, out_channels, kernel_size=1, stride=1,
padding=None, dilation=1, bias=True, w_init_gain='linear'):
super(ConvNorm, self).__init__()
if padding is None:
assert (kernel_size % 2 == 1)
padding = int(dilation * (kernel_size - 1) / 2)
self.conv = torch.nn.Conv1d(in_channels, out_channels,
kernel_size=kernel_size, stride=stride,
padding=padding, dilation=dilation,
bias=bias)
torch.nn.init.xavier_uniform_(
self.conv.weight, gain=torch.nn.init.calculate_gain(w_init_gain))
def forward(self, signal):
conv_signal = self.conv(signal)
return conv_signal
class CBHG(nn.Module):
def __init__(self, K, in_channels, channels, proj_channels, num_highways):
super().__init__()
# List of all rnns to call `flatten_parameters()` on
self._to_flatten = []
self.bank_kernels = [i for i in range(1, K + 1)]
self.conv1d_bank = nn.ModuleList()
for k in self.bank_kernels:
conv = BatchNormConv(in_channels, channels, k)
self.conv1d_bank.append(conv)
self.maxpool = nn.MaxPool1d(kernel_size=2, stride=1, padding=1)
self.conv_project1 = BatchNormConv(len(self.bank_kernels) * channels, proj_channels[0], 3)
self.conv_project2 = BatchNormConv(proj_channels[0], proj_channels[1], 3, relu=False)
# Fix the highway input if necessary
if proj_channels[-1] != channels:
self.highway_mismatch = True
self.pre_highway = nn.Linear(proj_channels[-1], channels, bias=False)
else:
self.highway_mismatch = False
self.highways = nn.ModuleList()
for i in range(num_highways):
hn = HighwayNetwork(channels)
self.highways.append(hn)
self.rnn = nn.GRU(channels, channels, batch_first=True, bidirectional=True)
self._to_flatten.append(self.rnn)
# Avoid fragmentation of RNN parameters and associated warning
self._flatten_parameters()
def forward(self, x):
# Although we `_flatten_parameters()` on init, when using DataParallel
# the model gets replicated, making it no longer guaranteed that the
# weights are contiguous in GPU memory. Hence, we must call it again
self._flatten_parameters()
# Save these for later
residual = x
seq_len = x.size(-1)
conv_bank = []
# Convolution Bank
for conv in self.conv1d_bank:
c = conv(x) # Convolution
conv_bank.append(c[:, :, :seq_len])
# Stack along the channel axis
conv_bank = torch.cat(conv_bank, dim=1)
# dump the last padding to fit residual
x = self.maxpool(conv_bank)[:, :, :seq_len]
# Conv1d projections
x = self.conv_project1(x)
x = self.conv_project2(x)
# Residual Connect
x = x + residual
# Through the highways
x = x.transpose(1, 2)
if self.highway_mismatch is True:
x = self.pre_highway(x)
for h in self.highways:
x = h(x)
# And then the RNN
x, _ = self.rnn(x)
return x
def _flatten_parameters(self):
"""Calls `flatten_parameters` on all the rnns used by the WaveRNN. Used
to improve efficiency and avoid PyTorch yelling at us."""
[m.flatten_parameters() for m in self._to_flatten]
class TacotronEncoder(nn.Module):
def __init__(self, embed_dims, num_chars, cbhg_channels, K, num_highways, dropout):
super().__init__()
self.embedding = nn.Embedding(num_chars, embed_dims)
self.pre_net = PreNet(embed_dims, embed_dims, embed_dims, dropout=dropout)
self.cbhg = CBHG(K=K, in_channels=cbhg_channels, channels=cbhg_channels,
proj_channels=[cbhg_channels, cbhg_channels],
num_highways=num_highways)
self.proj_out = nn.Linear(cbhg_channels * 2, cbhg_channels)
def forward(self, x):
x = self.embedding(x)
x = self.pre_net(x)
x.transpose_(1, 2)
x = self.cbhg(x)
x = self.proj_out(x)
return x
class RNNEncoder(nn.Module):
def __init__(self, num_chars, embedding_dim, n_convolutions=3, kernel_size=5):
super(RNNEncoder, self).__init__()
self.embedding = nn.Embedding(num_chars, embedding_dim, padding_idx=0)
convolutions = []
for _ in range(n_convolutions):
conv_layer = nn.Sequential(
ConvNorm(embedding_dim,
embedding_dim,
kernel_size=kernel_size, stride=1,
padding=int((kernel_size - 1) / 2),
dilation=1, w_init_gain='relu'),
nn.BatchNorm1d(embedding_dim))
convolutions.append(conv_layer)
self.convolutions = nn.ModuleList(convolutions)
self.lstm = nn.LSTM(embedding_dim, int(embedding_dim / 2), 1,
batch_first=True, bidirectional=True)
def forward(self, x):
input_lengths = (x > 0).sum(-1)
input_lengths = input_lengths.cpu().numpy()
x = self.embedding(x)
x = x.transpose(1, 2) # [B, H, T]
for conv in self.convolutions:
x = F.dropout(F.relu(conv(x)), 0.5, self.training) + x
x = x.transpose(1, 2) # [B, T, H]
# pytorch tensor are not reversible, hence the conversion
x = nn.utils.rnn.pack_padded_sequence(x, input_lengths, batch_first=True, enforce_sorted=False)
self.lstm.flatten_parameters()
outputs, _ = self.lstm(x)
outputs, _ = nn.utils.rnn.pad_packed_sequence(outputs, batch_first=True)
return outputs
class DecoderRNN(torch.nn.Module):
def __init__(self, hidden_size, decoder_rnn_dim, dropout):
super(DecoderRNN, self).__init__()
self.in_conv1d = nn.Sequential(
torch.nn.Conv1d(
in_channels=hidden_size,
out_channels=hidden_size,
kernel_size=9, padding=4,
),
torch.nn.ReLU(),
torch.nn.Conv1d(
in_channels=hidden_size,
out_channels=hidden_size,
kernel_size=9, padding=4,
),
)
self.ln = nn.LayerNorm(hidden_size)
if decoder_rnn_dim == 0:
decoder_rnn_dim = hidden_size * 2
self.rnn = torch.nn.LSTM(
input_size=hidden_size,
hidden_size=decoder_rnn_dim,
num_layers=1,
batch_first=True,
bidirectional=True,
dropout=dropout
)
self.rnn.flatten_parameters()
self.conv1d = torch.nn.Conv1d(
in_channels=decoder_rnn_dim * 2,
out_channels=hidden_size,
kernel_size=3,
padding=1,
)
def forward(self, x):
input_masks = x.abs().sum(-1).ne(0).data[:, :, None]
input_lengths = input_masks.sum([-1, -2])
input_lengths = input_lengths.cpu().numpy()
x = self.in_conv1d(x.transpose(1, 2)).transpose(1, 2)
x = self.ln(x)
x = nn.utils.rnn.pack_padded_sequence(x, input_lengths, batch_first=True, enforce_sorted=False)
self.rnn.flatten_parameters()
x, _ = self.rnn(x) # [B, T, C]
x, _ = nn.utils.rnn.pad_packed_sequence(x, batch_first=True)
x = x * input_masks
pre_mel = self.conv1d(x.transpose(1, 2)).transpose(1, 2) # [B, T, C]
pre_mel = pre_mel * input_masks
return pre_mel
@@ -0,0 +1,751 @@
import math
import torch
from torch import nn
from torch.nn import Parameter, Linear
from .layers import LayerNorm, Embedding
from ...utils.nn.seq_utils import (
get_incremental_state,
set_incremental_state,
softmax,
make_positions,
)
import torch.nn.functional as F
DEFAULT_MAX_SOURCE_POSITIONS = 2000
DEFAULT_MAX_TARGET_POSITIONS = 2000
class SinusoidalPositionalEmbedding(nn.Module):
"""This module produces sinusoidal positional embeddings of any length.
Padding symbols are ignored.
"""
def __init__(self, embedding_dim, padding_idx, init_size=1024):
super().__init__()
self.embedding_dim = embedding_dim
self.padding_idx = padding_idx
self.weights = SinusoidalPositionalEmbedding.get_embedding(
init_size,
embedding_dim,
padding_idx,
)
self.register_buffer('_float_tensor', torch.FloatTensor(1))
@staticmethod
def get_embedding(num_embeddings, embedding_dim, padding_idx=None):
"""Build sinusoidal embeddings.
This matches the implementation in tensor2tensor, but differs slightly
from the description in Section 3.5 of "Attention Is All You Need".
"""
half_dim = embedding_dim // 2
emb = math.log(10000) / (half_dim - 1)
emb = torch.exp(torch.arange(half_dim, dtype=torch.float) * -emb)
emb = torch.arange(num_embeddings, dtype=torch.float).unsqueeze(1) * emb.unsqueeze(0)
emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=1).view(num_embeddings, -1)
if embedding_dim % 2 == 1:
# zero pad
emb = torch.cat([emb, torch.zeros(num_embeddings, 1)], dim=1)
if padding_idx is not None:
emb[padding_idx, :] = 0
return emb
def forward(self, input, incremental_state=None, timestep=None, positions=None, **kwargs):
"""Input is expected to be of size [bsz x seqlen]."""
bsz, seq_len = input.shape[:2]
max_pos = self.padding_idx + 1 + seq_len
if self.weights is None or max_pos > self.weights.size(0):
# recompute/expand embeddings if needed
self.weights = SinusoidalPositionalEmbedding.get_embedding(
max_pos,
self.embedding_dim,
self.padding_idx,
)
self.weights = self.weights.to(self._float_tensor)
if incremental_state is not None:
# positions is the same for every token when decoding a single step
pos = timestep.view(-1)[0] + 1 if timestep is not None else seq_len
return self.weights[self.padding_idx + pos, :].expand(bsz, 1, -1)
positions = make_positions(input, self.padding_idx) if positions is None else positions
return self.weights.index_select(0, positions.view(-1)).view(bsz, seq_len, -1).detach()
def max_positions(self):
"""Maximum number of supported positions."""
return int(1e5) # an arbitrary large number
class TransformerFFNLayer(nn.Module):
def __init__(self, hidden_size, filter_size, padding="SAME", kernel_size=1, dropout=0., act='gelu'):
super().__init__()
self.kernel_size = kernel_size
self.dropout = dropout
self.act = act
if padding == 'SAME':
self.ffn_1 = nn.Conv1d(hidden_size, filter_size, kernel_size, padding=kernel_size // 2)
elif padding == 'LEFT':
self.ffn_1 = nn.Sequential(
nn.ConstantPad1d((kernel_size - 1, 0), 0.0),
nn.Conv1d(hidden_size, filter_size, kernel_size)
)
self.ffn_2 = Linear(filter_size, hidden_size)
def forward(self, x, incremental_state=None):
# x: T x B x C
if incremental_state is not None:
saved_state = self._get_input_buffer(incremental_state)
if 'prev_input' in saved_state:
prev_input = saved_state['prev_input']
x = torch.cat((prev_input, x), dim=0)
x = x[-self.kernel_size:]
saved_state['prev_input'] = x
self._set_input_buffer(incremental_state, saved_state)
x = self.ffn_1(x.permute(1, 2, 0)).permute(2, 0, 1)
x = x * self.kernel_size ** -0.5
if incremental_state is not None:
x = x[-1:]
if self.act == 'gelu':
x = F.gelu(x)
if self.act == 'relu':
x = F.relu(x)
x = F.dropout(x, self.dropout, training=self.training)
x = self.ffn_2(x)
return x
def _get_input_buffer(self, incremental_state):
return get_incremental_state(
self,
incremental_state,
'f',
) or {}
def _set_input_buffer(self, incremental_state, buffer):
set_incremental_state(
self,
incremental_state,
'f',
buffer,
)
def clear_buffer(self, incremental_state):
if incremental_state is not None:
saved_state = self._get_input_buffer(incremental_state)
if 'prev_input' in saved_state:
del saved_state['prev_input']
self._set_input_buffer(incremental_state, saved_state)
class MultiheadAttention(nn.Module):
def __init__(self, embed_dim, num_heads, kdim=None, vdim=None, dropout=0., bias=True,
add_bias_kv=False, add_zero_attn=False, self_attention=False,
encoder_decoder_attention=False):
super().__init__()
self.embed_dim = embed_dim
self.kdim = kdim if kdim is not None else embed_dim
self.vdim = vdim if vdim is not None else embed_dim
self.qkv_same_dim = self.kdim == embed_dim and self.vdim == embed_dim
self.num_heads = num_heads
self.dropout = dropout
self.head_dim = embed_dim // num_heads
assert self.head_dim * num_heads == self.embed_dim, "embed_dim must be divisible by num_heads"
self.scaling = self.head_dim ** -0.5
self.self_attention = self_attention
self.encoder_decoder_attention = encoder_decoder_attention
assert not self.self_attention or self.qkv_same_dim, 'Self-attention requires query, key and ' \
'value to be of the same size'
if self.qkv_same_dim:
self.in_proj_weight = Parameter(torch.Tensor(3 * embed_dim, embed_dim))
else:
self.k_proj_weight = Parameter(torch.Tensor(embed_dim, self.kdim))
self.v_proj_weight = Parameter(torch.Tensor(embed_dim, self.vdim))
self.q_proj_weight = Parameter(torch.Tensor(embed_dim, embed_dim))
if bias:
self.in_proj_bias = Parameter(torch.Tensor(3 * embed_dim))
else:
self.register_parameter('in_proj_bias', None)
self.out_proj = nn.Linear(embed_dim, embed_dim, bias=bias)
if add_bias_kv:
self.bias_k = Parameter(torch.Tensor(1, 1, embed_dim))
self.bias_v = Parameter(torch.Tensor(1, 1, embed_dim))
else:
self.bias_k = self.bias_v = None
self.add_zero_attn = add_zero_attn
self.reset_parameters()
self.enable_torch_version = False
if hasattr(F, "multi_head_attention_forward"):
self.enable_torch_version = True
else:
self.enable_torch_version = False
self.last_attn_probs = None
def reset_parameters(self):
if self.qkv_same_dim:
nn.init.xavier_uniform_(self.in_proj_weight)
else:
nn.init.xavier_uniform_(self.k_proj_weight)
nn.init.xavier_uniform_(self.v_proj_weight)
nn.init.xavier_uniform_(self.q_proj_weight)
nn.init.xavier_uniform_(self.out_proj.weight)
if self.in_proj_bias is not None:
nn.init.constant_(self.in_proj_bias, 0.)
nn.init.constant_(self.out_proj.bias, 0.)
if self.bias_k is not None:
nn.init.xavier_normal_(self.bias_k)
if self.bias_v is not None:
nn.init.xavier_normal_(self.bias_v)
def forward(
self,
query, key, value,
key_padding_mask=None,
incremental_state=None,
need_weights=True,
static_kv=False,
attn_mask=None,
before_softmax=False,
need_head_weights=False,
enc_dec_attn_constraint_mask=None,
reset_attn_weight=None
):
"""Input shape: Time x Batch x Channel
Args:
key_padding_mask (ByteTensor, optional): mask to exclude
keys that are pads, of shape `(batch, src_len)`, where
padding elements are indicated by 1s.
need_weights (bool, optional): return the attention weights,
averaged over heads (default: False).
attn_mask (ByteTensor, optional): typically used to
implement causal attention, where the mask prevents the
attention from looking forward in time (default: None).
before_softmax (bool, optional): return the raw attention
weights and values before the attention softmax.
need_head_weights (bool, optional): return the attention
weights for each head. Implies *need_weights*. Default:
return the average attention weights over all heads.
"""
if need_head_weights:
need_weights = True
tgt_len, bsz, embed_dim = query.size()
assert embed_dim == self.embed_dim
assert list(query.size()) == [tgt_len, bsz, embed_dim]
if self.enable_torch_version and incremental_state is None and not static_kv and reset_attn_weight is None:
if self.qkv_same_dim:
return F.multi_head_attention_forward(query, key, value,
self.embed_dim, self.num_heads,
self.in_proj_weight,
self.in_proj_bias, self.bias_k, self.bias_v,
self.add_zero_attn, self.dropout,
self.out_proj.weight, self.out_proj.bias,
self.training, key_padding_mask, need_weights,
attn_mask)
else:
return F.multi_head_attention_forward(query, key, value,
self.embed_dim, self.num_heads,
torch.empty([0]),
self.in_proj_bias, self.bias_k, self.bias_v,
self.add_zero_attn, self.dropout,
self.out_proj.weight, self.out_proj.bias,
self.training, key_padding_mask, need_weights,
attn_mask, use_separate_proj_weight=True,
q_proj_weight=self.q_proj_weight,
k_proj_weight=self.k_proj_weight,
v_proj_weight=self.v_proj_weight)
if incremental_state is not None:
saved_state = self._get_input_buffer(incremental_state)
if 'prev_key' in saved_state:
# previous time steps are cached - no need to recompute
# key and value if they are static
if static_kv:
assert self.encoder_decoder_attention and not self.self_attention
key = value = None
else:
saved_state = None
if self.self_attention:
# self-attention
q, k, v = self.in_proj_qkv(query)
elif self.encoder_decoder_attention:
# encoder-decoder attention
q = self.in_proj_q(query)
if key is None:
assert value is None
k = v = None
else:
k = self.in_proj_k(key)
v = self.in_proj_v(key)
else:
q = self.in_proj_q(query)
k = self.in_proj_k(key)
v = self.in_proj_v(value)
q *= self.scaling
if self.bias_k is not None:
assert self.bias_v is not None
k = torch.cat([k, self.bias_k.repeat(1, bsz, 1)])
v = torch.cat([v, self.bias_v.repeat(1, bsz, 1)])
if attn_mask is not None:
attn_mask = torch.cat([attn_mask, attn_mask.new_zeros(attn_mask.size(0), 1)], dim=1)
if key_padding_mask is not None:
key_padding_mask = torch.cat(
[key_padding_mask, key_padding_mask.new_zeros(key_padding_mask.size(0), 1)], dim=1)
q = q.contiguous().view(tgt_len, bsz * self.num_heads, self.head_dim).transpose(0, 1)
if k is not None:
k = k.contiguous().view(-1, bsz * self.num_heads, self.head_dim).transpose(0, 1)
if v is not None:
v = v.contiguous().view(-1, bsz * self.num_heads, self.head_dim).transpose(0, 1)
if saved_state is not None:
# saved states are stored with shape (bsz, num_heads, seq_len, head_dim)
if 'prev_key' in saved_state:
prev_key = saved_state['prev_key'].view(bsz * self.num_heads, -1, self.head_dim)
if static_kv:
k = prev_key
else:
k = torch.cat((prev_key, k), dim=1)
if 'prev_value' in saved_state:
prev_value = saved_state['prev_value'].view(bsz * self.num_heads, -1, self.head_dim)
if static_kv:
v = prev_value
else:
v = torch.cat((prev_value, v), dim=1)
if 'prev_key_padding_mask' in saved_state and saved_state['prev_key_padding_mask'] is not None:
prev_key_padding_mask = saved_state['prev_key_padding_mask']
if static_kv:
key_padding_mask = prev_key_padding_mask
else:
key_padding_mask = torch.cat((prev_key_padding_mask, key_padding_mask), dim=1)
saved_state['prev_key'] = k.view(bsz, self.num_heads, -1, self.head_dim)
saved_state['prev_value'] = v.view(bsz, self.num_heads, -1, self.head_dim)
saved_state['prev_key_padding_mask'] = key_padding_mask
self._set_input_buffer(incremental_state, saved_state)
src_len = k.size(1)
# This is part of a workaround to get around fork/join parallelism
# not supporting Optional types.
if key_padding_mask is not None and key_padding_mask.shape == torch.Size([]):
key_padding_mask = None
if key_padding_mask is not None:
assert key_padding_mask.size(0) == bsz
assert key_padding_mask.size(1) == src_len
if self.add_zero_attn:
src_len += 1
k = torch.cat([k, k.new_zeros((k.size(0), 1) + k.size()[2:])], dim=1)
v = torch.cat([v, v.new_zeros((v.size(0), 1) + v.size()[2:])], dim=1)
if attn_mask is not None:
attn_mask = torch.cat([attn_mask, attn_mask.new_zeros(attn_mask.size(0), 1)], dim=1)
if key_padding_mask is not None:
key_padding_mask = torch.cat(
[key_padding_mask, torch.zeros(key_padding_mask.size(0), 1).type_as(key_padding_mask)], dim=1)
attn_weights = torch.bmm(q, k.transpose(1, 2))
attn_weights = self.apply_sparse_mask(attn_weights, tgt_len, src_len, bsz)
assert list(attn_weights.size()) == [bsz * self.num_heads, tgt_len, src_len]
if attn_mask is not None:
if len(attn_mask.shape) == 2:
attn_mask = attn_mask.unsqueeze(0)
elif len(attn_mask.shape) == 3:
attn_mask = attn_mask[:, None].repeat([1, self.num_heads, 1, 1]).reshape(
bsz * self.num_heads, tgt_len, src_len)
attn_weights = attn_weights + attn_mask
if enc_dec_attn_constraint_mask is not None: # bs x head x L_kv
attn_weights = attn_weights.view(bsz, self.num_heads, tgt_len, src_len)
attn_weights = attn_weights.masked_fill(
enc_dec_attn_constraint_mask.unsqueeze(2).bool(),
-1e8,
)
attn_weights = attn_weights.view(bsz * self.num_heads, tgt_len, src_len)
if key_padding_mask is not None:
# don't attend to padding symbols
attn_weights = attn_weights.view(bsz, self.num_heads, tgt_len, src_len)
attn_weights = attn_weights.masked_fill(
key_padding_mask.unsqueeze(1).unsqueeze(2),
-1e8,
)
attn_weights = attn_weights.view(bsz * self.num_heads, tgt_len, src_len)
attn_logits = attn_weights.view(bsz, self.num_heads, tgt_len, src_len)
if before_softmax:
return attn_weights, v
attn_weights_float = softmax(attn_weights, dim=-1)
attn_weights = attn_weights_float.type_as(attn_weights)
attn_probs = F.dropout(attn_weights_float.type_as(attn_weights), p=self.dropout, training=self.training)
if reset_attn_weight is not None:
if reset_attn_weight:
self.last_attn_probs = attn_probs.detach()
else:
assert self.last_attn_probs is not None
attn_probs = self.last_attn_probs
attn = torch.bmm(attn_probs, v)
assert list(attn.size()) == [bsz * self.num_heads, tgt_len, self.head_dim]
attn = attn.transpose(0, 1).contiguous().view(tgt_len, bsz, embed_dim)
attn = self.out_proj(attn)
if need_weights:
attn_weights = attn_weights_float.view(bsz, self.num_heads, tgt_len, src_len).transpose(1, 0)
if not need_head_weights:
# average attention weights over heads
attn_weights = attn_weights.mean(dim=0)
else:
attn_weights = None
return attn, (attn_weights, attn_logits)
def in_proj_qkv(self, query):
return self._in_proj(query).chunk(3, dim=-1)
def in_proj_q(self, query):
if self.qkv_same_dim:
return self._in_proj(query, end=self.embed_dim)
else:
bias = self.in_proj_bias
if bias is not None:
bias = bias[:self.embed_dim]
return F.linear(query, self.q_proj_weight, bias)
def in_proj_k(self, key):
if self.qkv_same_dim:
return self._in_proj(key, start=self.embed_dim, end=2 * self.embed_dim)
else:
weight = self.k_proj_weight
bias = self.in_proj_bias
if bias is not None:
bias = bias[self.embed_dim:2 * self.embed_dim]
return F.linear(key, weight, bias)
def in_proj_v(self, value):
if self.qkv_same_dim:
return self._in_proj(value, start=2 * self.embed_dim)
else:
weight = self.v_proj_weight
bias = self.in_proj_bias
if bias is not None:
bias = bias[2 * self.embed_dim:]
return F.linear(value, weight, bias)
def _in_proj(self, input, start=0, end=None):
weight = self.in_proj_weight
bias = self.in_proj_bias
weight = weight[start:end, :]
if bias is not None:
bias = bias[start:end]
return F.linear(input, weight, bias)
def _get_input_buffer(self, incremental_state):
return get_incremental_state(
self,
incremental_state,
'attn_state',
) or {}
def _set_input_buffer(self, incremental_state, buffer):
set_incremental_state(
self,
incremental_state,
'attn_state',
buffer,
)
def apply_sparse_mask(self, attn_weights, tgt_len, src_len, bsz):
return attn_weights
def clear_buffer(self, incremental_state=None):
if incremental_state is not None:
saved_state = self._get_input_buffer(incremental_state)
if 'prev_key' in saved_state:
del saved_state['prev_key']
if 'prev_value' in saved_state:
del saved_state['prev_value']
self._set_input_buffer(incremental_state, saved_state)
class EncSALayer(nn.Module):
def __init__(self, c, num_heads, dropout, attention_dropout=0.1,
relu_dropout=0.1, kernel_size=9, padding='SAME', act='gelu'):
super().__init__()
self.c = c
self.dropout = dropout
self.num_heads = num_heads
if num_heads > 0:
self.layer_norm1 = LayerNorm(c)
self.self_attn = MultiheadAttention(
self.c, num_heads, self_attention=True, dropout=attention_dropout, bias=False)
self.layer_norm2 = LayerNorm(c)
self.ffn = TransformerFFNLayer(
c, 4 * c, kernel_size=kernel_size, dropout=relu_dropout, padding=padding, act=act)
def forward(self, x, encoder_padding_mask=None, **kwargs):
layer_norm_training = kwargs.get('layer_norm_training', None)
if layer_norm_training is not None:
self.layer_norm1.training = layer_norm_training
self.layer_norm2.training = layer_norm_training
if self.num_heads > 0:
residual = x
x = self.layer_norm1(x)
x, _, = self.self_attn(
query=x,
key=x,
value=x,
key_padding_mask=encoder_padding_mask
)
x = F.dropout(x, self.dropout, training=self.training)
x = residual + x
x = x * (1 - encoder_padding_mask.float()).transpose(0, 1)[..., None]
residual = x
x = self.layer_norm2(x)
x = self.ffn(x)
x = F.dropout(x, self.dropout, training=self.training)
x = residual + x
x = x * (1 - encoder_padding_mask.float()).transpose(0, 1)[..., None]
return x
class DecSALayer(nn.Module):
def __init__(self, c, num_heads, dropout, attention_dropout=0.1, relu_dropout=0.1,
kernel_size=9, act='gelu'):
super().__init__()
self.c = c
self.dropout = dropout
self.layer_norm1 = LayerNorm(c)
self.self_attn = MultiheadAttention(
c, num_heads, self_attention=True, dropout=attention_dropout, bias=False
)
self.layer_norm2 = LayerNorm(c)
self.encoder_attn = MultiheadAttention(
c, num_heads, encoder_decoder_attention=True, dropout=attention_dropout, bias=False,
)
self.layer_norm3 = LayerNorm(c)
self.ffn = TransformerFFNLayer(
c, 4 * c, padding='LEFT', kernel_size=kernel_size, dropout=relu_dropout, act=act)
def forward(
self,
x,
encoder_out=None,
encoder_padding_mask=None,
incremental_state=None,
self_attn_mask=None,
self_attn_padding_mask=None,
attn_out=None,
reset_attn_weight=None,
**kwargs,
):
layer_norm_training = kwargs.get('layer_norm_training', None)
if layer_norm_training is not None:
self.layer_norm1.training = layer_norm_training
self.layer_norm2.training = layer_norm_training
self.layer_norm3.training = layer_norm_training
residual = x
x = self.layer_norm1(x)
x, _ = self.self_attn(
query=x,
key=x,
value=x,
key_padding_mask=self_attn_padding_mask,
incremental_state=incremental_state,
attn_mask=self_attn_mask
)
x = F.dropout(x, self.dropout, training=self.training)
x = residual + x
attn_logits = None
if encoder_out is not None or attn_out is not None:
residual = x
x = self.layer_norm2(x)
if encoder_out is not None:
x, attn = self.encoder_attn(
query=x,
key=encoder_out,
value=encoder_out,
key_padding_mask=encoder_padding_mask,
incremental_state=incremental_state,
static_kv=True,
enc_dec_attn_constraint_mask=get_incremental_state(self, incremental_state,
'enc_dec_attn_constraint_mask'),
reset_attn_weight=reset_attn_weight
)
attn_logits = attn[1]
elif attn_out is not None:
x = self.encoder_attn.in_proj_v(attn_out)
if encoder_out is not None or attn_out is not None:
x = F.dropout(x, self.dropout, training=self.training)
x = residual + x
residual = x
x = self.layer_norm3(x)
x = self.ffn(x, incremental_state=incremental_state)
x = F.dropout(x, self.dropout, training=self.training)
x = residual + x
return x, attn_logits
def clear_buffer(self, input, encoder_out=None, encoder_padding_mask=None, incremental_state=None):
self.encoder_attn.clear_buffer(incremental_state)
self.ffn.clear_buffer(incremental_state)
def set_buffer(self, name, tensor, incremental_state):
return set_incremental_state(self, incremental_state, name, tensor)
class TransformerEncoderLayer(nn.Module):
def __init__(self, hidden_size, dropout, kernel_size=9, num_heads=2):
super().__init__()
self.hidden_size = hidden_size
self.dropout = dropout
self.num_heads = num_heads
self.op = EncSALayer(
hidden_size, num_heads, dropout=dropout,
attention_dropout=0.0, relu_dropout=dropout,
kernel_size=kernel_size)
def forward(self, x, **kwargs):
return self.op(x, **kwargs)
class TransformerDecoderLayer(nn.Module):
def __init__(self, hidden_size, dropout, kernel_size=9, num_heads=2):
super().__init__()
self.hidden_size = hidden_size
self.dropout = dropout
self.num_heads = num_heads
self.op = DecSALayer(
hidden_size, num_heads, dropout=dropout,
attention_dropout=0.0, relu_dropout=dropout,
kernel_size=kernel_size)
def forward(self, x, **kwargs):
return self.op(x, **kwargs)
def clear_buffer(self, *args):
return self.op.clear_buffer(*args)
def set_buffer(self, *args):
return self.op.set_buffer(*args)
class FFTBlocks(nn.Module):
def __init__(self, hidden_size, num_layers, ffn_kernel_size=9, dropout=0.0,
num_heads=2, use_pos_embed=True, use_last_norm=True,
use_pos_embed_alpha=True):
super().__init__()
self.num_layers = num_layers
embed_dim = self.hidden_size = hidden_size
self.dropout = dropout
self.use_pos_embed = use_pos_embed
self.use_last_norm = use_last_norm
if use_pos_embed:
self.max_source_positions = DEFAULT_MAX_TARGET_POSITIONS
self.padding_idx = 0
self.pos_embed_alpha = nn.Parameter(torch.Tensor([1])) if use_pos_embed_alpha else 1
self.embed_positions = SinusoidalPositionalEmbedding(
embed_dim, self.padding_idx, init_size=DEFAULT_MAX_TARGET_POSITIONS,
)
self.layers = nn.ModuleList([])
self.layers.extend([
TransformerEncoderLayer(self.hidden_size, self.dropout,
kernel_size=ffn_kernel_size, num_heads=num_heads)
for _ in range(self.num_layers)
])
if self.use_last_norm:
self.layer_norm = nn.LayerNorm(embed_dim)
else:
self.layer_norm = None
def forward(self, x, padding_mask=None, attn_mask=None, return_hiddens=False):
"""
:param x: [B, T, C]
:param padding_mask: [B, T]
:return: [B, T, C] or [L, B, T, C]
"""
padding_mask = x.abs().sum(-1).eq(0).data if padding_mask is None else padding_mask
nonpadding_mask_TB = 1 - padding_mask.transpose(0, 1).float()[:, :, None] # [T, B, 1]
if self.use_pos_embed:
positions = self.pos_embed_alpha * self.embed_positions(x[..., 0])
x = x + positions
x = F.dropout(x, p=self.dropout, training=self.training)
# B x T x C -> T x B x C
x = x.transpose(0, 1) * nonpadding_mask_TB
hiddens = []
for layer in self.layers:
x = layer(x, encoder_padding_mask=padding_mask, attn_mask=attn_mask) * nonpadding_mask_TB
hiddens.append(x)
if self.use_last_norm:
x = self.layer_norm(x) * nonpadding_mask_TB
if return_hiddens:
x = torch.stack(hiddens, 0) # [L, T, B, C]
x = x.transpose(1, 2) # [L, B, T, C]
else:
x = x.transpose(0, 1) # [B, T, C]
return x
class FastSpeechEncoder(FFTBlocks):
def __init__(self, dict_size, hidden_size=256, num_layers=4, kernel_size=9, num_heads=2,
dropout=0.0):
super().__init__(hidden_size, num_layers, kernel_size, num_heads=num_heads,
use_pos_embed=False, dropout=dropout) # use_pos_embed_alpha for compatibility
self.embed_tokens = Embedding(dict_size, hidden_size, 0)
self.embed_scale = math.sqrt(hidden_size)
self.padding_idx = 0
self.embed_positions = SinusoidalPositionalEmbedding(
hidden_size, self.padding_idx, init_size=DEFAULT_MAX_TARGET_POSITIONS,
)
def forward(self, txt_tokens, attn_mask=None):
"""
:param txt_tokens: [B, T]
:return: {
'encoder_out': [B x T x C]
}
"""
encoder_padding_mask = txt_tokens.eq(self.padding_idx).data
x = self.forward_embedding(txt_tokens) # [B, T, H]
if self.num_layers > 0:
x = super(FastSpeechEncoder, self).forward(x, encoder_padding_mask, attn_mask=attn_mask)
return x
def forward_embedding(self, txt_tokens):
# embed tokens and positions
x = self.embed_scale * self.embed_tokens(txt_tokens)
positions = self.embed_positions(txt_tokens)
x = x + positions
x = F.dropout(x, p=self.dropout, training=self.training)
return x
class FastSpeechDecoder(FFTBlocks):
def __init__(self, hidden_size=256, num_layers=4, kernel_size=9, num_heads=2):
super().__init__(hidden_size, num_layers, kernel_size, num_heads=num_heads)
@@ -0,0 +1,109 @@
import torch
from torch import nn
from packaging import version
def fused_add_tanh_sigmoid_multiply(input_a, input_b, n_channels):
n_channels_int = n_channels[0]
in_act = input_a + input_b
t_act = torch.tanh(in_act[:, :n_channels_int, :])
s_act = torch.sigmoid(in_act[:, n_channels_int:, :])
acts = t_act * s_act
return acts
jit_fused_add_tanh_sigmoid_multiply = fused_add_tanh_sigmoid_multiply
def script_function():
if version.parse(torch.__version__) >= version.parse('2.0'):
global jit_fused_add_tanh_sigmoid_multiply
jit_fused_add_tanh_sigmoid_multiply = torch.jit.script(fused_add_tanh_sigmoid_multiply)
class WN(torch.nn.Module):
def __init__(self, hidden_size, kernel_size, dilation_rate, n_layers, c_cond=0,
p_dropout=0, share_cond_layers=False, is_BTC=False):
super(WN, self).__init__()
assert (kernel_size % 2 == 1)
assert (hidden_size % 2 == 0)
self.is_BTC = is_BTC
self.hidden_size = hidden_size
self.kernel_size = kernel_size
self.dilation_rate = dilation_rate
self.n_layers = n_layers
self.gin_channels = c_cond
self.p_dropout = p_dropout
self.share_cond_layers = share_cond_layers
self.in_layers = torch.nn.ModuleList()
self.res_skip_layers = torch.nn.ModuleList()
self.drop = nn.Dropout(p_dropout)
if c_cond != 0 and not share_cond_layers:
cond_layer = torch.nn.Conv1d(c_cond, 2 * hidden_size * n_layers, 1)
self.cond_layer = torch.nn.utils.weight_norm(cond_layer, name='weight')
for i in range(n_layers):
dilation = dilation_rate ** i
padding = int((kernel_size * dilation - dilation) / 2)
in_layer = torch.nn.Conv1d(hidden_size, 2 * hidden_size, kernel_size,
dilation=dilation, padding=padding)
in_layer = torch.nn.utils.weight_norm(in_layer, name='weight')
self.in_layers.append(in_layer)
# last one is not necessary
if i < n_layers - 1:
res_skip_channels = 2 * hidden_size
else:
res_skip_channels = hidden_size
res_skip_layer = torch.nn.Conv1d(hidden_size, res_skip_channels, 1)
res_skip_layer = torch.nn.utils.weight_norm(res_skip_layer, name='weight')
self.res_skip_layers.append(res_skip_layer)
script_function()
def forward(self, x, nonpadding=None, cond=None):
if self.is_BTC:
x = x.transpose(1, 2)
cond = cond.transpose(1, 2) if cond is not None else None
nonpadding = nonpadding.transpose(1, 2) if nonpadding is not None else None
if nonpadding is None:
nonpadding = 1
output = torch.zeros_like(x)
n_channels_tensor = torch.IntTensor([self.hidden_size])
if cond is not None and not self.share_cond_layers:
cond = self.cond_layer(cond)
for i in range(self.n_layers):
x_in = self.in_layers[i](x)
x_in = self.drop(x_in)
if cond is not None:
cond_offset = i * 2 * self.hidden_size
cond_l = cond[:, cond_offset:cond_offset + 2 * self.hidden_size, :]
else:
cond_l = torch.zeros_like(x_in)
if version.parse(torch.__version__) >= version.parse('2.0'):
acts = jit_fused_add_tanh_sigmoid_multiply(x_in, cond_l, n_channels_tensor)
else:
acts = fused_add_tanh_sigmoid_multiply(x_in, cond_l, n_channels_tensor)
res_skip_acts = self.res_skip_layers[i](acts)
if i < self.n_layers - 1:
x = (x + res_skip_acts[:, :self.hidden_size, :]) * nonpadding
output = output + res_skip_acts[:, self.hidden_size:, :]
else:
output = output + res_skip_acts
output = output * nonpadding
if self.is_BTC:
output = output.transpose(1, 2)
return output
def remove_weight_norm(self):
def remove_weight_norm(m):
try:
nn.utils.remove_weight_norm(m)
except ValueError: # this module didn't have weight norm
return
self.apply(remove_weight_norm)
@@ -0,0 +1 @@
"""Pitch extractor modules for ROSVOT."""
@@ -0,0 +1,6 @@
from .constants import *
from .model import E2E0
from .utils import to_local_average_f0, to_viterbi_f0
from .inference import RMVPE
from .spec import MelSpectrogram
from .extractor import extract
@@ -0,0 +1,9 @@
SAMPLE_RATE = 16000
N_CLASS = 360
N_MELS = 128
MEL_FMIN = 30
MEL_FMAX = 8000
WINDOW_LENGTH = 1024
CONST = 1997.3794084376191
@@ -0,0 +1,173 @@
import torch
import torch.nn as nn
from .constants import N_MELS
class ConvBlockRes(nn.Module):
def __init__(self, in_channels, out_channels, momentum=0.01):
super(ConvBlockRes, self).__init__()
self.conv = nn.Sequential(
nn.Conv2d(in_channels=in_channels,
out_channels=out_channels,
kernel_size=(3, 3),
stride=(1, 1),
padding=(1, 1),
bias=False),
nn.BatchNorm2d(out_channels, momentum=momentum),
nn.ReLU(),
nn.Conv2d(in_channels=out_channels,
out_channels=out_channels,
kernel_size=(3, 3),
stride=(1, 1),
padding=(1, 1),
bias=False),
nn.BatchNorm2d(out_channels, momentum=momentum),
nn.ReLU(),
)
if in_channels != out_channels:
self.shortcut = nn.Conv2d(in_channels, out_channels, (1, 1))
self.is_shortcut = True
else:
self.is_shortcut = False
def forward(self, x):
if self.is_shortcut:
return self.conv(x) + self.shortcut(x)
else:
return self.conv(x) + x
class ResEncoderBlock(nn.Module):
def __init__(self, in_channels, out_channels, kernel_size, n_blocks=1, momentum=0.01):
super(ResEncoderBlock, self).__init__()
self.n_blocks = n_blocks
self.conv = nn.ModuleList()
self.conv.append(ConvBlockRes(in_channels, out_channels, momentum))
for i in range(n_blocks - 1):
self.conv.append(ConvBlockRes(out_channels, out_channels, momentum))
self.kernel_size = kernel_size
if self.kernel_size is not None:
self.pool = nn.AvgPool2d(kernel_size=kernel_size)
def forward(self, x):
for i in range(self.n_blocks):
x = self.conv[i](x)
if self.kernel_size is not None:
return x, self.pool(x)
else:
return x
class ResDecoderBlock(nn.Module):
def __init__(self, in_channels, out_channels, stride, n_blocks=1, momentum=0.01):
super(ResDecoderBlock, self).__init__()
out_padding = (0, 1) if stride == (1, 2) else (1, 1)
self.n_blocks = n_blocks
self.conv1 = nn.Sequential(
nn.ConvTranspose2d(in_channels=in_channels,
out_channels=out_channels,
kernel_size=(3, 3),
stride=stride,
padding=(1, 1),
output_padding=out_padding,
bias=False),
nn.BatchNorm2d(out_channels, momentum=momentum),
nn.ReLU(),
)
self.conv2 = nn.ModuleList()
self.conv2.append(ConvBlockRes(out_channels * 2, out_channels, momentum))
for i in range(n_blocks-1):
self.conv2.append(ConvBlockRes(out_channels, out_channels, momentum))
def forward(self, x, concat_tensor):
x = self.conv1(x)
x = torch.cat((x, concat_tensor), dim=1)
for i in range(self.n_blocks):
x = self.conv2[i](x)
return x
class Encoder(nn.Module):
def __init__(self, in_channels, in_size, n_encoders, kernel_size, n_blocks, out_channels=16, momentum=0.01):
super(Encoder, self).__init__()
self.n_encoders = n_encoders
self.bn = nn.BatchNorm2d(in_channels, momentum=momentum)
self.layers = nn.ModuleList()
self.latent_channels = []
for i in range(self.n_encoders):
self.layers.append(ResEncoderBlock(in_channels, out_channels, kernel_size, n_blocks, momentum=momentum))
self.latent_channels.append([out_channels, in_size])
in_channels = out_channels
out_channels *= 2
in_size //= 2
self.out_size = in_size
self.out_channel = out_channels
def forward(self, x):
concat_tensors = []
x = self.bn(x)
for i in range(self.n_encoders):
_, x = self.layers[i](x)
concat_tensors.append(_)
return x, concat_tensors
class Intermediate(nn.Module):
def __init__(self, in_channels, out_channels, n_inters, n_blocks, momentum=0.01):
super(Intermediate, self).__init__()
self.n_inters = n_inters
self.layers = nn.ModuleList()
self.layers.append(ResEncoderBlock(in_channels, out_channels, None, n_blocks, momentum))
for i in range(self.n_inters-1):
self.layers.append(ResEncoderBlock(out_channels, out_channels, None, n_blocks, momentum))
def forward(self, x):
for i in range(self.n_inters):
x = self.layers[i](x)
return x
class Decoder(nn.Module):
def __init__(self, in_channels, n_decoders, stride, n_blocks, momentum=0.01):
super(Decoder, self).__init__()
self.layers = nn.ModuleList()
self.n_decoders = n_decoders
for i in range(self.n_decoders):
out_channels = in_channels // 2
self.layers.append(ResDecoderBlock(in_channels, out_channels, stride, n_blocks, momentum))
in_channels = out_channels
def forward(self, x, concat_tensors):
for i in range(self.n_decoders):
x = self.layers[i](x, concat_tensors[-1-i])
return x
class TimbreFilter(nn.Module):
def __init__(self, latent_rep_channels):
super(TimbreFilter, self).__init__()
self.layers = nn.ModuleList()
for latent_rep in latent_rep_channels:
self.layers.append(ConvBlockRes(latent_rep[0], latent_rep[0]))
def forward(self, x_tensors):
out_tensors = []
for i, layer in enumerate(self.layers):
out_tensors.append(layer(x_tensors[i]))
return out_tensors
class DeepUnet0(nn.Module):
def __init__(self, kernel_size, n_blocks, en_de_layers=5, inter_layers=4, in_channels=1, en_out_channels=16):
super(DeepUnet0, self).__init__()
self.encoder = Encoder(in_channels, N_MELS, en_de_layers, kernel_size, n_blocks, en_out_channels)
self.intermediate = Intermediate(self.encoder.out_channel // 2, self.encoder.out_channel, inter_layers, n_blocks)
self.tf = TimbreFilter(self.encoder.latent_channels)
self.decoder = Decoder(self.encoder.out_channel, en_de_layers, kernel_size, n_blocks)
def forward(self, x):
x, concat_tensors = self.encoder(x)
x = self.intermediate(x)
x = self.decoder(x, concat_tensors)
return x
@@ -0,0 +1,183 @@
import math
import os
from tqdm import tqdm
import librosa
import numpy as np
import torch
from torch.utils.data import Dataset, DataLoader, DistributedSampler
import torch.multiprocessing as mp
from torch.distributed import init_process_group
import torch.distributed as dist
from .inference import RMVPE
from ....utils.commons.dataset_utils import batch_by_size, build_dataloader
# import utils
from ....utils.audio import get_wav_num_frames
"""
A convenient API for batch inference
update: add ddp
"""
class RMVPEInferDataset(Dataset):
def __init__(self, wav_fns: list, id_and_sizes=None, sr=24000, hop_size=128, num_workers=0):
if id_and_sizes is None:
id_and_sizes = []
if type(wav_fns[0]) == str: # wav_paths
for idx, wav_path in enumerate(wav_fns):
total_frames = get_wav_num_frames(wav_path, sr)
id_and_sizes.append((idx, round(total_frames / hop_size)))
else: # numpy arrays, mono wavs
for idx, wav in enumerate(wav_fns):
id_and_sizes.append((idx, round(wav.shape[-1] / hop_size)))
self.wav_fns = wav_fns
self.id_and_sizes = id_and_sizes
self.sr = sr
self.num_workers = num_workers
def __getitem__(self, idx):
if type(self.wav_fns[idx]) == str:
wav_fn = self.wav_fns[idx]
wav, _ = librosa.core.load(wav_fn, sr=self.sr)
else:
wav = self.wav_fns[idx]
return idx, wav
def collater(self, samples: list):
return samples
def __len__(self):
return len(self.wav_fns)
def ordered_indices(self):
"""Return an ordered list of indices. Batches will be constructed based
on this order."""
return np.arange(len(self))
def num_tokens(self, index):
return self.id_and_sizes[index][1]
@torch.no_grad()
def extract(wav_fns: list, id_and_sizes=None, ckpt=None, sr=24000, hop_size=128, bsz=128, max_tokens=100000,
fmax=900, fmin=50, ds_workers=0):
all_gpu_ids = [int(x) for x in os.environ.get("CUDA_VISIBLE_DEVICES", "").split(",") if x != '']
num_gpus = len(all_gpu_ids)
dist_config = {
"dist_backend": "nccl",
"dist_url": "tcp://localhost:54189",
"world_size": 1
}
# https://discuss.pytorch.org/t/how-to-fix-a-sigsegv-in-pytorch-when-using-distributed-training-e-g-ddp/113518/10#:~:text=Using%20start%20and%20join%20avoids
# https://github.com/pytorch/pytorch/issues/40403#issuecomment-648515174
# mp.set_start_method('spawn')
if num_gpus > 1:
result_queue = mp.Queue()
for rank in range(num_gpus):
mp.Process(target=extract_worker, args=(rank, wav_fns, id_and_sizes, ckpt, sr, hop_size, bsz, max_tokens, fmax,
fmin, dist_config, num_gpus, ds_workers, result_queue,)).start()
f0_res = [None] * len(wav_fns)
for _ in range(num_gpus):
f0_res_dict = result_queue.get()
for idx in f0_res_dict:
f0_res[idx] = f0_res_dict[idx]
del f0_res_dict
else:
# f0_res = extract_one_process(wav_fns, id_and_sizes, ckpt, sr, hop_size, bsz, max_tokens, fmax, fmin)
f0_res_dict = extract_worker(0, wav_fns, id_and_sizes, ckpt, sr, hop_size, bsz, max_tokens, fmax,
fmin, dist_config, num_gpus, ds_workers, None)
f0_res = [None] * len(wav_fns)
for idx in f0_res_dict:
f0_res[idx] = f0_res_dict[idx]
return f0_res
@torch.no_grad()
def extract_worker(rank, wav_fns: list, id_and_sizes=None, ckpt=None, sr=24000, hop_size=128, bsz=128, max_tokens=100000,
fmax=900, fmin=50, dist_config=None, num_gpus=1, ds_workers=0, q=None):
# print(f"rank: {rank}")
if num_gpus > 1:
init_process_group(backend=dist_config['dist_backend'], init_method=dist_config['dist_url'],
world_size=dist_config['world_size'] * num_gpus, rank=rank)
dataset = RMVPEInferDataset(wav_fns, id_and_sizes, sr, hop_size, num_workers=ds_workers)
# ds_sampler = DistributedSampler(dataset, shuffle=False) if num_gpus > 1 else None
# loader = DataLoader(dataset, sampler=ds_sampler, collate_fn=dataset.collator, batch_size=1, num_workers=40, drop_last=False)
loader = build_dataloader(dataset, shuffle=False, max_tokens=max_tokens, max_sentences=bsz, use_ddp=num_gpus > 1)
loader = tqdm(loader, desc=f'| Processing f0 in [n_ranks={num_gpus}; max_tokens={max_tokens}; max_sentences={bsz}]') if rank == 0 else loader
device = torch.device(f"cuda:{int(rank)}")
model = RMVPE(ckpt, device=device)
f0_res_dict = {}
for batch in loader:
if batch is None or len(batch) == 0:
continue
idxs = [item[0] for item in batch]
wavs = [item[1] for item in batch]
lengths = [(wav.shape[0] + hop_size - 1) // hop_size for wav in wavs]
with torch.no_grad():
f0s, uvs = model.get_pitch_batch(
wavs, sample_rate=sr,
hop_size=hop_size,
lengths=lengths,
fmax=fmax,
fmin=fmin
)
for i, idx in enumerate(idxs):
f0_res_dict[idx] = f0s[i]
if q is not None:
q.put(f0_res_dict)
else:
return f0_res_dict
# old version
def extract_one_process(wav_fns: list, id_and_sizes=None, ckpt=None, sr=24000, hop_size=128, bsz=128, max_tokens=100000,
fmax=900, fmin=50, device='cuda'):
assert ckpt is not None
rmvpe = RMVPE(ckpt, device=device)
if id_and_sizes is None:
id_and_sizes = []
if type(wav_fns[0]) == str: # wav_paths
for idx, wav_path in enumerate(wav_fns):
total_frames = get_wav_num_frames(wav_path, sr)
id_and_sizes.append((idx, round(total_frames / hop_size)))
else: # numpy arrays, mono wavs
for idx, wav in enumerate(wav_fns):
id_and_sizes.append((idx, round(wav.shape[-1] / hop_size)))
get_size = lambda x: x[1]
bs = batch_by_size(id_and_sizes, get_size, max_tokens=max_tokens, max_sentences=bsz)
for i in range(len(bs)):
bs[i] = [bs[i][j][0] for j in range(len(bs[i]))]
f0_res = [None] * len(wav_fns)
for batch in tqdm(bs, total=len(bs), desc=f'| Processing f0 in [max_tokens={max_tokens}; max_sentences={bsz}]'):
wavs, mel_lengths, lengths = [], [], []
for idx in batch:
if type(wav_fns[idx]) == str:
wav_fn = wav_fns[idx]
wav, _ = librosa.core.load(wav_fn, sr=sr)
else:
wav = wav_fns[idx]
wavs.append(wav)
mel_lengths.append(math.ceil((wav.shape[0] + 1) / hop_size))
lengths.append((wav.shape[0] + hop_size - 1) // hop_size)
with torch.no_grad():
f0s, uvs = rmvpe.get_pitch_batch(
wavs, sample_rate=sr,
hop_size=hop_size,
lengths=lengths,
fmax=fmax,
fmin=fmin
)
for i, idx in enumerate(batch):
f0_res[idx] = f0s[i]
if rmvpe is not None:
rmvpe.release_cuda()
torch.cuda.empty_cache()
rmvpe = None
return f0_res
@@ -0,0 +1,134 @@
import math
import numpy as np
import torch
import torch.nn.functional as F
from torchaudio.transforms import Resample
import pyworld as pw
from ....utils.audio.pitch_utils import interp_f0, resample_align_curve
from .constants import *
from .model import E2E0
from .spec import MelSpectrogram
from .utils import to_local_average_f0, to_viterbi_f0
class RMVPE:
def __init__(self, model_path, hop_length=160, device=None):
self.resample_kernel = {}
if device is None:
self.device = 'cuda' if torch.cuda.is_available() else 'cpu'
else:
self.device = device
self.model = E2E0(4, 1, (2, 2)).eval().to(self.device)
ckpt = torch.load(model_path, map_location=self.device)
self.model.load_state_dict(ckpt['model'], strict=False)
self.mel_extractor = MelSpectrogram(
N_MELS, SAMPLE_RATE, WINDOW_LENGTH, hop_length, None, MEL_FMIN, MEL_FMAX
).to(self.device)
self.hop_length = hop_length
@torch.no_grad()
def mel2hidden(self, mel):
n_frames = mel.shape[-1]
mel = F.pad(mel, (0, 32 * ((n_frames - 1) // 32 + 1) - n_frames), mode='constant')
hidden = self.model(mel)
return hidden[:, :n_frames]
def decode(self, hidden, thred=0.03, use_viterbi=False):
if use_viterbi:
f0 = to_viterbi_f0(hidden, thred=thred)
else:
f0 = to_local_average_f0(hidden, thred=thred)
return f0
def postprocess(self, f0, fmin=50, fmax=1000, audio=None, min_gap=2):
if audio is not None:
# this doesn't work. deprecated
t = np.arange(0, f0.shape[0] * self.hop_length / 16000, self.hop_length / 16000)
f0 = pw.stonemask(audio.astype(np.float64), f0.astype(np.float64), t, 16000).astype(float)
f0[f0 < fmin] = 0
f0[f0 > fmax] = 0
# eliminate glitch
# min_gap: if successive positive f0 positions < min_gap, zero these positions
# eg: if min_gap=2, [0, 500, 500, 0] => [0, 0, 0, 0]
for idx in range(f0.shape[0] - min_gap - 1):
if f0[idx] == 0 and f0[idx + min_gap + 1] == 0 and np.sum(f0[idx: idx + min_gap + 2]) > 0:
f0[idx: idx + min_gap + 2] = 0
return f0
def infer_from_audio(self, audio, sample_rate=16000, thred=0.03, use_viterbi=False):
audio = torch.from_numpy(audio).float().unsqueeze(0).to(self.device)
if sample_rate == 16000:
audio_res = audio
else:
key_str = str(sample_rate)
if key_str not in self.resample_kernel:
self.resample_kernel[key_str] = Resample(sample_rate, 16000, lowpass_filter_width=128)
self.resample_kernel[key_str] = self.resample_kernel[key_str].to(self.device)
audio_res = self.resample_kernel[key_str](audio)
mel = self.mel_extractor(audio_res, center=True)
hidden = self.mel2hidden(mel)
f0 = self.decode(hidden, thred=thred, use_viterbi=use_viterbi).squeeze(0)
return f0
def get_pitch(self, waveform, sample_rate, hop_size, length, interp_uv=False, fmin=50, fmax=1000):
f0 = self.infer_from_audio(waveform, sample_rate=sample_rate)
f0 = self.postprocess(f0, fmin, fmax)
uv = f0 == 0
time_step = hop_size / sample_rate
f0_res = resample_align_curve(f0, 0.01, time_step, length)
uv_res = resample_align_curve(uv.astype(np.float32), 0.01, time_step, length) > 0.5
if not interp_uv:
f0_res[uv_res] = 0
return f0_res, uv_res
def infer_from_audio_batch(self, audios, sample_rate=16000, thred=0.03, use_viterbi=False):
from ....utils.commons.dataset_utils import collate_1d_or_2d
if isinstance(audios, list):
audios = [torch.from_numpy(audio).float() for audio in audios]
sizes = [math.ceil((audio.shape[0] + 1) / self.hop_length) for audio in audios]
audios = collate_1d_or_2d(audios, 0.0).to(self.device)
elif isinstance(audios, torch.Tensor):
sizes = None
if audios.device != self.device:
audios = audios.to(self.device)
else:
raise NotImplementedError
if sample_rate == 16000:
audios_res = audios
else:
key_str = str(sample_rate)
if key_str not in self.resample_kernel:
self.resample_kernel[key_str] = Resample(sample_rate, 16000, lowpass_filter_width=128)
self.resample_kernel[key_str] = self.resample_kernel[key_str].to(self.device)
audios_res = self.resample_kernel[key_str](audios)
mels = self.mel_extractor(audios_res, center=True)
hiddens = self.mel2hidden(mels)
f0 = self.decode(hiddens, thred=thred, use_viterbi=use_viterbi)
f0s = []
for i in range(f0.shape[0]):
f = f0[i, :sizes[i]] if sizes is not None else f0[i, :]
f0s.append(f)
return f0s
def get_pitch_batch(self, waveforms, sample_rate, hop_size, lengths, interp_uv=False, fmin=50, fmax=1000):
# hop_size, sample_rate: tgt params
f0s = self.infer_from_audio_batch(waveforms, sample_rate=sample_rate)
f0s_res, uvs_res = [], []
for idx, f0 in enumerate(f0s):
f0 = self.postprocess(f0, fmin, fmax, min_gap=6)
uv = f0 == 0
length = lengths[idx]
time_step = hop_size / sample_rate
f0_res = resample_align_curve(f0, 0.01, time_step, length)
uv_res = resample_align_curve(uv.astype(np.float32), 0.01, time_step, length) > 0.5
if not interp_uv:
f0_res[uv_res] = 0
f0s_res.append(f0_res)
uvs_res.append(uv_res)
return f0s_res, uvs_res
def release_cuda(self):
self.model = self.model.cpu()
self.mel_extractor = self.mel_extractor.cpu()
@@ -0,0 +1,32 @@
from torch import nn
from .constants import *
from .deepunet import DeepUnet0
from .seq import BiGRU
class E2E0(nn.Module):
def __init__(self, n_blocks, n_gru, kernel_size, en_de_layers=5, inter_layers=4, in_channels=1,
en_out_channels=16):
super(E2E0, self).__init__()
self.unet = DeepUnet0(kernel_size, n_blocks, en_de_layers, inter_layers, in_channels, en_out_channels)
self.cnn = nn.Conv2d(en_out_channels, 3, (3, 3), padding=(1, 1))
if n_gru:
self.fc = nn.Sequential(
BiGRU(3 * N_MELS, 256, n_gru),
nn.Linear(512, N_CLASS),
nn.Dropout(0.25),
nn.Sigmoid()
)
else:
self.fc = nn.Sequential(
nn.Linear(3 * N_MELS, N_CLASS),
nn.Dropout(0.25),
nn.Sigmoid()
)
def forward(self, mel):
mel = mel.transpose(-1, -2).unsqueeze(1)
x = self.cnn(self.unet(mel)).transpose(1, 2).flatten(-2)
x = self.fc(x)
return x
@@ -0,0 +1,10 @@
import torch.nn as nn
class BiGRU(nn.Module):
def __init__(self, input_features, hidden_features, num_layers):
super(BiGRU, self).__init__()
self.gru = nn.GRU(input_features, hidden_features, num_layers=num_layers, batch_first=True, bidirectional=True)
def forward(self, x):
return self.gru(x)[0]
@@ -0,0 +1,72 @@
import torch
import numpy as np
import torch.nn.functional as F
from librosa.filters import mel
class MelSpectrogram(torch.nn.Module):
def __init__(
self,
n_mel_channels,
sampling_rate,
win_length,
hop_length,
n_fft=None,
mel_fmin=0,
mel_fmax=None,
clamp=1e-5
):
super().__init__()
n_fft = win_length if n_fft is None else n_fft
self.hann_window = {}
mel_basis = mel(
sr=sampling_rate,
n_fft=n_fft,
n_mels=n_mel_channels,
fmin=mel_fmin,
fmax=mel_fmax,
htk=True)
mel_basis = torch.from_numpy(mel_basis).float()
self.register_buffer("mel_basis", mel_basis)
self.n_fft = win_length if n_fft is None else n_fft
self.hop_length = hop_length
self.win_length = win_length
self.sampling_rate = sampling_rate
self.n_mel_channels = n_mel_channels
self.clamp = clamp
def forward(self, audio, keyshift=0, speed=1, center=True):
factor = 2 ** (keyshift / 12)
n_fft_new = int(np.round(self.n_fft * factor))
win_length_new = int(np.round(self.win_length * factor))
hop_length_new = int(np.round(self.hop_length * speed))
keyshift_key = str(keyshift) + '_' + str(audio.device)
if keyshift_key not in self.hann_window:
self.hann_window[keyshift_key] = torch.hann_window(win_length_new).to(audio.device)
if center:
pad_left = win_length_new // 2
pad_right = (win_length_new + 1) // 2
audio = F.pad(audio, (pad_left, pad_right))
fft = torch.stft(
audio,
n_fft=n_fft_new,
hop_length=hop_length_new,
win_length=win_length_new,
window=self.hann_window[keyshift_key],
center=False,
return_complex=True
)
magnitude = fft.abs()
if keyshift != 0:
size = self.n_fft // 2 + 1
resize = magnitude.size(1)
if resize < size:
magnitude = F.pad(magnitude, (0, 0, 0, size - resize))
magnitude = magnitude[:, :size, :] * self.win_length / win_length_new
mel_output = torch.matmul(self.mel_basis, magnitude)
log_mel_spec = torch.log(torch.clamp(mel_output, min=self.clamp))
return log_mel_spec
@@ -0,0 +1,43 @@
import librosa
import numpy as np
import torch
from .constants import *
def to_local_average_f0(hidden, center=None, thred=0.03):
idx = torch.arange(N_CLASS, device=hidden.device)[None, None, :] # [B=1, T=1, N]
idx_cents = idx * 20 + CONST # [B=1, N]
if center is None:
center = torch.argmax(hidden, dim=2, keepdim=True) # [B, T, 1]
start = torch.clip(center - 4, min=0) # [B, T, 1]
end = torch.clip(center + 5, max=N_CLASS) # [B, T, 1]
idx_mask = (idx >= start) & (idx < end) # [B, T, N]
weights = hidden * idx_mask # [B, T, N]
product_sum = torch.sum(weights * idx_cents, dim=2) # [B, T]
weight_sum = torch.sum(weights, dim=2) # [B, T]
cents = product_sum / (weight_sum + (weight_sum == 0)) # avoid dividing by zero, [B, T]
f0 = 10 * 2 ** (cents / 1200)
uv = hidden.max(dim=2)[0] < thred # [B, T]
f0 = f0 * ~uv
return f0.cpu().numpy()
def to_viterbi_f0(hidden, thred=0.03):
# Create viterbi transition matrix
if not hasattr(to_viterbi_f0, 'transition'):
xx, yy = np.meshgrid(range(N_CLASS), range(N_CLASS))
transition = np.maximum(30 - abs(xx - yy), 0)
transition = transition / transition.sum(axis=1, keepdims=True)
to_viterbi_f0.transition = transition
# Convert to probability
prob = hidden.squeeze(0).cpu().numpy()
prob = prob.T
prob = prob / prob.sum(axis=0)
# Perform viterbi decoding
path = librosa.sequence.viterbi(prob, to_viterbi_f0.transition).astype(np.int64)
center = torch.from_numpy(path).unsqueeze(0).unsqueeze(-1).to(hidden.device)
return to_local_average_f0(hidden, center=center, thred=thred)
@@ -0,0 +1 @@
"""Core ROSVOT model components."""
@@ -0,0 +1,295 @@
from copy import deepcopy
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
from ...utils.commons.hparams import hparams
from ...utils.commons.gpu_mem_track import MemTracker
from ..commons.layers import Embedding
from ..commons.conv import ResidualBlock, ConvBlocks
from ..commons.conformer.conformer import ConformerLayers
from .unet import Unet
def regulate_boundary(bd_logits, threshold, min_gap=18, ref_bd=None, ref_bd_min_gap=8, non_padding=None):
# this doesn't preserve gradient
device = bd_logits.device
bd_logits = torch.sigmoid(bd_logits).data.cpu()
# bd_logits[0] = bd_logits[-1] = 1e-5 # avoid itv invalid problem
bd = (bd_logits > threshold).long()
bd_res = torch.zeros_like(bd).long()
for i in range(bd.shape[0]):
bd_i = bd[i]
last_bd_idx = -1
start = -1
for j in range(bd_i.shape[0]):
if bd_i[j] == 1:
if 0 <= start < j:
continue
elif start < 0:
start = j
else:
if 0 <= start < j:
if j - 1 > start:
bd_idx = start + int(torch.argmax(bd_logits[i, start: j]).item())
else:
bd_idx = start
if bd_idx - last_bd_idx < min_gap and last_bd_idx > 0:
bd_idx = round((bd_idx + last_bd_idx) / 2)
bd_res[i, last_bd_idx] = 0
bd_res[i, bd_idx] = 1
last_bd_idx = bd_idx
start = -1
# assert ref_bd_min_gap <= min_gap // 2
if ref_bd is not None and ref_bd_min_gap > 0:
ref = ref_bd.data.cpu()
for i in range(bd_res.shape[0]):
ref_bd_i = ref[i]
ref_bd_i_js = []
for j in range(ref_bd_i.shape[0]):
if ref_bd_i[j] == 1:
ref_bd_i_js.append(j)
seg_sum = torch.sum(bd_res[i, max(0, j - ref_bd_min_gap): j + ref_bd_min_gap])
if seg_sum == 0:
bd_res[i, j] = 1
elif seg_sum == 1 and bd_res[i, j] != 1:
bd_res[i, max(0, j - ref_bd_min_gap): j + ref_bd_min_gap] = \
ref_bd_i[max(0, j - ref_bd_min_gap): j + ref_bd_min_gap]
elif seg_sum > 1:
for k in range(1, ref_bd_min_gap+1):
if bd_res[i, max(0, j - k)] == 1 and ref_bd_i[max(0, j - k)] != 1:
bd_res[i, max(0, j - k)] = 0
break
if bd_res[i, min(bd_res.shape[1] - 1, j + k)] == 1 and ref_bd_i[min(bd_res.shape[1] - 1, j + k)] != 1:
bd_res[i, min(bd_res.shape[1] - 1, j + k)] = 0
break
bd_res[i, j] = 1
# final check
assert torch.sum(bd_res[i, ref_bd_i_js]) == len(ref_bd_i_js), \
f"{torch.sum(bd_res[i, ref_bd_i_js])} {len(ref_bd_i_js)}"
bd_res = bd_res.to(device)
# force valid begin and end
bd_res[:, 0] = 0
if non_padding is not None:
for i in range(bd_res.shape[0]):
bd_res[i, sum(non_padding[i]) - 1:] = 0
else:
bd_res[:, -1] = 0
return bd_res
class BackboneNet(nn.Module):
def __init__(self, hparams):
super().__init__()
self.hidden_size = hidden_size = hparams['hidden_size']
self.dropout = hparams.get('dropout', 0.0)
updown_rates = [2, 2, 2]
channel_multiples = [1, 1, 1]
if hparams.get('updown_rates', None) is not None:
updown_rates = [int(i) for i in hparams.get('updown_rates', None).split('-')]
if hparams.get('channel_multiples', None) is not None:
channel_multiples = [float(i) for i in hparams.get('channel_multiples', None).split('-')]
assert len(updown_rates) == len(channel_multiples)
# convs
if hparams.get('bkb_net', 'conv') == 'conv':
self.net = Unet(hidden_size, down_layers=len(updown_rates), mid_layers=hparams.get('bkb_layers', 12),
up_layers=len(updown_rates), kernel_size=3, updown_rates=updown_rates,
channel_multiples=channel_multiples, dropout=0, is_BTC=True,
constant_channels=False, mid_net=None, use_skip_layer=hparams.get('unet_skip_layer', False))
# conformer
elif hparams.get('bkb_net', 'conv') == 'conformer':
mid_net = ConformerLayers(
hidden_size, num_layers=hparams.get('bkb_layers', 12), kernel_size=hparams.get('conformer_kernel', 9),
dropout=self.dropout, num_heads=4)
self.net = Unet(hidden_size, down_layers=len(updown_rates), up_layers=len(updown_rates), kernel_size=3,
updown_rates=updown_rates, channel_multiples=channel_multiples, dropout=0,
is_BTC=True, constant_channels=False, mid_net=mid_net,
use_skip_layer=hparams.get('unet_skip_layer', False))
def forward(self, x):
return self.net(x)
class PitchDecoder(nn.Module):
def __init__(self, hparams):
super().__init__()
self.hidden_size = hidden_size = hparams['hidden_size']
self.dropout = hparams.get('dropout', 0.0)
self.note_bd_out = nn.Linear(hidden_size, 1)
self.note_bd_temperature = max(1e-7, hparams.get('note_bd_temperature', 1.0))
# note prediction
self.pitch_attn_num_head = hparams.get('pitch_attn_num_head', 1)
self.multihead_dot_attn = nn.Linear(hidden_size, self.pitch_attn_num_head)
self.post = ConvBlocks(hidden_size, out_dims=hidden_size, dilations=None, kernel_size=3,
layers_in_block=1, c_multiple=1, dropout=self.dropout, num_layers=1,
post_net_kernel=3, act_type='leakyrelu')
self.pitch_out = nn.Linear(hidden_size, hparams.get('note_num', 100) + 4)
self.note_num = hparams.get('note_num', 100)
self.note_start = hparams.get('note_start', 30)
self.pitch_temperature = max(1e-7, hparams.get('note_pitch_temperature', 1.0))
def forward(self, feat, note_bd, train=True):
bsz, T, _ = feat.shape
attn = torch.sigmoid(self.multihead_dot_attn(feat)) # [B, T, C] -> [B, T, num_head]
attn = F.dropout(attn, self.dropout, train)
attn_feat = feat.unsqueeze(3) * attn.unsqueeze(2) # [B, T, C, 1] x [B, T, 1, num_head] -> [B, T, C, num_head]
attn_feat = torch.mean(attn_feat, dim=-1) # [B, T, C, num_head] -> [B, T, C]
mel2note = torch.cumsum(note_bd, 1)
note_length = torch.max(torch.sum(note_bd, dim=1)).item() + 1 # max length
note_lengths = torch.sum(note_bd, dim=1) + 1 # [B]
# print('note_length', note_length)
attn = torch.mean(attn, dim=-1, keepdim=True) # [B, T, num_head] -> [B, T, 1]
denom = mel2note.new_zeros(bsz, note_length, dtype=attn.dtype).scatter_add_(
dim=1, index=mel2note, src=attn.squeeze(-1)
) # [B, T] -> [B, note_length] count the note frames of each note (with padding excluded)
frame2note = mel2note.unsqueeze(-1).repeat(1, 1, self.hidden_size) # [B, T] -> [B, T, C], with padding included
note_aggregate = frame2note.new_zeros(bsz, note_length, self.hidden_size, dtype=attn_feat.dtype).scatter_add_(
dim=1, index=frame2note, src=attn_feat
) # [B, T, C] -> [B, note_length, C]
note_aggregate = note_aggregate / (denom.unsqueeze(-1) + 1e-5)
note_aggregate = F.dropout(note_aggregate, self.dropout, train)
note_logits = self.post(note_aggregate)
note_logits = self.pitch_out(note_logits) / self.pitch_temperature
# note_logits = torch.clamp(note_logits, min=-16., max=16.) # don't know need it or not
note_pred = torch.softmax(note_logits, dim=-1) # [B, note_length, note_num]
note_pred = torch.argmax(note_pred, dim=-1) # [B, note_length]
# for some reason, note idx maybe 130 (why?)
note_pred[note_pred > self.note_num] = 0
note_pred[note_pred < self.note_start] = 0
return note_lengths, note_logits, note_pred
class MidiExtractor(nn.Module):
def __init__(self, hparams):
super(MidiExtractor, self).__init__()
self.hparams = deepcopy(hparams)
self.hidden_size = hidden_size = hparams['hidden_size']
self.dropout = hparams.get('dropout', 0.0)
self.note_bd_threshold = hparams.get('note_bd_threshold', 0.5)
self.note_bd_min_gap = round(hparams.get('note_bd_min_gap', 100) * hparams['audio_sample_rate'] / 1000 / hparams['hop_size'])
self.note_bd_ref_min_gap = round(hparams.get('note_bd_ref_min_gap', 50) * hparams['audio_sample_rate'] / 1000 / hparams['hop_size'])
self.mel_proj = nn.Conv1d(hparams['use_mel_bins'], hidden_size, kernel_size=3, padding=1)
self.mel_encoder = ConvBlocks(hidden_size, out_dims=hidden_size, dilations=None, kernel_size=3,
layers_in_block=2, c_multiple=1, dropout=self.dropout, num_layers=1,
post_net_kernel=3, act_type='leakyrelu')
self.use_pitch = hparams.get('use_pitch_embed', True)
if self.use_pitch:
self.pitch_embed = Embedding(300, hidden_size, 0, 'kaiming')
self.uv_embed = Embedding(3, hidden_size, 0, 'kaiming')
self.use_wbd = hparams.get('use_wbd', True)
if self.use_wbd:
self.word_bd_embed = Embedding(3, hidden_size, 0, 'kaiming')
self.cond_encoder = ConvBlocks(hidden_size, out_dims=hidden_size, dilations=None, kernel_size=3,
layers_in_block=1, c_multiple=1, dropout=self.dropout, num_layers=1,
post_net_kernel=3, act_type='leakyrelu')
# backbone
self.net = BackboneNet(hparams)
# note bd prediction
self.note_bd_out = nn.Linear(hidden_size, 1)
self.note_bd_temperature = max(1e-7, hparams.get('note_bd_temperature', 1.0))
# note prediction
self.pitch_decoder = PitchDecoder(hparams)
self.reset_parameters()
def run_encoder(self, mel=None, word_bd=None, pitch=None, uv=None, non_padding=None):
mel_embed = self.mel_proj(mel.transpose(1, 2)).transpose(1, 2)
mel_embed = self.mel_encoder(mel_embed)
pitch_embed = word_bd_embed = 0
if self.use_pitch and pitch is not None and uv is not None:
pitch_embed = self.pitch_embed(pitch) + self.uv_embed(uv) # [B, T, C]
if self.use_wbd and word_bd is not None:
word_bd_embed = self.word_bd_embed(word_bd)
feat = self.cond_encoder(mel_embed + pitch_embed + word_bd_embed)
return feat
def forward(self, mel=None, word_bd=None, note_bd=None, pitch=None, uv=None, non_padding=None, train=True):
ret = {}
bsz, T, _ = mel.shape
feat = self.run_encoder(mel, word_bd, pitch, uv, non_padding)
feat = self.net(feat) # [B, T, C]
# note bd prediction
note_bd_logits = self.note_bd_out(F.dropout(feat, self.dropout, train)).squeeze(-1) / self.note_bd_temperature
note_bd_logits = torch.clamp(note_bd_logits, min=-16., max=16.)
ret['note_bd_logits'] = note_bd_logits # [B, T]
if note_bd is None or not train:
note_bd = regulate_boundary(note_bd_logits, self.note_bd_threshold, self.note_bd_min_gap,
word_bd, self.note_bd_ref_min_gap, non_padding)
ret['note_bd_pred'] = note_bd # [B, T]
# note pitch prediction
note_lengths, note_logits, note_pred = self.pitch_decoder(feat, note_bd, train)
ret['note_lengths'], ret['note_logits'], ret['note_pred'] = note_lengths, note_logits, note_pred
return ret
def reset_parameters(self):
nn.init.kaiming_normal_(self.pitch_decoder.multihead_dot_attn.weight, mode='fan_in')
nn.init.kaiming_normal_(self.note_bd_out.weight, mode='fan_in')
nn.init.kaiming_normal_(self.pitch_decoder.pitch_out.weight, mode='fan_in')
nn.init.kaiming_normal_(self.mel_proj.weight, mode='fan_in')
nn.init.constant_(self.pitch_decoder.multihead_dot_attn.bias, 0.0)
nn.init.constant_(self.note_bd_out.bias, 0.0)
nn.init.constant_(self.pitch_decoder.pitch_out.bias, 0.0)
class WordbdExtractor(MidiExtractor):
def __init__(self, hparams):
super().__init__(hparams)
self.use_wbd = False
self.word_bd_embed = None
self.note_bd_out = self.note_bd_temperature = self.pitch_decoder = None
self.word_bd_threshold = hparams.get('word_bd_threshold', 0.5)
self.word_bd_min_gap = round(
hparams.get('word_bd_min_gap', 100) * hparams['audio_sample_rate'] / 1000 / hparams['hop_size'])
self.word_bd_out = nn.Linear(self.hidden_size, 1)
self.word_bd_temperature = max(1e-7, hparams.get('word_bd_temperature', 1.0))
nn.init.kaiming_normal_(self.word_bd_out.weight, mode='fan_in')
nn.init.constant_(self.word_bd_out.bias, 0.0)
def forward(self, mel=None, pitch=None, uv=None, non_padding=None, train=True):
# gpu_tracker.track()
ret = {}
bsz, T, _ = mel.shape
feat = self.run_encoder(mel=mel, pitch=pitch, uv=uv, non_padding=non_padding)
feat = self.net(feat) # [B, T, C]
word_bd_logits = self.word_bd_out(F.dropout(feat, self.dropout, train)).squeeze(-1) / self.word_bd_temperature
word_bd_logits = torch.clamp(word_bd_logits, min=-16., max=16.)
ret['word_bd_logits'] = word_bd_logits # [B, T]
if not train:
word_bd = regulate_boundary(word_bd_logits, self.word_bd_threshold, self.word_bd_min_gap,
non_padding=non_padding)
ret['word_bd_pred'] = word_bd # [B, T]
return ret
def reset_parameters(self):
if self.use_pitch:
nn.init.kaiming_normal_(self.pitch_embed.weight, mode='fan_in')
nn.init.kaiming_normal_(self.uv_embed.weight, mode='fan_in')
nn.init.kaiming_normal_(self.mel_proj.weight, mode='fan_in')
if self.use_pitch:
nn.init.constant_(self.pitch_embed.weight[self.pitch_embed.padding_idx], 0.0)
nn.init.constant_(self.uv_embed.weight[self.uv_embed.padding_idx], 0.0)
@@ -0,0 +1,172 @@
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from ..commons.layers import LayerNorm, Embedding
from ..commons.conv import ConvBlocks, ResidualBlock, get_norm_builder, get_act_builder
class UnetDown(nn.Module):
def __init__(self, hidden_size, n_layers, kernel_size, down_rates, channel_multiples=None, dropout=0.0,
is_BTC=True, constant_channels=False):
super(UnetDown, self).__init__()
assert n_layers == len(down_rates) # downs, down sample rate
down_rates = [int(i) for i in down_rates]
self.n_layers = n_layers
self.hidden_size = hidden_size
self.is_BTC = is_BTC
channel_multiples = channel_multiples if channel_multiples is not None else down_rates
self.layers = nn.ModuleList()
self.downs = nn.ModuleList()
in_channels = hidden_size
for i in range(self.n_layers):
out_channels = int(in_channels * channel_multiples[i]) if not constant_channels else in_channels
self.layers.append(nn.Sequential(
ResidualBlock(in_channels, kernel_size, dilation=1, n=1, norm_type='ln', dropout=dropout,
c_multiple=1, ln_eps=1e-5, act_type='leakyrelu'),
nn.Conv1d(in_channels, out_channels, kernel_size, padding=(kernel_size - 1) // 2),
ResidualBlock(out_channels, kernel_size, dilation=1, n=1, norm_type='ln',
dropout=dropout, c_multiple=1, ln_eps=1e-5, act_type='leakyrelu')
))
self.downs.append(nn.Sequential(
nn.AvgPool1d(down_rates[i])
))
in_channels = out_channels
self.last_norm = get_norm_builder('ln', out_channels)()
self.post_net = nn.Conv1d(out_channels, out_channels, kernel_size=kernel_size,
padding=kernel_size // 2)
def forward(self, x, **kwargs):
# x [B, T, C]
if self.is_BTC:
x = x.transpose(1, 2) # [B, C, T]
skip_xs = []
for i in range(self.n_layers):
skip_x = self.layers[i](x)
x = self.downs[i](skip_x)
if self.is_BTC:
skip_xs.append(skip_x.transpose(1, 2)) # [B, T, C]
else:
skip_xs.append(skip_x)
x = self.post_net(self.last_norm(x))
if self.is_BTC:
x = x.transpose(1, 2)
return x, skip_xs
class UnetMid(nn.Module):
def __init__(self, hidden_size, kernel_size, n_layers=None, in_dims=None, out_dims=None,
dropout=0.0, is_BTC=True, net=None):
super(UnetMid, self).__init__()
in_dims = in_dims if in_dims is not None else hidden_size
out_dims = out_dims if out_dims is not None else hidden_size
self.pre = nn.Conv1d(in_dims, hidden_size, kernel_size, padding=kernel_size // 2)
self.post = nn.Conv1d(hidden_size, out_dims, kernel_size, padding=kernel_size // 2)
self.is_BTC = is_BTC
if net is not None:
self.net = net
else:
self.net = ConvBlocks(hidden_size, out_dims=hidden_size, dilations=None, kernel_size=kernel_size,
layers_in_block=2, c_multiple=2, dropout=dropout, num_layers=n_layers,
post_net_kernel=3, act_type='leakyrelu', is_BTC=is_BTC)
def forward(self, x, cond=None, **kwargs):
# x [B, T, C]
if self.is_BTC:
x = self.pre(x.transpose(1, 2)).transpose(1, 2)
else:
x = self.pre(x)
if cond is None:
cond = 0
x = self.net(x + cond)
if self.is_BTC:
x = self.post(x.transpose(1, 2)).transpose(1, 2)
else:
x = self.post(x)
return x
class UnetUp(nn.Module):
def __init__(self, hidden_size, n_layers, kernel_size, up_rates, channel_multiples=None, dropout=0.0,
is_BTC=True, constant_channels=False, use_skip_layer=False, skip_scale=1.0):
super(UnetUp, self).__init__()
assert n_layers == len(up_rates) # this is reversed in up module, from the output to the interface with middle
up_rates = [int(i) for i in up_rates]
self.n_layers = n_layers
self.hidden_size = hidden_size
self.is_BTC = is_BTC
self.skip_scale = skip_scale
channel_multiples = channel_multiples if channel_multiples is not None else up_rates
# in_channels = int(np.cumprod(channel_multiples)[-1] * hidden_size) if not constant_channels else hidden_size
self.in_channels_lst = (np.cumprod([1] + channel_multiples) * hidden_size).astype(int) if not constant_channels \
else [hidden_size for _ in range(self.n_layers + 1)]
in_channels = self.in_channels_lst[-1]
self.ups = nn.ModuleList()
self.skip_layers = nn.ModuleList()
self.layers = nn.ModuleList()
for i in range(self.n_layers-1, -1, -1):
out_channels = self.in_channels_lst[i] if not constant_channels else in_channels
self.ups.append(nn.Sequential(
nn.ConvTranspose1d(in_channels, in_channels, kernel_size=kernel_size, stride=up_rates[i],
padding=kernel_size//2, output_padding=up_rates[i]-1),
get_norm_builder('ln', in_channels)(),
get_act_builder('leakyrelu')()
))
self.layers.append(nn.Sequential(
# ResidualBlock(in_channels*2, kernel_size, dilation=1, n=1, norm_type='ln', dropout=dropout,
# c_multiple=1, ln_eps=1e-5, act_type='leakyrelu'),
nn.Conv1d(in_channels*2, out_channels, kernel_size, padding=(kernel_size - 1) // 2),
ResidualBlock(out_channels, kernel_size, dilation=1, n=1, norm_type='ln',
dropout=dropout, c_multiple=1, ln_eps=1e-5, act_type='leakyrelu')
))
if use_skip_layer:
self.skip_layers.append(
ResidualBlock(in_channels, kernel_size, dilation=1, n=1, norm_type='ln', dropout=dropout,
c_multiple=1, ln_eps=1e-5, act_type='leakyrelu')
)
else:
self.skip_layers.append(nn.Identity())
in_channels = out_channels
self.out_channels = out_channels
self.last_norm = get_norm_builder('ln', out_channels)()
self.post_net = nn.Conv1d(out_channels, out_channels, kernel_size=kernel_size,
padding=kernel_size // 2)
def forward(self, x, skips, **kwargs):
# x [B, T, C]
if self.is_BTC:
x = x.transpose(1, 2) # [B, C, T]
for i in range(self.n_layers):
x = self.ups[i](x)
skip_x = skips[self.n_layers - i - 1] if not self.is_BTC \
else skips[self.n_layers - i - 1].transpose(1, 2) # [B, T, C] -> [B, C, T]
skip_x = self.skip_layers[i](skip_x) * self.skip_scale
x = torch.cat((x, skip_x), dim=1) # [B, C, T]
x = self.layers[i](x)
x = self.post_net(self.last_norm(x))
if self.is_BTC:
x = x.transpose(1, 2)
return x
class Unet(nn.Module):
def __init__(self, hidden_size, down_layers, up_layers, kernel_size,
updown_rates, mid_layers=None, channel_multiples=None, dropout=0.0,
is_BTC=True, constant_channels=False, mid_net=None, use_skip_layer=False, skip_scale=1.0):
super(Unet, self).__init__()
assert len(updown_rates) == down_layers == up_layers, f"{len(updown_rates)}, {down_layers}, {up_layers}"
if channel_multiples is not None:
assert len(channel_multiples) == len(updown_rates)
else:
channel_multiples = updown_rates
self.down = UnetDown(hidden_size, down_layers, kernel_size, updown_rates,
channel_multiples, dropout, is_BTC, constant_channels)
down_out_dims = int(np.cumprod(channel_multiples)[-1] * hidden_size) if not constant_channels else hidden_size
self.mid = UnetMid(hidden_size, kernel_size, mid_layers,
in_dims=down_out_dims, out_dims=down_out_dims, dropout=dropout, is_BTC=is_BTC, net=mid_net)
self.up = UnetUp(hidden_size, up_layers, kernel_size, updown_rates,
channel_multiples, dropout, is_BTC, constant_channels, use_skip_layer, skip_scale)
def forward(self, x, mid_cond=None, **kwargs):
x, skips = self.down(x)
x = self.mid(x, mid_cond)
x = self.up(x, skips)
return x
@@ -0,0 +1,15 @@
def seed_everything(seed: int, seed_cudnn=False):
import random, os
import numpy as np
import torch
random.seed(seed)
os.environ['PYTHONHASHSEED'] = str(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed(seed)
if seed_cudnn:
torch.backends.cudnn.deterministic = True
torch.backends.cudnn.benchmark = True
@@ -0,0 +1,100 @@
import librosa
import numpy as np
import wave
import soundfile as sf
def librosa_pad_lr(x, fsize, fshift, pad_sides=1):
'''compute right padding (final frame) or both sides padding (first and final frames)
'''
assert pad_sides in (1, 2)
# return int(fsize // 2)
pad = (x.shape[0] // fshift + 1) * fshift - x.shape[0]
if pad_sides == 1:
return 0, pad
else:
return pad // 2, pad // 2 + pad % 2
def amp_to_db(x):
return 20 * np.log10(np.maximum(1e-5, x))
def db_to_amp(x):
return 10.0 ** (x * 0.05)
def normalize(S, min_level_db):
return (S - min_level_db) / -min_level_db
def denormalize(D, min_level_db):
return (D * -min_level_db) + min_level_db
def librosa_wav2spec(wav_path,
fft_size=1024,
hop_size=256,
win_length=1024,
window="hann",
num_mels=80,
fmin=80,
fmax=-1,
eps=1e-6,
sample_rate=22050,
loud_norm=False,
trim_long_sil=False):
import pyloudnorm as pyln
if isinstance(wav_path, str):
if trim_long_sil:
from .vad import trim_long_silences
wav, _, _ = trim_long_silences(wav_path, sample_rate)
else:
wav, _ = librosa.core.load(wav_path, sr=sample_rate)
else:
wav = wav_path
wav_orig = np.copy(wav)
if loud_norm:
meter = pyln.Meter(sample_rate) # create BS.1770 meter
loudness = meter.integrated_loudness(wav)
wav = pyln.normalize.loudness(wav, loudness, -22.0)
if np.abs(wav).max() > 1:
wav = wav / np.abs(wav).max()
# get amplitude spectrogram
x_stft = librosa.stft(wav, n_fft=fft_size, hop_length=hop_size,
win_length=win_length, window=window, pad_mode="constant")
linear_spc = np.abs(x_stft) # (n_bins, T)
# get mel basis
fmin = 0 if fmin == -1 else fmin
fmax = sample_rate / 2 if fmax == -1 else fmax
mel_basis = librosa.filters.mel(sr=sample_rate, n_fft=fft_size, n_mels=num_mels, fmin=fmin, fmax=fmax)
# calculate mel spec
mel = mel_basis @ linear_spc
mel = np.log10(np.maximum(eps, mel)) # (n_mel_bins, T)
l_pad, r_pad = librosa_pad_lr(wav, fft_size, hop_size, 1)
wav = np.pad(wav, (l_pad, r_pad), mode='constant', constant_values=0.0)
wav = wav[:mel.shape[1] * hop_size]
# log linear spec
linear_spc = np.log10(np.maximum(eps, linear_spc))
return {'wav': wav, 'mel': mel.T, 'linear': linear_spc.T, 'mel_basis': mel_basis, 'wav_orig': wav_orig}
def get_wav_num_frames(path, sr=None):
try:
with wave.open(path, 'rb') as f:
sr_ = f.getframerate()
if sr is None:
sr = sr_
return int(f.getnframes() / (sr_ / sr))
except wave.Error:
wav_file, sr_ = sf.read(path, dtype='float32')
if sr is None:
sr = sr_
return int(len(wav_file) / (sr_ / sr))
except:
wav_file, sr_ = librosa.core.load(path, sr=sr)
return len(wav_file)
@@ -0,0 +1,90 @@
import re
import torch
import numpy as np
from ..text.text_encoder import is_sil_phoneme
def get_mel2ph(tg_fn, ph, mel, hop_size, audio_sample_rate, min_sil_duration=0):
from textgrid import TextGrid
ph_list = ph.split(" ")
itvs = TextGrid.fromFile(tg_fn)[1]
itvs_ = []
for i in range(len(itvs)):
if itvs[i].maxTime - itvs[i].minTime < min_sil_duration and i > 0 and is_sil_phoneme(itvs[i].mark):
itvs_[-1].maxTime = itvs[i].maxTime
else:
itvs_.append(itvs[i])
itvs.intervals = itvs_
itv_marks = [itv.mark for itv in itvs]
tg_len = len([x for x in itvs if not is_sil_phoneme(x.mark)])
ph_len = len([x for x in ph_list if not is_sil_phoneme(x)])
assert tg_len == ph_len, (tg_len, ph_len, itv_marks, ph_list, tg_fn)
mel2ph = np.zeros([mel.shape[0]], int)
i_itv = 0
i_ph = 0
while i_itv < len(itvs):
itv = itvs[i_itv]
ph = ph_list[i_ph]
itv_ph = itv.mark
start_frame = int(itv.minTime * audio_sample_rate / hop_size + 0.5)
end_frame = int(itv.maxTime * audio_sample_rate / hop_size + 0.5)
if is_sil_phoneme(itv_ph) and not is_sil_phoneme(ph):
mel2ph[start_frame:end_frame] = i_ph
i_itv += 1
elif not is_sil_phoneme(itv_ph) and is_sil_phoneme(ph):
i_ph += 1
else:
if not ((is_sil_phoneme(itv_ph) and is_sil_phoneme(ph)) \
or re.sub(r'\d+', '', itv_ph.lower()) == re.sub(r'\d+', '', ph.lower())):
print(f"| WARN: {tg_fn} phs are not same: ", itv_ph, ph, itv_marks, ph_list)
mel2ph[start_frame:end_frame] = i_ph + 1
i_ph += 1
i_itv += 1
mel2ph[-1] = mel2ph[-2]
assert not np.any(mel2ph == 0)
T_t = len(ph_list)
dur = mel2token_to_dur(mel2ph, T_t)
return mel2ph.tolist(), dur.tolist()
def split_audio_by_mel2ph(audio, mel2ph, hop_size, audio_num_mel_bins):
if isinstance(audio, torch.Tensor):
audio = audio.numpy()
if isinstance(mel2ph, torch.Tensor):
mel2ph = mel2ph.numpy()
assert len(audio.shape) == 1, len(mel2ph.shape) == 1
split_locs = []
for i in range(1, len(mel2ph)):
if mel2ph[i] != mel2ph[i - 1]:
split_loc = i * hop_size
split_locs.append(split_loc)
new_audio = []
for i in range(len(split_locs) - 1):
new_audio.append(audio[split_locs[i]:split_locs[i + 1]])
new_audio.append(np.zeros([0.5 * audio_num_mel_bins]))
return np.concatenate(new_audio)
def mel2token_to_dur(mel2token, T_txt=None, max_dur=None):
is_torch = isinstance(mel2token, torch.Tensor)
has_batch_dim = True
if not is_torch:
mel2token = torch.LongTensor(mel2token)
if T_txt is None:
T_txt = mel2token.max()
if len(mel2token.shape) == 1:
mel2token = mel2token[None, ...]
has_batch_dim = False
B, _ = mel2token.shape
dur = mel2token.new_zeros(B, T_txt + 1).scatter_add(1, mel2token, torch.ones_like(mel2token))
dur = dur[:, 1:]
if max_dur is not None:
dur = dur.clamp(max=max_dur)
if not is_torch:
dur = dur.numpy()
if not has_batch_dim:
dur = dur[0]
return dur
@@ -0,0 +1,22 @@
import subprocess
import numpy as np
from scipy.io import wavfile
def save_wav(wav, path, sr, norm=False):
if norm:
wav = wav / np.abs(wav).max()
wav = wav * 32767
wavfile.write(path[:-4] + '.wav', sr, wav.astype(np.int16))
if path[-4:] == '.mp3':
to_mp3(path[:-4])
def to_mp3(out_path):
if out_path[-4:] == '.wav':
out_path = out_path[:-4]
subprocess.check_call(
f'ffmpeg -threads 1 -loglevel error -i "{out_path}.wav" -vn -b:a 192k -y -hide_banner -async 1 "{out_path}.mp3"',
shell=True, stdin=subprocess.PIPE)
subprocess.check_call(f'rm -f "{out_path}.wav"', shell=True)
@@ -0,0 +1,139 @@
import math
import numpy as np
import torch
import torch.utils.data
from librosa.filters import mel as librosa_mel_fn
from scipy.io.wavfile import read
import torch
import torch.nn as nn
MAX_WAV_VALUE = 32768.0
def load_wav(full_path):
sampling_rate, data = read(full_path)
return data, sampling_rate
def dynamic_range_compression(x, C=1, clip_val=1e-5):
return np.log10(np.clip(x, a_min=clip_val, a_max=None) * C)
def dynamic_range_decompression(x, C=1):
return np.exp(x) / C
def dynamic_range_compression_torch(x, C=1, clip_val=1e-5):
return torch.log10(torch.clamp(x, min=clip_val) * C)
def dynamic_range_decompression_torch(x, C=1):
return torch.exp(x) / C
def spectral_normalize_torch(magnitudes):
output = dynamic_range_compression_torch(magnitudes)
return output
def spectral_de_normalize_torch(magnitudes):
output = dynamic_range_decompression_torch(magnitudes)
return output
class MelNet(nn.Module):
def __init__(self, hparams, device='cpu') -> None:
super().__init__()
self.n_fft = hparams['fft_size']
self.num_mels = hparams['audio_num_mel_bins']
self.sampling_rate = hparams['audio_sample_rate']
self.hop_size = hparams['hop_size']
self.win_size = hparams['win_size']
self.fmin = hparams['fmin']
self.fmax = hparams['fmax']
self.device = device
mel = librosa_mel_fn(sr=self.sampling_rate, n_fft=self.n_fft, n_mels=self.num_mels, fmin=self.fmin,
fmax=self.fmax)
self.mel_basis = torch.from_numpy(mel).float().to(self.device)
self.hann_window = torch.hann_window(self.win_size).to(self.device)
def to(self, device, **kwagrs):
super().to(device=device, **kwagrs)
self.mel_basis = self.mel_basis.to(device)
self.hann_window = self.hann_window.to(device)
self.device = device
def forward(self, y, center=False, complex=False):
if isinstance(y, np.ndarray):
y = torch.FloatTensor(y)
if len(y.shape) == 1:
y = y.unsqueeze(0)
y = y.clamp(min=-1., max=1.).to(self.device)
pad_length = math.ceil(y.shape[1] / self.hop_size) * self.hop_size - y.shape[1]
y = torch.nn.functional.pad(y.unsqueeze(1),
[int((self.n_fft - self.hop_size) / 2),
int((self.n_fft - self.hop_size) / 2 + pad_length)],
mode='reflect')
y = y.squeeze(1)
spec = torch.stft(y, self.n_fft, hop_length=self.hop_size, win_length=self.win_size, window=self.hann_window,
center=center, pad_mode='reflect', normalized=False, onesided=True, return_complex=True)
if not complex:
spec = torch.view_as_real(spec)
spec = torch.sqrt(spec.pow(2).sum(-1) + (1e-9)) # [B, n_fft, T]
spec = torch.matmul(self.mel_basis, spec)
spec = spectral_normalize_torch(spec)
spec = spec.transpose(1, 2) # [B, T, n_fft]
else:
B, C, T, _ = spec.shape
spec = spec.transpose(1, 2) # [B, T, n_fft, 2]
return spec
## below can be used in one gpu, but not ddp
mel_basis = {}
hann_window = {}
def mel_spectrogram(y, hparams, center=False, complex=False): # y should be a tensor with shape (b,wav_len)
# hop_size: 512 # For 22050Hz, 275 ~= 12.5 ms (0.0125 * sample_rate)
# win_size: 2048 # For 22050Hz, 1100 ~= 50 ms (If None, win_size: fft_size) (0.05 * sample_rate)
# fmin: 55 # Set this to 55 if your speaker is male! if female, 95 should help taking off noise. (To test depending on dataset. Pitch info: male~[65, 260], female~[100, 525])
# fmax: 10000 # To be increased/reduced depending on data.
# fft_size: 2048 # Extra window size is filled with 0 paddings to match this parameter
# n_fft, num_mels, sampling_rate, hop_size, win_size, fmin, fmax,
n_fft = hparams['fft_size']
num_mels = hparams['audio_num_mel_bins']
sampling_rate = hparams['audio_sample_rate']
hop_size = hparams['hop_size']
win_size = hparams['win_size']
fmin = hparams['fmin']
fmax = hparams['fmax']
if isinstance(y, np.ndarray):
y = torch.FloatTensor(y)
if len(y.shape) == 1:
y = y.unsqueeze(0)
y = y.clamp(min=-1., max=1.)
global mel_basis, hann_window
if fmax not in mel_basis:
mel = librosa_mel_fn(sampling_rate, n_fft, num_mels, fmin, fmax)
mel_basis[str(fmax) + '_' + str(y.device)] = torch.from_numpy(mel).float().to(y.device)
hann_window[str(y.device)] = torch.hann_window(win_size).to(y.device)
y = torch.nn.functional.pad(y.unsqueeze(1), [int((n_fft - hop_size) / 2), int((n_fft - hop_size) / 2)],
mode='reflect')
y = y.squeeze(1)
spec = torch.stft(y, n_fft, hop_length=hop_size, win_length=win_size, window=hann_window[str(y.device)],
center=center, pad_mode='reflect', normalized=False, onesided=True, return_complex=complex)
if not complex:
spec = torch.sqrt(spec.pow(2).sum(-1) + (1e-9))
spec = torch.matmul(mel_basis[str(fmax) + '_' + str(y.device)], spec)
spec = spectral_normalize_torch(spec)
else:
B, C, T, _ = spec.shape
spec = spec.transpose(1, 2) # [B, T, n_fft, 2]
return spec
@@ -0,0 +1,60 @@
import math
import numpy as np
PITCH_EXTRACTOR = {}
def register_pitch_extractor(name):
def register_pitch_extractor_(cls):
PITCH_EXTRACTOR[name] = cls
return cls
return register_pitch_extractor_
def get_pitch_extractor(name):
return PITCH_EXTRACTOR[name]
def extract_pitch_simple(wav):
from ..commons.hparams import hparams
return extract_pitch(hparams['pitch_extractor'], wav,
hparams['hop_size'], hparams['audio_sample_rate'],
f0_min=hparams['f0_min'], f0_max=hparams['f0_max'])
def extract_pitch(extractor_name, wav_data, hop_size, audio_sample_rate, f0_min=75, f0_max=800, **kwargs):
return get_pitch_extractor(extractor_name)(wav_data, hop_size, audio_sample_rate, f0_min, f0_max, **kwargs)
@register_pitch_extractor('parselmouth')
def parselmouth_pitch(wav_data, hop_size, audio_sample_rate, f0_min, f0_max,
voicing_threshold=0.6, *args, **kwargs):
import parselmouth
time_step = hop_size / audio_sample_rate * 1000
n_mel_frames = int(len(wav_data) // hop_size)
f0_pm = parselmouth.Sound(wav_data, audio_sample_rate).to_pitch_ac(
time_step=time_step / 1000, voicing_threshold=voicing_threshold,
pitch_floor=f0_min, pitch_ceiling=f0_max).selected_array['frequency']
pad_size = (n_mel_frames - len(f0_pm) + 1) // 2
f0 = np.pad(f0_pm, [[pad_size, n_mel_frames - len(f0_pm) - pad_size]], mode='constant')
return f0
@register_pitch_extractor('pyworld')
def pyworld_pitch(wav_data, hop_size, audio_sample_rate, f0_min, f0_max,
voicing_threshold=0.6, *args, **kwargs):
import pyworld as pw
# f0, _ = pw.harvest(wav_data.astype(np.double), audio_sample_rate, f0_floor=f0_min, f0_ceil=f0_max,
# frame_period=hop_size * 1000 / audio_sample_rate)
f0, _ = pw.dio(wav_data.astype(np.double), audio_sample_rate, f0_floor=f0_min, f0_ceil=f0_max, frame_period=hop_size * 1000 / audio_sample_rate)
f0[f0 < f0_min] = 0.0
f0[f0 > f0_max] = 0.0
n_mel_frames = math.ceil(len(wav_data) / hop_size)
if n_mel_frames > len(f0):
pad_size = (n_mel_frames - len(f0) + 1) // 2
f0 = np.pad(f0, [[pad_size, n_mel_frames - len(f0) - pad_size]], mode='constant')
elif n_mel_frames < len(f0):
left_del = (len(f0) - n_mel_frames + 1) // 2
right_del = len(f0) - n_mel_frames - left_del
f0 = f0[left_del: (-right_del if right_del > 0 else len(f0))]
return f0
@@ -0,0 +1,303 @@
import numpy as np
import torch
import pretty_midi
def to_lf0(f0):
f0[f0 < 1.0e-5] = 1.0e-6
lf0 = f0.log() if isinstance(f0, torch.Tensor) else np.log(f0)
lf0[f0 < 1.0e-5] = - 1.0E+10
return lf0
def to_f0(lf0):
f0 = np.where(lf0 <= 0, 0.0, np.exp(lf0))
return f0.flatten()
def f0_to_coarse(f0, f0_bin=256, f0_max=900.0, f0_min=50.0):
f0_mel_min = 1127 * np.log(1 + f0_min / 700)
f0_mel_max = 1127 * np.log(1 + f0_max / 700)
is_torch = isinstance(f0, torch.Tensor)
f0_mel = 1127 * (1 + f0 / 700).log() if is_torch else 1127 * np.log(1 + f0 / 700)
f0_mel[f0_mel > 0] = (f0_mel[f0_mel > 0] - f0_mel_min) * (f0_bin - 2) / (f0_mel_max - f0_mel_min) + 1
f0_mel[f0_mel <= 1] = 1
f0_mel[f0_mel > f0_bin - 1] = f0_bin - 1
f0_coarse = (f0_mel + 0.5).long() if is_torch else np.rint(f0_mel).astype(int)
assert f0_coarse.max() <= f0_bin-1 and f0_coarse.min() >= 1, (f0_coarse.max(), f0_coarse.min(), f0.min(), f0.max())
return f0_coarse
def coarse_to_f0(f0_coarse, f0_bin=256, f0_max=900.0, f0_min=50.0):
f0_mel_min = 1127 * np.log(1 + f0_min / 700)
f0_mel_max = 1127 * np.log(1 + f0_max / 700)
uv = f0_coarse == 1
f0 = f0_mel_min + (f0_coarse - 1) * (f0_mel_max - f0_mel_min) / (f0_bin - 2)
f0 = ((f0 / 1127).exp() - 1) * 700
f0[uv] = 0
return f0
def norm_f0(f0, uv, pitch_norm='log', f0_mean=400, f0_std=100):
is_torch = isinstance(f0, torch.Tensor)
if pitch_norm == 'standard':
f0 = (f0 - f0_mean) / f0_std
if pitch_norm == 'log':
f0 = torch.log2(f0 + 1e-8) if is_torch else np.log2(f0 + 1e-8)
if uv is not None:
f0[uv > 0] = 0
return f0
def norm_interp_f0(f0, pitch_norm='log', f0_mean=None, f0_std=None):
is_torch = isinstance(f0, torch.Tensor)
if is_torch:
device = f0.device
f0 = f0.data.cpu().numpy()
uv = f0 == 0
f0 = norm_f0(f0, uv, pitch_norm, f0_mean, f0_std)
if sum(uv) == len(f0):
f0[uv] = 0
elif sum(uv) > 0:
f0[uv] = np.interp(np.where(uv)[0], np.where(~uv)[0], f0[~uv])
if is_torch:
uv = torch.FloatTensor(uv)
f0 = torch.FloatTensor(f0)
f0 = f0.to(device)
uv = uv.to(device)
return f0, uv
def denorm_f0(f0, uv, pitch_norm='log', f0_mean=400, f0_std=100, pitch_padding=None, min=50, max=900):
is_torch = isinstance(f0, torch.Tensor)
if pitch_norm == 'standard':
f0 = f0 * f0_std + f0_mean
if pitch_norm == 'log':
f0 = 2 ** f0
f0 = f0.clamp(min=min, max=max) if is_torch else np.clip(f0, a_min=min, a_max=max)
if uv is not None:
f0[uv > 0] = 0
if pitch_padding is not None:
f0[pitch_padding] = 0
return f0
def interp_f0(f0, uv=None):
if uv is None:
uv = f0 == 0
f0 = norm_f0(f0, uv)
if uv.any() and not uv.all():
f0[uv] = np.interp(np.where(uv)[0], np.where(~uv)[0], f0[~uv])
return denorm_f0(f0, uv=None), uv
def resample_align_curve(points: np.ndarray, original_timestep: float, target_timestep: float, align_length=-1):
t_max = (len(points) - 1) * original_timestep
curve_interp = np.interp(
np.arange(0, t_max, target_timestep),
original_timestep * np.arange(len(points)),
points
).astype(points.dtype)
if align_length > 0:
delta_l = align_length - len(curve_interp)
if delta_l < 0:
curve_interp = curve_interp[:align_length]
elif delta_l > 0:
curve_interp = np.concatenate((curve_interp, np.full(delta_l, fill_value=curve_interp[-1])), axis=0)
return curve_interp
def midi_to_hz(midi):
if type(midi) == np.ndarray:
non_mask = midi == 0
freq_hz = 440.0 * 2.0 ** ((midi - 69.0) / 12.0)
freq_hz[non_mask] = 0
else:
freq_hz = 440.0 * 2.0 ** ((midi - 69.0) / 12.0)
return freq_hz
def hz_to_midi(hz):
if type(hz) == torch.Tensor:
non_mask = hz == 0
midi = 69.0 + 12.0 * (torch.log2(hz) - torch.log2(torch.Tensor(440.0)))
midi[non_mask] = 0
elif type(hz) == np.ndarray:
non_mask = hz == 0
midi = 69.0 + 12.0 * (np.log2(hz) - np.log2(440.0))
midi[non_mask] = 0
else:
midi = 69.0 + 12.0 * (np.log2(hz) - np.log2(440.0))
if hz == 0:
midi = 0
return midi
def boundary2Interval(bd):
# bd has a shape of [T] with T frames
is_torch = isinstance(bd, torch.Tensor)
if is_torch:
device = bd.device
bd = bd.data.cpu().numpy()
assert len(bd.shape) == 1
# force valid begin and end
# bd[0] = 0 # took care of in regulate_boundary()
# bd[-1] = 0
ret = np.zeros(shape=(bd.sum() + 1, 2), dtype=int)
ret_idx = 0
ret[0, 0] = 0
for i, u in enumerate(bd):
if i == 0:
continue
if u == 1:
ret[ret_idx, 1] = i
ret[ret_idx+1, 0] = i
ret_idx += 1
ret[-1, 1] = bd.shape[0] - 1
if is_torch:
ret = torch.LongTensor(ret).to(device)
return ret
def validate_pitch_and_itv(notes, note_itv):
# notes [T]
# note_itv [T, 2]
assert notes.shape[0] == note_itv.shape[0]
res_notes = []
res_note_itv = []
for idx in range(notes.shape[0]):
pitch, itv = notes[idx], note_itv[idx]
if itv[0] >= itv[1]:
raise RuntimeError("The note duration should be positive")
if pitch == 0:
continue
res_notes.append(pitch)
res_note_itv.append([itv[0], itv[1]])
res_notes = np.array(res_notes)
res_note_itv = np.array(res_note_itv)
return res_notes, res_note_itv
def save_midi(notes, note_itv, midi_path):
# notes [T]
# note_itv [T, 2]
notes, note_itv = validate_pitch_and_itv(notes, note_itv)
if notes.shape == (0,):
return None
assert notes.shape[0] == note_itv.shape[0]
piano_chord = pretty_midi.PrettyMIDI()
piano_program = pretty_midi.instrument_name_to_program('Acoustic Grand Piano')
piano = pretty_midi.Instrument(program=piano_program)
for idx in range(notes.shape[0]):
pitch, itv = notes[idx], note_itv[idx]
note = pretty_midi.Note(velocity=120, pitch=pitch, start=itv[0], end=itv[1])
piano.notes.append(note)
piano_chord.remove_invalid_notes()
piano_chord.instruments.append(piano)
piano_chord.write(midi_path)
return piano_chord
def midi2NoteInterval(mid):
assert type(mid) == pretty_midi.PrettyMIDI
if len(mid.instruments) == 0 or len(mid.instruments[0].notes) == 0:
return None
ret = np.zeros(shape=(len(mid.instruments[0].notes), 2))
for i, note in enumerate(mid.instruments[0].notes):
ret[i, 0] = note.start
ret[i, 1] = note.end
return ret
def midi2NotePitch(mid):
assert type(mid) == pretty_midi.PrettyMIDI
if len(mid.instruments) == 0 or len(mid.instruments[0].notes) == 0:
return None
ret = np.zeros(shape=len(mid.instruments[0].notes))
for i, note in enumerate(mid.instruments[0].notes):
ret[i] = note.pitch
return ret
def midi_onset_eval(mid_gt, mid_pred):
import mir_eval
interval_true = midi2NoteInterval(mid_gt)
if interval_true is None:
raise RuntimeError('Midi ground truth is None')
interval_pred = midi2NoteInterval(mid_pred)
if interval_pred is None:
return 0, 0, 0
onset_p, onset_r, onset_f = mir_eval.transcription.onset_precision_recall_f1(
interval_true, interval_pred, onset_tolerance=0.05, strict=False, beta=1.0)
return onset_p, onset_r, onset_f
def midi_offset_eval(mid_gt, mid_pred):
import mir_eval
interval_true = midi2NoteInterval(mid_gt)
if interval_true is None:
raise RuntimeError('Midi ground truth is None')
interval_pred = midi2NoteInterval(mid_pred)
if interval_pred is None:
return 0, 0, 0
offset_p, offset_r, offset_f = mir_eval.transcription.offset_precision_recall_f1(
interval_true, interval_pred, offset_ratio=0.2, offset_min_tolerance=0.05, strict=False, beta=1.0)
return offset_p, offset_r, offset_f
def midi_pitch_eval(mid_gt, mid_pred, offset_ratio=0.2):
import mir_eval
interval_true = midi2NoteInterval(mid_gt)
pitch_true = midi_to_hz(midi2NotePitch(mid_gt))
if interval_true is None or pitch_true is None:
raise RuntimeError('Midi ground truth is None')
interval_pred = midi2NoteInterval(mid_pred)
pitch_pred = midi2NotePitch(mid_pred)
if interval_pred is None:
return 0, 0, 0, 0
if pitch_pred is None:
pitch_pred = np.zeros(interval_pred.shape[0])
pitch_pred = midi_to_hz(pitch_pred)
overlap_p, overlap_r, overlap_f, avg_overlap_ratio = mir_eval.transcription.precision_recall_f1_overlap(
interval_true, pitch_true, interval_pred, pitch_pred, onset_tolerance=0.05, pitch_tolerance=50.0,
offset_ratio=offset_ratio, offset_min_tolerance=0.05, strict=False, beta=1.0)
return overlap_p, overlap_r, overlap_f, avg_overlap_ratio
def midi_COn_eval(mid_gt, mid_pred):
return midi_onset_eval(mid_gt, mid_pred)
def midi_COnP_eval(mid_gt, mid_pred):
return midi_pitch_eval(mid_gt, mid_pred, offset_ratio=None)
def midi_COnPOff_eval(mid_gt, mid_pred):
return midi_pitch_eval(mid_gt, mid_pred)
def midi_melody_eval(mid_gt, mid_pred, hop_size=256, sample_rate=48000):
interval_true = midi2NoteInterval(mid_gt)
pitch_true = midi_to_hz(midi2NotePitch(mid_gt))
if interval_true is None or pitch_true is None:
raise RuntimeError('Midi ground truth is None')
interval_pred = midi2NoteInterval(mid_pred)
pitch_pred = midi2NotePitch(mid_pred)
if interval_pred is None:
return 0, 0, 0, 0
if pitch_pred is None:
pitch_pred = np.zeros(interval_pred.shape[0])
pitch_pred = midi_to_hz(pitch_pred)
vr, vfa, rpa, rca, oa = melody_eval_pitch_and_itv(
pitch_true, interval_true, pitch_pred, interval_pred, hop_size, sample_rate)
return vr, vfa, rpa, rca, oa
def melody_eval_pitch_and_itv(pitch_true, interval_true, pitch_pred, interval_pred, hop_size=256, sample_rate=48000):
import mir_eval
t_gt = np.arange(0, interval_true[-1][1], hop_size / sample_rate)
freq_gt = np.zeros_like(t_gt)
for idx in range(len(pitch_true)):
freq_gt[min(len(freq_gt) - 1, round(interval_true[idx][0] * sample_rate / hop_size)): round(
interval_true[idx][1] * sample_rate / hop_size)] = pitch_true[idx]
t_pred = np.arange(0, interval_pred[-1][1], hop_size / sample_rate)
freq_pred = np.zeros_like(t_pred)
for idx in range(len(pitch_pred)):
freq_pred[min(len(freq_pred) - 1, round(interval_pred[idx][0] * sample_rate / hop_size)): round(
interval_pred[idx][1] * sample_rate / hop_size)] = pitch_pred[idx]
ref_voicing, ref_cent, est_voicing, est_cent = mir_eval.melody.to_cent_voicing(t_gt, freq_gt,
t_pred, freq_pred)
vr, vfa = mir_eval.melody.voicing_measures(ref_voicing,
est_voicing) # voicing recall, voicing false alarm
rpa = mir_eval.melody.raw_pitch_accuracy(ref_voicing, ref_cent, est_voicing, est_cent)
rca = mir_eval.melody.raw_chroma_accuracy(ref_voicing, ref_cent, est_voicing, est_cent)
oa = mir_eval.melody.overall_accuracy(ref_voicing, ref_cent, est_voicing, est_cent)
return vr, vfa, rpa, rca, oa
@@ -0,0 +1,78 @@
from skimage.transform import resize
import struct
import webrtcvad
from scipy.ndimage.morphology import binary_dilation
import librosa
import numpy as np
import pyloudnorm as pyln
import warnings
warnings.filterwarnings("ignore", message="Possible clipped samples in output")
int16_max = (2 ** 15) - 1
def trim_long_silences(path, sr=None, return_raw_wav=False, norm=True, vad_max_silence_length=12):
"""
Ensures that segments without voice in the waveform remain no longer than a
threshold determined by the VAD parameters in params.py.
:param wav: the raw waveform as a numpy array of floats
:param vad_max_silence_length: Maximum number of consecutive silent frames a segment can have.
:return: the same waveform with silences trimmed away (length <= original wav length)
"""
## Voice Activation Detection
# Window size of the VAD. Must be either 10, 20 or 30 milliseconds.
# This sets the granularity of the VAD. Should not need to be changed.
sampling_rate = 16000
wav_raw, sr = librosa.core.load(path, sr=sr)
if norm:
meter = pyln.Meter(sr) # create BS.1770 meter
loudness = meter.integrated_loudness(wav_raw)
wav_raw = pyln.normalize.loudness(wav_raw, loudness, -20.0)
if np.abs(wav_raw).max() > 1.0:
wav_raw = wav_raw / np.abs(wav_raw).max()
wav = librosa.resample(wav_raw, sr, sampling_rate, res_type='kaiser_best')
vad_window_length = 30 # In milliseconds
# Number of frames to average together when performing the moving average smoothing.
# The larger this value, the larger the VAD variations must be to not get smoothed out.
vad_moving_average_width = 8
# Compute the voice detection window size
samples_per_window = (vad_window_length * sampling_rate) // 1000
# Trim the end of the audio to have a multiple of the window size
wav = wav[:len(wav) - (len(wav) % samples_per_window)]
# Convert the float waveform to 16-bit mono PCM
pcm_wave = struct.pack("%dh" % len(wav), *(np.round(wav * int16_max)).astype(np.int16))
# Perform voice activation detection
voice_flags = []
vad = webrtcvad.Vad(mode=3)
for window_start in range(0, len(wav), samples_per_window):
window_end = window_start + samples_per_window
voice_flags.append(vad.is_speech(pcm_wave[window_start * 2:window_end * 2],
sample_rate=sampling_rate))
voice_flags = np.array(voice_flags)
# Smooth the voice detection with a moving average
def moving_average(array, width):
array_padded = np.concatenate((np.zeros((width - 1) // 2), array, np.zeros(width // 2)))
ret = np.cumsum(array_padded, dtype=float)
ret[width:] = ret[width:] - ret[:-width]
return ret[width - 1:] / width
audio_mask = moving_average(voice_flags, vad_moving_average_width)
audio_mask = np.round(audio_mask).astype(np.bool)
# Dilate the voiced regions
audio_mask = binary_dilation(audio_mask, np.ones(vad_max_silence_length + 1))
audio_mask = np.repeat(audio_mask, samples_per_window)
audio_mask = resize(audio_mask, (len(wav_raw),)) > 0
if return_raw_wav:
return wav_raw, audio_mask, sr
return wav_raw[audio_mask], audio_mask, sr
@@ -0,0 +1,235 @@
import logging
import os
import random
import subprocess
import sys
from datetime import datetime
import numpy as np
import torch.utils.data
from torch import nn
from torch.utils.tensorboard import SummaryWriter
from .dataset_utils import data_loader
from .hparams import hparams
from .meters import AvgrageMeter
from .tensor_utils import tensors_to_scalars
from .trainer import Trainer
torch.multiprocessing.set_sharing_strategy(os.getenv('TORCH_SHARE_STRATEGY', 'file_system'))
log_format = '%(asctime)s %(message)s'
logging.basicConfig(stream=sys.stdout, level=logging.INFO,
format=log_format, datefmt='%m/%d %I:%M:%S %p')
class BaseTask(nn.Module):
def __init__(self, *args, **kwargs):
super(BaseTask, self).__init__()
self.current_epoch = 0
self.global_step = 0
self.trainer = None
self.use_ddp = False
self.gradient_clip_norm = hparams['clip_grad_norm']
self.gradient_clip_val = hparams.get('clip_grad_value', 0)
self.model = None
self.training_losses_meter = None
self.logger: SummaryWriter = None
######################
# build model, dataloaders, optimizer, scheduler and tensorboard
######################
def build_model(self):
raise NotImplementedError
@data_loader
def train_dataloader(self):
raise NotImplementedError
@data_loader
def test_dataloader(self):
raise NotImplementedError
@data_loader
def val_dataloader(self):
raise NotImplementedError
def build_scheduler(self, optimizer):
return None
def build_optimizer(self, model):
raise NotImplementedError
def configure_optimizers(self):
optm = self.build_optimizer(self.model)
self.scheduler = self.build_scheduler(optm)
if isinstance(optm, (list, tuple)):
return optm
return [optm]
def build_tensorboard(self, save_dir, name, **kwargs):
log_dir = os.path.join(save_dir, name)
os.makedirs(log_dir, exist_ok=True)
self.logger = SummaryWriter(log_dir=log_dir, **kwargs)
######################
# training
######################
def on_train_start(self):
pass
def on_train_end(self):
pass
def on_epoch_start(self):
self.training_losses_meter = {'total_loss': AvgrageMeter()}
def on_epoch_end(self):
loss_outputs = {k: round(v.avg, 4) for k, v in self.training_losses_meter.items()}
print(f"Epoch {self.current_epoch} ended. Steps: {self.global_step}. {loss_outputs}")
def _training_step(self, sample, batch_idx, optimizer_idx):
"""
:param sample:
:param batch_idx:
:return: total loss: torch.Tensor, loss_log: dict
"""
raise NotImplementedError
def training_step(self, sample, batch_idx, optimizer_idx=-1):
"""
:param sample:
:param batch_idx:
:param optimizer_idx:
:return: {'loss': torch.Tensor, 'progress_bar': dict, 'tb_log': dict}
"""
loss_ret = self._training_step(sample, batch_idx, optimizer_idx)
if loss_ret is None:
return {'loss': None}
total_loss, log_outputs = loss_ret
log_outputs = tensors_to_scalars(log_outputs)
for k, v in log_outputs.items():
if k not in self.training_losses_meter:
self.training_losses_meter[k] = AvgrageMeter()
if not np.isnan(v):
self.training_losses_meter[k].update(v)
self.training_losses_meter['total_loss'].update(total_loss.item())
if optimizer_idx >= 0:
log_outputs[f'lr_{optimizer_idx}'] = self.trainer.optimizers[optimizer_idx].param_groups[0]['lr']
progress_bar_log = log_outputs
tb_log = {f'tr/{k}': v for k, v in log_outputs.items()}
return {
'loss': total_loss,
'progress_bar': progress_bar_log,
'tb_log': tb_log
}
def on_before_optimization(self, opt_idx):
if self.gradient_clip_norm > 0:
torch.nn.utils.clip_grad_norm_(self.parameters(), self.gradient_clip_norm)
if self.gradient_clip_val > 0:
torch.nn.utils.clip_grad_value_(self.parameters(), self.gradient_clip_val)
def on_after_optimization(self, epoch, batch_idx, optimizer, optimizer_idx):
if self.scheduler is not None:
# self.scheduler.step(self.global_step // hparams['accumulate_grad_batches'])
# the code above causes EPOCH_DEPRECATION_WARNING, changed it and changed the optimizer init with
# step_size divided by accumulate_grad_batches
self.scheduler.step()
######################
# validation
######################
def validation_start(self):
pass
def validation_step(self, sample, batch_idx):
"""
:param sample:
:param batch_idx:
:return: output: {"losses": {...}, "total_loss": float, ...} or (total loss: torch.Tensor, loss_log: dict)
"""
raise NotImplementedError
def validation_end(self, outputs):
"""
:param outputs:
:return: loss_output: dict
"""
all_losses_meter = {'total_loss': AvgrageMeter()}
for output in outputs:
if len(output) == 0 or output is None:
continue
if isinstance(output, dict):
assert 'losses' in output, 'Key "losses" should exist in validation output.'
n = output.pop('nsamples', 1)
losses = tensors_to_scalars(output['losses'])
total_loss = output.get('total_loss', sum(losses.values()))
else:
assert len(output) == 2, 'Validation output should only consist of two elements: (total_loss, losses)'
n = 1
total_loss, losses = output
losses = tensors_to_scalars(losses)
if isinstance(total_loss, torch.Tensor):
total_loss = total_loss.item()
for k, v in losses.items():
if k not in all_losses_meter:
all_losses_meter[k] = AvgrageMeter()
all_losses_meter[k].update(v, n)
all_losses_meter['total_loss'].update(total_loss, n)
loss_output = {k: round(v.avg, 4) for k, v in all_losses_meter.items()}
print(f"| Validation results@{self.global_step}: {loss_output}")
return {
'tb_log': {f'val/{k}': v for k, v in loss_output.items()},
'val_loss': loss_output['total_loss']
}
######################
# testing
######################
def test_start(self):
pass
def test_step(self, sample, batch_idx):
return self.validation_step(sample, batch_idx)
def test_end(self, outputs):
return self.validation_end(outputs)
######################
# start training/testing
######################
@classmethod
def start(cls):
os.environ['MASTER_PORT'] = str(random.randint(15000, 30000))
random.seed(hparams['seed'])
np.random.seed(hparams['seed'])
work_dir = hparams['work_dir']
trainer = Trainer(
work_dir=work_dir,
val_check_interval=hparams['val_check_interval'],
tb_log_interval=hparams['tb_log_interval'],
max_updates=hparams['max_updates'],
num_sanity_val_steps=hparams['num_sanity_val_steps'] if not hparams['validate'] else 10000,
accumulate_grad_batches=hparams['accumulate_grad_batches'],
print_nan_grads=hparams['print_nan_grads'],
resume_from_checkpoint=hparams.get('resume_from_checkpoint', 0),
amp=hparams['amp'],
monitor_key=hparams['valid_monitor_key'],
monitor_mode=hparams['valid_monitor_mode'],
num_ckpt_keep=hparams['num_ckpt_keep'],
save_best=hparams['save_best'],
seed=hparams['seed'],
debug=hparams['debug']
)
if not hparams['infer']: # train
trainer.fit(cls)
else:
trainer.test(cls)
def on_keyboard_interrupt(self):
pass
@@ -0,0 +1,68 @@
import glob
import os
import re
import torch
def get_last_checkpoint(work_dir, steps=None):
checkpoint = None
last_ckpt_path = None
ckpt_paths = get_all_ckpts(work_dir, steps)
if len(ckpt_paths) > 0:
last_ckpt_path = ckpt_paths[0]
checkpoint = torch.load(last_ckpt_path, map_location='cpu')
return checkpoint, last_ckpt_path
def get_all_ckpts(work_dir, steps=None):
if steps is None:
ckpt_path_pattern = f'{work_dir}/model_ckpt_steps_*.ckpt'
else:
ckpt_path_pattern = f'{work_dir}/model_ckpt_steps_{steps}.ckpt'
return sorted(glob.glob(ckpt_path_pattern),
key=lambda x: -int(re.findall('.*steps\_(\d+)\.ckpt', x)[0]))
def load_ckpt(cur_model, ckpt_base_dir, model_name='model', force=True, strict=True, verbose=True):
if os.path.isfile(ckpt_base_dir):
base_dir = os.path.dirname(ckpt_base_dir)
ckpt_path = ckpt_base_dir
checkpoint = torch.load(ckpt_base_dir, map_location='cpu')
else:
base_dir = ckpt_base_dir
checkpoint, ckpt_path = get_last_checkpoint(ckpt_base_dir)
if checkpoint is not None:
state_dict = checkpoint["state_dict"]
if len([k for k in state_dict.keys() if '.' in k]) > 0:
state_dict = {k[len(model_name) + 1:]: v for k, v in state_dict.items()
if k.startswith(f'{model_name}.')}
else:
if '.' not in model_name:
state_dict = state_dict[model_name]
else:
base_model_name = model_name.split('.')[0]
rest_model_name = model_name[len(base_model_name) + 1:]
state_dict = {
k[len(rest_model_name) + 1:]: v for k, v in state_dict[base_model_name].items()
if k.startswith(f'{rest_model_name}.')}
if not strict:
cur_model_state_dict = cur_model.state_dict()
unmatched_keys = []
for key, param in state_dict.items():
if key in cur_model_state_dict:
new_param = cur_model_state_dict[key]
if new_param.shape != param.shape:
unmatched_keys.append(key)
print("| Unmatched keys: ", key, new_param.shape, param.shape)
for key in unmatched_keys:
del state_dict[key]
# print(state_dict)
cur_model.load_state_dict(state_dict, strict=strict)
if verbose:
print(f"| load '{model_name}' from '{ckpt_path}'.")
else:
e_msg = f"| ckpt not found in {base_dir}."
if force:
assert False, e_msg
else:
print(e_msg)
@@ -0,0 +1,372 @@
import os
import sys
import traceback
import types
from functools import wraps
from itertools import chain
import numpy as np
import torch.utils.data
import torch.nn.functional as F
from torch.utils.data import ConcatDataset
from .hparams import hparams
def collate_1d_or_2d(values, pad_idx=0, left_pad=False, shift_right=False, max_len=None, shift_id=1):
if len(values[0].shape) == 1:
return collate_1d(values, pad_idx, left_pad, shift_right, max_len, shift_id)
else:
return collate_2d(values, pad_idx, left_pad, shift_right, max_len)
def collate_1d(values, pad_idx=0, left_pad=False, shift_right=False, max_len=None, shift_id=1):
"""Convert a list of 1d tensors into a padded 2d tensor."""
size = max(v.size(0) for v in values) if max_len is None else max_len
res = values[0].new(len(values), size).fill_(pad_idx)
def copy_tensor(src, dst):
assert dst.numel() == src.numel()
if shift_right:
dst[1:] = src[:-1]
dst[0] = shift_id
else:
dst.copy_(src)
for i, v in enumerate(values):
copy_tensor(v, res[i][size - len(v):] if left_pad else res[i][:len(v)])
return res
def collate_2d(values, pad_idx=0, left_pad=False, shift_right=False, max_len=None):
"""Convert a list of 2d tensors into a padded 3d tensor."""
size = max(v.size(0) for v in values) if max_len is None else max_len
res = values[0].new(len(values), size, values[0].shape[1]).fill_(pad_idx)
def copy_tensor(src, dst):
assert dst.numel() == src.numel()
if shift_right:
dst[1:] = src[:-1]
else:
dst.copy_(src)
for i, v in enumerate(values):
copy_tensor(v, res[i][size - len(v):] if left_pad else res[i][:len(v)])
return res
def collate_xd(values, pad_value=0, max_len=None):
size = ((max(v.size(0) for v in values) if max_len is None else max_len), *values[0].shape[1:])
res = torch.full((len(values), *size), fill_value=pad_value, dtype=values[0].dtype, device=values[0].device)
for i, v in enumerate(values):
res[i, :len(v), ...] = v
return res
def pad_or_cut_1d(values: torch.tensor, tgt_len, pad_value=0):
src_len = values.shape[0]
if src_len < tgt_len:
res = F.pad(values, [0, tgt_len - src_len], value=pad_value)
else:
res = values[:tgt_len]
return res
def pad_or_cut_2d(values: torch.tensor, tgt_len, dim=-1, pad_value=0):
if dim == 0 or dim == -2:
src_len = values.shape[0]
if src_len < tgt_len:
res = F.pad(values, [0, 0, 0, tgt_len - src_len], value=pad_value)
else:
res = values[:tgt_len]
elif dim == 1 or dim == -1:
src_len = values.shape[1]
if src_len < tgt_len:
res = F.pad(values, [0, tgt_len - src_len], value=pad_value)
else:
res = values[:, :tgt_len]
else:
raise RuntimeError(f"Wrong dim number {dim} while the tensor only has {len(values.shape)} dimensions.")
return res
def pad_or_cut_3d(values: torch.tensor, tgt_len, dim=-1, pad_value=0):
if dim == 0 or dim == -3:
src_len = values.shape[0]
if src_len < tgt_len:
res = F.pad(values, [0, 0, 0, 0, 0, tgt_len - src_len], value=pad_value)
else:
res = values[:tgt_len]
elif dim == 1 or dim == -2:
src_len = values.shape[1]
if src_len < tgt_len:
res = F.pad(values, [0, 0, 0, tgt_len - src_len], value=pad_value)
else:
res = values[:, :tgt_len]
elif dim == 2 or dim == -1:
src_len = values.shape[2]
if src_len < tgt_len:
res = F.pad(values, [0, tgt_len - src_len], value=pad_value)
else:
res = values[:, :, :tgt_len]
else:
raise RuntimeError(f"Wrong dim number {dim} while the tensor only has {len(values.shape)} dimensions.")
return res
def pad_or_cut_xd(values, tgt_len, dim=-1, pad_value=0):
if len(values.shape) == 1:
return pad_or_cut_1d(values, tgt_len, pad_value)
elif len(values.shape) == 2:
return pad_or_cut_2d(values, tgt_len, dim, pad_value)
elif len(values.shape) == 3:
return pad_or_cut_3d(values, tgt_len, dim, pad_value)
else:
raise NotImplementedError
def _is_batch_full(batch, num_tokens, max_tokens, max_sentences):
if len(batch) == 0:
return 0
if len(batch) == max_sentences:
return 1
if num_tokens > max_tokens:
return 1
return 0
def batch_by_size(
indices, num_tokens_fn, max_tokens=None, max_sentences=None,
required_batch_size_multiple=1, distributed=False
):
"""
Yield mini-batches of indices bucketed by size. Batches may contain
sequences of different lengths.
Args:
indices (List[int]): ordered list of dataset indices
num_tokens_fn (callable): function that returns the number of tokens at
a given index
max_tokens (int, optional): max number of tokens in each batch
(default: None).
max_sentences (int, optional): max number of sentences in each
batch (default: None).
required_batch_size_multiple (int, optional): require batch size to
be a multiple of N (default: 1).
"""
max_tokens = max_tokens if max_tokens is not None else sys.maxsize
max_sentences = max_sentences if max_sentences is not None else sys.maxsize
bsz_mult = required_batch_size_multiple
if isinstance(indices, types.GeneratorType):
indices = np.fromiter(indices, dtype=np.int64, count=-1)
sample_len = 0
sample_lens = []
batch = []
batches = []
for i in range(len(indices)):
idx = indices[i]
num_tokens = num_tokens_fn(idx)
sample_lens.append(num_tokens)
sample_len = max(sample_len, num_tokens)
assert sample_len <= max_tokens, (
"sentence at index {} of size {} exceeds max_tokens "
"limit of {}!".format(idx, sample_len, max_tokens)
)
num_tokens = (len(batch) + 1) * sample_len
if _is_batch_full(batch, num_tokens, max_tokens, max_sentences):
mod_len = max(
bsz_mult * (len(batch) // bsz_mult),
len(batch) % bsz_mult,
)
batches.append(batch[:mod_len])
batch = batch[mod_len:]
sample_lens = sample_lens[mod_len:]
sample_len = max(sample_lens) if len(sample_lens) > 0 else 0
batch.append(idx)
if len(batch) > 0:
batches.append(batch)
return batches
def build_dataloader(dataset, shuffle, max_tokens=None, max_sentences=None,
required_batch_size_multiple=-1, endless=False, apply_batch_by_size=True, pin_memory=False, use_ddp=False):
import torch.distributed as dist
devices_cnt = torch.cuda.device_count()
if devices_cnt == 0:
devices_cnt = 1
if not use_ddp:
devices_cnt = 1
if required_batch_size_multiple == -1:
required_batch_size_multiple = devices_cnt
def shuffle_batches(batches):
np.random.shuffle(batches)
return batches
if max_tokens is not None:
max_tokens *= devices_cnt
if max_sentences is not None:
max_sentences *= devices_cnt
indices = dataset.ordered_indices()
if apply_batch_by_size:
batch_sampler = batch_by_size(
indices, dataset.num_tokens, max_tokens=max_tokens, max_sentences=max_sentences,
required_batch_size_multiple=required_batch_size_multiple,
)
else:
batch_sampler = []
for i in range(0, len(indices), max_sentences):
batch_sampler.append(indices[i:i + max_sentences])
if shuffle:
batches = shuffle_batches(list(batch_sampler))
if endless:
batches = [b for _ in range(1000) for b in shuffle_batches(list(batch_sampler))]
else:
batches = batch_sampler
if endless:
batches = [b for _ in range(1000) for b in batches]
num_workers = dataset.num_workers
if use_ddp:
num_replicas = dist.get_world_size()
rank = dist.get_rank()
# batches = [x[rank::num_replicas] for x in batches if len(x) % num_replicas == 0]
# ensure that every sample in the dataset is covered
batches_ = []
for x in batches:
if len(x) % num_replicas == 0:
batches_.append(x[rank::num_replicas])
else:
x_ = x + [x[-1]] * (len(x) - len(x) // num_replicas * num_replicas)
batches_.append(x_[rank::num_replicas])
batches = batches_
return torch.utils.data.DataLoader(dataset,
collate_fn=dataset.collater,
batch_sampler=batches,
num_workers=num_workers,
pin_memory=pin_memory)
def unpack_dict_to_list(samples):
samples_ = []
bsz = samples.get('outputs').size(0)
for i in range(bsz):
res = {}
for k, v in samples.items():
try:
res[k] = v[i]
except:
pass
samples_.append(res)
return samples_
def remove_padding(x, padding_idx=0):
if x is None:
return None
assert len(x.shape) in [1, 2]
if len(x.shape) == 2: # [T, H]
return x[np.abs(x).sum(-1) != padding_idx]
elif len(x.shape) == 1: # [T]
return x[x != padding_idx]
def data_loader(fn):
"""
Decorator to make any fx with this use the lazy property
:param fn:
:return:
"""
wraps(fn)
attr_name = '_lazy_' + fn.__name__
def _get_data_loader(self):
try:
value = getattr(self, attr_name)
except AttributeError:
try:
value = fn(self) # Lazy evaluation, done only once.
except AttributeError as e:
# Guard against AttributeError suppression. (Issue #142)
traceback.print_exc()
error = f'{fn.__name__}: An AttributeError was encountered: ' + str(e)
raise RuntimeError(error) from e
setattr(self, attr_name, value) # Memoize evaluation.
return value
return _get_data_loader
class BaseDataset(torch.utils.data.Dataset):
def __init__(self, shuffle):
super().__init__()
self.hparams = hparams
self.shuffle = shuffle
self.sort_by_len = hparams['sort_by_len']
self.sizes = None
@property
def _sizes(self):
return self.sizes
def __getitem__(self, index):
raise NotImplementedError
def collater(self, samples):
raise NotImplementedError
def __len__(self):
return len(self._sizes)
def num_tokens(self, index):
return self.size(index)
def size(self, index):
"""Return an example's size as a float or tuple. This value is used when
filtering a dataset with ``--max-positions``."""
return min(self._sizes[index], hparams['max_frames'])
def ordered_indices(self):
"""Return an ordered list of indices. Batches will be constructed based
on this order."""
if self.shuffle:
indices = np.random.permutation(len(self))
if self.sort_by_len:
indices = indices[np.argsort(np.array(self._sizes)[indices], kind='mergesort')]
else:
indices = np.arange(len(self))
return indices.tolist()
@property
def num_workers(self):
return int(os.getenv('NUM_WORKERS', hparams['ds_workers']))
class BaseConcatDataset(ConcatDataset):
def collater(self, samples):
return self.datasets[0].collater(samples)
@property
def _sizes(self):
if not hasattr(self, 'sizes'):
self.sizes = list(chain.from_iterable([d._sizes for d in self.datasets]))
return self.sizes
def size(self, index):
return min(self._sizes[index], hparams['max_frames'])
def num_tokens(self, index):
return self.size(index)
def ordered_indices(self):
"""Return an ordered list of indices. Batches will be constructed based
on this order."""
if self.datasets[0].shuffle:
indices = np.random.permutation(len(self))
if self.datasets[0].sort_by_len:
indices = indices[np.argsort(np.array(self._sizes)[indices], kind='mergesort')]
else:
indices = np.arange(len(self))
return indices
@property
def num_workers(self):
return self.datasets[0].num_workers
@@ -0,0 +1,164 @@
from torch.nn.parallel import DistributedDataParallel
from torch.nn.parallel.distributed import _find_tensors
import torch.optim
import torch.utils.data
import torch
from packaging import version
class DDP(DistributedDataParallel):
"""
Override the forward call in lightning so it goes to training and validation step respectively
"""
def forward(self, *inputs, **kwargs): # pragma: no cover
# if version.parse(torch.__version__[:6]) < version.parse("1.11"):
if version.parse(torch.__version__) < version.parse("1.11"): # fix the hard [:6] problem
self._sync_params()
inputs, kwargs = self.scatter(inputs, kwargs, self.device_ids)
assert len(self.device_ids) == 1
if self.module.training:
output = self.module.training_step(*inputs[0], **kwargs[0])
elif self.module.testing:
output = self.module.test_step(*inputs[0], **kwargs[0])
else:
output = self.module.validation_step(*inputs[0], **kwargs[0])
if torch.is_grad_enabled():
# We'll return the output object verbatim since it is a freeform
# object. We need to find any tensors in this object, though,
# because we need to figure out which parameters were used during
# this forward pass, to ensure we short circuit reduction for any
# unused parameters. Only if `find_unused_parameters` is set.
if self.find_unused_parameters:
self.reducer.prepare_for_backward(list(_find_tensors(output)))
else:
self.reducer.prepare_for_backward([])
elif version.parse("1.11") <= version.parse(torch.__version__) < version.parse("2.0"):
from torch.nn.parallel.distributed import \
logging, Join, _DDPSink, _tree_flatten_with_rref, _tree_unflatten_with_rref
with torch.autograd.profiler.record_function("DistributedDataParallel.forward"):
if torch.is_grad_enabled() and self.require_backward_grad_sync:
self.logger.set_runtime_stats_and_log()
self.num_iterations += 1
self.reducer.prepare_for_forward()
# Notify the join context that this process has not joined, if
# needed
work = Join.notify_join_context(self)
if work:
self.reducer._set_forward_pass_work_handle(
work, self._divide_by_initial_world_size
)
# Calling _rebuild_buckets before forward compuation,
# It may allocate new buckets before deallocating old buckets
# inside _rebuild_buckets. To save peak memory usage,
# call _rebuild_buckets before the peak memory usage increases
# during forward computation.
# This should be called only once during whole training period.
if torch.is_grad_enabled() and self.reducer._rebuild_buckets():
logging.info("Reducer buckets have been rebuilt in this iteration.")
self._has_rebuilt_buckets = True
# sync params according to location (before/after forward) user
# specified as part of hook, if hook was specified.
buffer_hook_registered = hasattr(self, 'buffer_hook')
if self._check_sync_bufs_pre_fwd():
self._sync_buffers()
if self._join_config.enable:
# Notify joined ranks whether they should sync in backwards pass or not.
self._check_global_requires_backward_grad_sync(is_joined_rank=False)
# modified part
inputs, kwargs = self.scatter(inputs, kwargs, self.device_ids)
if self.module.training:
output = self.module.training_step(*inputs[0], **kwargs[0])
elif self.module.testing:
output = self.module.test_step(*inputs[0], **kwargs[0])
else:
output = self.module.validation_step(*inputs[0], **kwargs[0])
# sync params according to location (before/after forward) user
# specified as part of hook, if hook was specified.
if self._check_sync_bufs_post_fwd():
self._sync_buffers()
if torch.is_grad_enabled() and self.require_backward_grad_sync:
self.require_forward_param_sync = True
# We'll return the output object verbatim since it is a freeform
# object. We need to find any tensors in this object, though,
# because we need to figure out which parameters were used during
# this forward pass, to ensure we short circuit reduction for any
# unused parameters. Only if `find_unused_parameters` is set.
if self.find_unused_parameters and not self.static_graph:
# Do not need to populate this for static graph.
self.reducer.prepare_for_backward(list(_find_tensors(output)))
else:
self.reducer.prepare_for_backward([])
else:
self.require_forward_param_sync = False
# TODO: DDPSink is currently enabled for unused parameter detection and
# static graph training for first iteration.
if (self.find_unused_parameters and not self.static_graph) or (
self.static_graph and self.num_iterations == 1
):
state_dict = {
'static_graph': self.static_graph,
'num_iterations': self.num_iterations,
}
output_tensor_list, treespec, output_is_rref = _tree_flatten_with_rref(
output
)
output_placeholders = [None for _ in range(len(output_tensor_list))]
# Do not touch tensors that have no grad_fn, which can cause issues
# such as https://github.com/pytorch/pytorch/issues/60733
for i, output in enumerate(output_tensor_list):
if torch.is_tensor(output) and output.grad_fn is None:
output_placeholders[i] = output
# When find_unused_parameters=True, makes tensors which require grad
# run through the DDPSink backward pass. When not all outputs are
# used in loss, this makes those corresponding tensors receive
# undefined gradient which the reducer then handles to ensure
# param.grad field is not touched and we don't error out.
passthrough_tensor_list = _DDPSink.apply(
self.reducer,
state_dict,
*output_tensor_list,
)
for i in range(len(output_placeholders)):
if output_placeholders[i] is None:
output_placeholders[i] = passthrough_tensor_list[i]
# Reconstruct output data structure.
output = _tree_unflatten_with_rref(
output_placeholders, treespec, output_is_rref
)
else:
# now pytorch version >= 2.0
with torch.autograd.profiler.record_function("DistributedDataParallel.forward"):
inputs, kwargs = self._pre_forward(*inputs, **kwargs)
output = (
# self.module.forward(*inputs, **kwargs)
# if self._delay_all_reduce_all_params
# else self._run_ddp_forward(*inputs, **kwargs)
# modified: delete 'delay_all_reduce_named_params' function
self._run_ddp_forward(*inputs, **kwargs)
)
return self._post_forward(output)
return output
def _run_ddp_forward(self, *inputs, **kwargs):
if version.parse(torch.__version__) >= version.parse("2.0"):
with self._inside_ddp_forward():
if self.module.training:
output = self.module.training_step(*inputs, **kwargs)
elif self.module.testing:
output = self.module.test_step(*inputs, **kwargs)
else:
output = self.module.validation_step(*inputs, **kwargs)
return output # type: ignore[index]
else:
return super(DDP, self)._run_ddp_forward(*inputs, **kwargs)
@@ -0,0 +1,113 @@
import gc
import datetime
import inspect
import torch
import numpy as np
dtype_memory_size_dict = {
torch.float64: 64/8,
torch.double: 64/8,
torch.float32: 32/8,
torch.float: 32/8,
torch.float16: 16/8,
torch.half: 16/8,
torch.int64: 64/8,
torch.long: 64/8,
torch.int32: 32/8,
torch.int: 32/8,
torch.int16: 16/8,
torch.short: 16/6,
torch.uint8: 8/8,
torch.int8: 8/8,
}
# compatibility of torch1.0
if getattr(torch, "bfloat16", None) is not None:
dtype_memory_size_dict[torch.bfloat16] = 16/8
if getattr(torch, "bool", None) is not None:
dtype_memory_size_dict[torch.bool] = 8/8 # pytorch use 1 byte for a bool, see https://github.com/pytorch/pytorch/issues/41571
def get_mem_space(x):
try:
ret = dtype_memory_size_dict[x]
except KeyError:
print(f"dtype {x} is not supported!")
return ret
class MemTracker(object):
"""
Class used to track pytorch memory usage
Arguments:
detail(bool, default True): whether the function shows the detail gpu memory usage
path(str): where to save log file
verbose(bool, default False): whether show the trivial exception
device(int): GPU number, default is 0
"""
def __init__(self, detail=True, path='', verbose=False, device=0):
self.print_detail = detail
self.last_tensor_sizes = set()
self.gpu_profile_fn = path + f'{datetime.datetime.now():%d-%b-%y-%H:%M:%S}-gpu_mem_track.txt'
self.verbose = verbose
self.begin = True
self.device = device
def get_tensors(self):
for obj in gc.get_objects():
try:
if torch.is_tensor(obj) or (hasattr(obj, 'data') and torch.is_tensor(obj.data)):
tensor = obj
else:
continue
if tensor.is_cuda:
yield tensor
except Exception as e:
if self.verbose:
print('A trivial exception occured: {}'.format(e))
def get_tensor_usage(self):
sizes = [np.prod(np.array(tensor.size())) * get_mem_space(tensor.dtype) for tensor in self.get_tensors()]
return np.sum(sizes) / 1024**2
def get_allocate_usage(self):
return torch.cuda.memory_allocated() / 1024**2
def clear_cache(self):
gc.collect()
torch.cuda.empty_cache()
def print_all_gpu_tensor(self, file=None):
for x in self.get_tensors():
print(x.size(), x.dtype, np.prod(np.array(x.size()))*get_mem_space(x.dtype)/1024**2, file=file)
def track(self):
"""
Track the GPU memory usage
"""
frameinfo = inspect.stack()[1]
where_str = frameinfo.filename + ' line ' + str(frameinfo.lineno) + ': ' + frameinfo.function
with open(self.gpu_profile_fn, 'a+') as f:
if self.begin:
f.write(f"GPU Memory Track | {datetime.datetime.now():%d-%b-%y-%H:%M:%S} |"
f" Total Tensor Used Memory:{self.get_tensor_usage():<7.1f}Mb"
f" Total Allocated Memory:{self.get_allocate_usage():<7.1f}Mb\n\n")
self.begin = False
if self.print_detail is True:
ts_list = [(tensor.size(), tensor.dtype) for tensor in self.get_tensors()]
new_tensor_sizes = {(type(x),
tuple(x.size()),
ts_list.count((x.size(), x.dtype)),
np.prod(np.array(x.size()))*get_mem_space(x.dtype)/1024**2,
x.dtype) for x in self.get_tensors()}
for t, s, n, m, data_type in new_tensor_sizes - self.last_tensor_sizes:
f.write(f'+ | {str(n)} * Size:{str(s):<20} | Memory: {str(m*n)[:6]} M | {str(t):<20} | {data_type}\n')
for t, s, n, m, data_type in self.last_tensor_sizes - new_tensor_sizes:
f.write(f'- | {str(n)} * Size:{str(s):<20} | Memory: {str(m*n)[:6]} M | {str(t):<20} | {data_type}\n')
self.last_tensor_sizes = new_tensor_sizes
f.write(f"\nAt {where_str:<50}"
f" Total Tensor Used Memory:{self.get_tensor_usage():<7.1f}Mb"
f" Total Allocated Memory:{self.get_allocate_usage():<7.1f}Mb\n\n")
@@ -0,0 +1,131 @@
import argparse
import os
import yaml
global_print_hparams = True
hparams = {}
class Args:
def __init__(self, **kwargs):
for k, v in kwargs.items():
self.__setattr__(k, v)
def override_config(old_config: dict, new_config: dict):
for k, v in new_config.items():
if isinstance(v, dict) and k in old_config:
override_config(old_config[k], new_config[k])
else:
old_config[k] = v
def set_hparams(config='', exp_name='', hparams_str='', print_hparams=True, global_hparams=True, root_dir=''):
if config == '' and exp_name == '':
parser = argparse.ArgumentParser(description='')
parser.add_argument('--config', type=str, default='',
help='location of the data corpus')
parser.add_argument('--exp_name', type=str, default='', help='exp_name')
parser.add_argument('-hp', '--hparams', type=str, default='',
help='location of the data corpus')
parser.add_argument('--infer', action='store_true', help='infer')
parser.add_argument('--validate', action='store_true', help='validate')
parser.add_argument('--reset', action='store_true', help='reset hparams')
parser.add_argument('--remove', action='store_true', help='remove old ckpt')
parser.add_argument('--debug', action='store_true', help='debug')
parser.add_argument('--root_dir', type=str, default='', help='root directory of the project.')
args, unknown = parser.parse_known_args()
print("| Unknow hparams: ", unknown)
else:
args = Args(config=config, exp_name=exp_name, hparams=hparams_str,
infer=False, validate=False, reset=False, debug=False, remove=False, root_dir=root_dir)
global hparams
assert args.config != '' or args.exp_name != ''
root_dir = args.root_dir
if args.config != '':
assert os.path.exists(os.path.join(root_dir, args.config)), f'| Wrong config path! root_dir: {root_dir}, config_path: {args.config}'
config_chains = []
loaded_config = set()
def load_config(config_fn):
# deep first inheritance and avoid the second visit of one node
if not os.path.exists(os.path.join(root_dir, config_fn)):
return {}
with open(os.path.join(root_dir, config_fn)) as f:
hparams_ = yaml.safe_load(f)
loaded_config.add(config_fn)
if 'base_config' in hparams_:
ret_hparams = {}
if not isinstance(hparams_['base_config'], list):
hparams_['base_config'] = [hparams_['base_config']]
for c in hparams_['base_config']:
if c.startswith('.'):
c = f'{os.path.dirname(config_fn)}/{c}'
c = os.path.normpath(c)
if c not in loaded_config:
override_config(ret_hparams, load_config(c))
override_config(ret_hparams, hparams_)
else:
ret_hparams = hparams_
config_chains.append(config_fn)
return ret_hparams
saved_hparams = {}
args_work_dir = ''
if args.exp_name != '':
args_work_dir = os.path.join(root_dir, f'checkpoints/{args.exp_name}')
ckpt_config_path = f'{args_work_dir}/config.yaml'
if os.path.exists(ckpt_config_path):
with open(ckpt_config_path) as f:
saved_hparams_ = yaml.safe_load(f)
if saved_hparams_ is not None:
saved_hparams.update(saved_hparams_)
hparams_ = {}
if args.config != '':
hparams_.update(load_config(args.config))
if not args.reset:
hparams_.update(saved_hparams)
hparams_['work_dir'] = args_work_dir
# Support config overriding in command line. Support list type config overriding.
# Examples: --hparams="a=1,b.c=2,d=[1 1 1]"
if args.hparams != "":
for new_hparam in args.hparams.split(","):
k, v = new_hparam.split("=")
v = v.strip("\'\" ")
config_node = hparams_
for k_ in k.split(".")[:-1]:
config_node = config_node[k_]
k = k.split(".")[-1]
if v in ['True', 'False'] or type(config_node[k]) in [bool, list, dict]:
if type(config_node[k]) == list:
v = v.replace(" ", ",")
config_node[k] = eval(v)
else:
config_node[k] = type(config_node[k])(v)
if args_work_dir != '' and args.remove:
answer = input("REMOVE old checkpoint? Y/N [Default: N]: ")
if answer.lower() == "y":
pass
if args_work_dir != '' and (not os.path.exists(ckpt_config_path) or args.reset) and not args.infer:
os.makedirs(hparams_['work_dir'], exist_ok=True)
with open(ckpt_config_path, 'w') as f:
yaml.safe_dump(hparams_, f)
hparams_['infer'] = args.infer
hparams_['debug'] = args.debug
hparams_['validate'] = args.validate
hparams_['exp_name'] = args.exp_name
global global_print_hparams
if global_hparams:
hparams.clear()
hparams.update(hparams_)
if print_hparams and global_print_hparams and global_hparams:
# print('| Hparams chains: ', config_chains)
# print('| Hparams: ')
# for i, (k, v) in enumerate(sorted(hparams_.items())):
# print(f"\033[;33;m{k}\033[0m: {v}, ", end="\n" if i % 5 == 4 else "")
# print("")
global_print_hparams = False
return hparams_
@@ -0,0 +1,71 @@
import pickle
from copy import deepcopy
import numpy as np
class IndexedDataset:
def __init__(self, path, num_cache=1):
super().__init__()
self.path = path
self.data_file = None
self.data_offsets = np.load(f"{path}.idx", allow_pickle=True).item()['offsets']
self.data_file = open(f"{path}.data", 'rb', buffering=-1)
self.cache = []
self.num_cache = num_cache
def check_index(self, i):
if i < 0 or i >= len(self.data_offsets) - 1:
raise IndexError('index out of range')
def __del__(self):
if self.data_file:
self.data_file.close()
def __getitem__(self, i):
self.check_index(i)
if self.num_cache > 0:
for c in self.cache:
if c[0] == i:
return c[1]
self.data_file.seek(self.data_offsets[i])
b = self.data_file.read(self.data_offsets[i + 1] - self.data_offsets[i])
item = pickle.loads(b)
if self.num_cache > 0:
self.cache = [(i, deepcopy(item))] + self.cache[:-1]
return item
def __len__(self):
return len(self.data_offsets) - 1
class IndexedDatasetBuilder:
def __init__(self, path):
self.path = path
self.out_file = open(f"{path}.data", 'wb')
self.byte_offsets = [0]
def add_item(self, item):
s = pickle.dumps(item)
bytes = self.out_file.write(s)
self.byte_offsets.append(self.byte_offsets[-1] + bytes)
def finalize(self):
self.out_file.close()
np.save(open(f"{self.path}.idx", 'wb'), {'offsets': self.byte_offsets})
if __name__ == "__main__":
import random
from tqdm import tqdm
ds_path = '/tmp/indexed_ds_example'
size = 100
items = [{"a": np.random.normal(size=[10000, 10]),
"b": np.random.normal(size=[10000, 10])} for i in range(size)]
builder = IndexedDatasetBuilder(ds_path)
for i in tqdm(range(size)):
builder.add_item(items[i])
builder.finalize()
ds = IndexedDataset(ds_path)
for i in tqdm(range(10000)):
idx = random.randint(0, size - 1)
assert (ds[idx]['a'] == items[idx]['a']).all()
@@ -0,0 +1,53 @@
import torch
import torch.nn.functional as F
def sigmoid_focal_loss(
inputs: torch.Tensor,
targets: torch.Tensor,
alpha: float = 0.25,
gamma: float = 2,
reduction: str = "none",
) -> torch.Tensor:
"""
Loss used in RetinaNet for dense detection: https://arxiv.org/abs/1708.02002.
Args:
inputs (Tensor): A float tensor of arbitrary shape.
The predictions for each example.
targets (Tensor): A float tensor with the same shape as inputs. Stores the binary
classification label for each element in inputs
(0 for the negative class and 1 for the positive class).
alpha (float): Weighting factor in range (0,1) to balance
positive vs negative examples or -1 for ignore. Default: ``0.25``.
gamma (float): Exponent of the modulating factor (1 - p_t) to
balance easy vs hard examples. Default: ``2``.
reduction (string): ``'none'`` | ``'mean'`` | ``'sum'``
``'none'``: No reduction will be applied to the output.
``'mean'``: The output will be averaged.
``'sum'``: The output will be summed. Default: ``'none'``.
Returns:
Loss tensor with the reduction option applied.
"""
# Original implementation from https://github.com/facebookresearch/fvcore/blob/master/fvcore/nn/focal_loss.py
p = torch.sigmoid(inputs)
ce_loss = F.binary_cross_entropy_with_logits(inputs, targets, reduction="none")
p_t = p * targets + (1 - p) * (1 - targets)
loss = ce_loss * ((1 - p_t) ** gamma)
if alpha >= 0: # decrease the importance of negative samples
alpha_t = alpha * targets + (1 - alpha) * (1 - targets)
loss = alpha_t * loss
# Check reduction option and return loss accordingly
if reduction == "none":
pass
elif reduction == "mean":
loss = loss.mean()
elif reduction == "sum":
loss = loss.sum()
else:
raise ValueError(
f"Invalid Value for arg 'reduction': '{reduction} \n Supported reduction modes: 'none', 'mean', 'sum'"
)
return loss
@@ -0,0 +1,42 @@
import time
import torch
class AvgrageMeter(object):
def __init__(self):
self.reset()
def reset(self):
self.avg = 0
self.sum = 0
self.cnt = 0
def update(self, val, n=1):
self.sum += val * n
self.cnt += n
self.avg = self.sum / self.cnt
class Timer:
timer_map = {}
def __init__(self, name, enable=False):
if name not in Timer.timer_map:
Timer.timer_map[name] = 0
self.name = name
self.enable = enable
def __enter__(self):
if self.enable:
if torch.cuda.is_available():
torch.cuda.synchronize()
self.t = time.time()
def __exit__(self, exc_type, exc_val, exc_tb):
if self.enable:
if torch.cuda.is_available():
torch.cuda.synchronize()
Timer.timer_map[self.name] += time.time() - self.t
if self.enable:
print(f'[Timer] {self.name}: {Timer.timer_map[self.name]}')
@@ -0,0 +1,180 @@
import os
import traceback
from functools import partial
from tqdm import tqdm
import torch
def chunked_worker(worker_id, args_queue=None, results_queue=None, init_ctx_func=None):
ctx = init_ctx_func(worker_id) if init_ctx_func is not None else None
while True:
args = args_queue.get()
if args == '<KILL>':
return
job_idx, map_func, arg = args
try:
map_func_ = partial(map_func, ctx=ctx) if ctx is not None else map_func
if isinstance(arg, dict):
res = map_func_(**arg)
elif isinstance(arg, (list, tuple)):
res = map_func_(*arg)
else:
res = map_func_(arg)
results_queue.put((job_idx, res))
except:
traceback.print_exc()
results_queue.put((job_idx, None))
class MultiprocessManager:
def __init__(self, num_workers=None, init_ctx_func=None, multithread=False, queue_max=-1):
if multithread:
from multiprocessing.dummy import Queue, Process
else:
from multiprocessing import Queue, Process
if num_workers is None:
num_workers = int(os.getenv('N_PROC', os.cpu_count()))
self.num_workers = num_workers
self.results_queue = Queue(maxsize=-1)
self.jobs_pending = []
self.args_queue = Queue(maxsize=queue_max)
self.workers = []
self.total_jobs = 0
self.multithread = multithread
for i in range(num_workers):
if multithread:
p = Process(target=chunked_worker,
args=(i, self.args_queue, self.results_queue, init_ctx_func))
else:
p = Process(target=chunked_worker,
args=(i, self.args_queue, self.results_queue, init_ctx_func),
daemon=True)
self.workers.append(p)
p.start()
def add_job(self, func, args):
if not self.args_queue.full():
self.args_queue.put((self.total_jobs, func, args))
else:
self.jobs_pending.append((self.total_jobs, func, args))
self.total_jobs += 1
def get_results(self):
self.n_finished = 0
while self.n_finished < self.total_jobs:
while len(self.jobs_pending) > 0 and not self.args_queue.full():
self.args_queue.put(self.jobs_pending[0])
self.jobs_pending = self.jobs_pending[1:]
job_id, res = self.results_queue.get()
yield job_id, res
self.n_finished += 1
for w in range(self.num_workers):
self.args_queue.put("<KILL>")
for w in self.workers:
w.join()
def close(self):
if not self.multithread:
for w in self.workers:
w.terminate()
def __len__(self):
return self.total_jobs
def multiprocess_run_tqdm(map_func, args, num_workers=None, ordered=True, init_ctx_func=None,
multithread=False, queue_max=-1, desc=None):
for i, res in tqdm(
multiprocess_run(map_func, args, num_workers, ordered, init_ctx_func, multithread,
queue_max=queue_max),
total=len(args), desc=desc):
yield i, res
def multiprocess_run(map_func, args, num_workers=None, ordered=True, init_ctx_func=None, multithread=False,
queue_max=-1):
"""
Multiprocessing running chunked jobs.
Examples:
>>> for res in tqdm(multiprocess_run(job_func, args):
>>> print(res)
:param map_func:
:param args:
:param num_workers:
:param ordered:
:param init_ctx_func:
:param q_max_size:
:param multithread:
:return:
"""
if num_workers is None:
num_workers = int(os.getenv('N_PROC', os.cpu_count()))
manager = MultiprocessManager(num_workers, init_ctx_func, multithread, queue_max=queue_max)
for arg in args:
manager.add_job(map_func, arg)
if ordered:
n_jobs = len(args)
results = ['<WAIT>' for _ in range(n_jobs)]
i_now = 0
for job_i, res in manager.get_results():
results[job_i] = res
while i_now < n_jobs and (not isinstance(results[i_now], str) or results[i_now] != '<WAIT>'):
yield i_now, results[i_now]
results[i_now] = None
i_now += 1
else:
for job_i, res in manager.get_results():
yield job_i, res
manager.close()
# #### this is the old version of chunked_multiprocess_run
def chunked_worker_old(worker_id, map_func, args, results_queue=None, init_ctx_func=None):
ctx = init_ctx_func(worker_id) if init_ctx_func is not None else None
for job_idx, arg in args:
try:
if not isinstance(arg, tuple) and not isinstance(arg, list):
arg = [arg]
if ctx is not None:
res = map_func(*arg, ctx=ctx)
else:
res = map_func(*arg)
results_queue.put((job_idx, res))
except:
traceback.print_exc()
results_queue.put((job_idx, None))
def chunked_multiprocess_run(
map_func, args, num_workers=None, ordered=True,
init_ctx_func=None, q_max_size=1000, multithread=False):
if multithread:
from multiprocessing.dummy import Queue, Process
else:
from multiprocessing import Queue, Process
args = zip(range(len(args)), args)
args = list(args)
n_jobs = len(args)
if num_workers is None:
num_workers = int(os.getenv('N_PROC', os.cpu_count()))
results_queues = []
if ordered:
for i in range(num_workers):
results_queues.append(Queue(maxsize=q_max_size // num_workers))
else:
results_queue = Queue(maxsize=q_max_size)
for i in range(num_workers):
results_queues.append(results_queue)
workers = []
for i in range(num_workers):
args_worker = args[i::num_workers]
p = Process(target=chunked_worker_old, args=(
i, map_func, args_worker, results_queues[i], init_ctx_func), daemon=True)
workers.append(p)
p.start()
for n_finished in range(n_jobs):
results_queue = results_queues[n_finished % num_workers]
job_idx, res = results_queue.get()
assert job_idx == n_finished or not ordered, (job_idx, n_finished)
yield res
for w in workers:
w.join()
@@ -0,0 +1,74 @@
import numpy as np
import torch
import torch.nn as nn
def get_filter_2d(kernel, kernel_size, channels, no_grad=True):
# Reshape to 2d depthwise convolutional weight
kernel = kernel.view(1, 1, kernel_size, kernel_size)
kernel = kernel.repeat(channels, 1, 1, 1)
filter = nn.Conv2d(in_channels=channels, out_channels=channels, kernel_size=kernel_size, groups=channels,
bias=False, padding=kernel_size // 2)
filter.weight.data = kernel
if no_grad:
filter.weight.requires_grad = False
return filter
def get_filter_1d(kernel, kernel_size, channels, no_grad=True):
kernel = kernel.view(1, 1, kernel_size)
kernel = kernel.repeat(channels, 1, 1)
filter = nn.Conv1d(in_channels=channels, out_channels=channels, kernel_size=kernel_size, groups=channels,
bias=False, padding=kernel_size // 2)
filter.weight.data = kernel
if no_grad:
filter.weight.requires_grad = False
return filter
def get_gaussian_kernel_2d(kernel_size, sigma):
# Create a x, y coordinate grid of shape (kernel_size, kernel_size, 2)
x_coord = torch.arange(kernel_size)
x_grid = x_coord.repeat(kernel_size).view(kernel_size, kernel_size)
y_grid = x_grid.t()
xy_grid = torch.stack([x_grid, y_grid], dim=-1).float()
mean = (kernel_size - 1) / 2.
variance = sigma ** 2.
# Calculate the 2-dimensional gaussian kernel which is
# the product of two gaussian distributions for two different
# variables (in this case called x and y)
gaussian_kernel = (1. / (2. * np.pi * variance)) * torch.exp(
-torch.sum((xy_grid - mean) ** 2., dim=-1) / (2 * variance))
# Make sure sum of values in gaussian kernel equals 1.
gaussian_kernel = gaussian_kernel / torch.sum(gaussian_kernel)
return gaussian_kernel
def get_gaussian_kernel_1d(kernel_size, sigma):
x_grid = torch.arange(kernel_size)
mean = (kernel_size - 1) / 2.
variance = sigma ** 2.
gaussian_kernel = (1. / ((2. * np.pi) ** 0.5 * sigma)) * torch.exp(-(x_grid - mean) ** 2. / (2 * variance))
gaussian_kernel = gaussian_kernel / torch.sum(gaussian_kernel)
return gaussian_kernel
def get_hann_kernel_1d(kernel_size, periodic=False):
# periodic=False gives symmetric kernel, otherwise equivalent to hann(kernel_size + 1)
return torch.hann_window(kernel_size, periodic)
def get_triangle_kernel_1d(kernel_size):
kernel = torch.zeros(kernel_size)
for idx in range(kernel_size):
kernel[idx] = 1 - abs((idx - (kernel_size - 1) / 2) / ((kernel_size - 1) / 2))
return kernel
def add_gaussian_noise(tensor, mean=0, std=1):
noise = torch.randn(tensor.size()) * std + mean
noisy_tensor = tensor + noise
return noisy_tensor
@@ -0,0 +1,5 @@
import os
os.environ["OMP_NUM_THREADS"] = "1"
os.environ['TF_NUM_INTEROP_THREADS'] = '1'
os.environ['TF_NUM_INTRAOP_THREADS'] = '1'
@@ -0,0 +1,92 @@
import torch
import torch.distributed as dist
def reduce_tensors(metrics):
new_metrics = {}
for k, v in metrics.items():
if isinstance(v, torch.Tensor):
dist.all_reduce(v)
v = v / dist.get_world_size()
if type(v) is dict:
v = reduce_tensors(v)
new_metrics[k] = v
return new_metrics
def tensors_to_scalars(tensors):
if isinstance(tensors, torch.Tensor):
tensors = tensors.item()
return tensors
elif isinstance(tensors, dict):
new_tensors = {}
for k, v in tensors.items():
v = tensors_to_scalars(v)
new_tensors[k] = v
return new_tensors
elif isinstance(tensors, list):
return [tensors_to_scalars(v) for v in tensors]
else:
return tensors
def tensors_to_np(tensors):
if isinstance(tensors, dict):
new_np = {}
for k, v in tensors.items():
if isinstance(v, torch.Tensor):
v = v.cpu().numpy()
if type(v) is dict:
v = tensors_to_np(v)
new_np[k] = v
elif isinstance(tensors, list):
new_np = []
for v in tensors:
if isinstance(v, torch.Tensor):
v = v.cpu().numpy()
if type(v) is dict:
v = tensors_to_np(v)
new_np.append(v)
elif isinstance(tensors, torch.Tensor):
v = tensors
if isinstance(v, torch.Tensor):
v = v.cpu().numpy()
if type(v) is dict:
v = tensors_to_np(v)
new_np = v
else:
raise Exception(f'tensors_to_np does not support type {type(tensors)}.')
return new_np
def move_to_cpu(tensors):
ret = {}
for k, v in tensors.items():
if isinstance(v, torch.Tensor):
v = v.cpu()
if type(v) is dict:
v = move_to_cpu(v)
ret[k] = v
return ret
def move_to_cuda(batch, gpu_id=0):
# base case: object can be directly moved using `cuda` or `to`
if callable(getattr(batch, 'cuda', None)):
return batch.cuda(gpu_id, non_blocking=True)
elif callable(getattr(batch, 'to', None)):
return batch.to(torch.device('cuda', gpu_id), non_blocking=True)
elif isinstance(batch, list):
for i, x in enumerate(batch):
batch[i] = move_to_cuda(x, gpu_id)
return batch
elif isinstance(batch, tuple):
batch = list(batch)
for i, x in enumerate(batch):
batch[i] = move_to_cuda(x, gpu_id)
return tuple(batch)
elif isinstance(batch, dict):
for k, v in batch.items():
batch[k] = move_to_cuda(v, gpu_id)
return batch
return batch
@@ -0,0 +1,557 @@
import random
import subprocess
import traceback
from datetime import datetime
from torch.cuda.amp import GradScaler, autocast
import numpy as np
import torch.optim
import torch.utils.data
import copy
import logging
import os
import re
import sys
import torch
import torch.distributed as dist
import torch.multiprocessing as mp
import tqdm
from .ckpt_utils import get_last_checkpoint, get_all_ckpts
from .ddp_utils import DDP
from .hparams import hparams
from .tensor_utils import move_to_cuda
class Tee(object):
def __init__(self, name, mode):
self.file = open(name, mode)
self.stdout = sys.stdout
sys.stdout = self
def __del__(self):
sys.stdout = self.stdout
self.file.close()
def write(self, data):
self.file.write(data)
self.stdout.write(data)
def flush(self):
self.file.flush()
class Trainer:
def __init__(
self,
work_dir,
default_save_path=None,
accumulate_grad_batches=1,
max_updates=160000,
print_nan_grads=False,
val_check_interval=2000,
num_sanity_val_steps=5,
amp=False,
# tb logger
log_save_interval=100,
tb_log_interval=10,
# checkpoint
monitor_key='val_loss',
monitor_mode='min',
num_ckpt_keep=5,
save_best=True,
resume_from_checkpoint=0,
seed=1234,
debug=False,
):
os.makedirs(work_dir, exist_ok=True)
self.work_dir = work_dir
self.accumulate_grad_batches = accumulate_grad_batches
self.max_updates = max_updates
self.num_sanity_val_steps = num_sanity_val_steps
self.print_nan_grads = print_nan_grads
self.default_save_path = default_save_path
self.resume_from_checkpoint = resume_from_checkpoint if resume_from_checkpoint > 0 else None
self.seed = seed
self.debug = debug
# model and optm
self.task = None
self.optimizers = []
# trainer state
self.testing = False
self.global_step = 0
self.current_epoch = 0
self.total_batches = 0
# configure checkpoint
self.monitor_key = monitor_key
self.num_ckpt_keep = num_ckpt_keep
self.save_best = save_best
self.monitor_op = np.less if monitor_mode == 'min' else np.greater
self.best_val_results = np.Inf if monitor_mode == 'min' else -np.Inf
self.mode = 'min'
# allow int, string and gpu list
self.all_gpu_ids = [
int(x) for x in os.environ.get("CUDA_VISIBLE_DEVICES", "").split(",") if x != '']
self.num_gpus = len(self.all_gpu_ids)
self.on_gpu = self.num_gpus > 0
self.root_gpu = 0
logging.info(f'GPU available: {torch.cuda.is_available()}, GPU used: {self.all_gpu_ids}')
self.use_ddp = self.num_gpus > 1
self.proc_rank = 0
# Tensorboard logging
self.log_save_interval = log_save_interval
self.val_check_interval = val_check_interval
self.tb_log_interval = tb_log_interval
self.amp = amp
self.amp_scalar = GradScaler()
def test(self, task_cls):
self.testing = True
self.fit(task_cls)
def fit(self, task_cls):
if len(self.all_gpu_ids) > 1:
mp.spawn(self.ddp_run, nprocs=self.num_gpus, args=(task_cls, copy.deepcopy(hparams)))
else:
self.task = task_cls()
self.task.trainer = self
self.run_single_process(self.task)
return 1
def ddp_run(self, gpu_idx, task_cls, hparams_):
hparams.update(hparams_)
self.proc_rank = gpu_idx
self.init_ddp_connection(self.proc_rank, self.num_gpus)
if dist.get_rank() != 0 and not self.debug:
sys.stdout = open(os.devnull, "w")
sys.stderr = open(os.devnull, "w")
task = task_cls()
task.trainer = self
torch.cuda.set_device(gpu_idx)
self.root_gpu = gpu_idx
self.task = task
self.run_single_process(task)
def run_single_process(self, task):
"""Sanity check a few things before starting actual training.
:param task:
"""
# build model, optm and load checkpoint
if self.proc_rank == 0:
self.save_terminal_logs()
if not self.testing:
self.save_codes()
model = task.build_model()
if model is not None:
task.model = model
checkpoint, _ = get_last_checkpoint(self.work_dir, self.resume_from_checkpoint)
if checkpoint is not None:
self.restore_weights(checkpoint)
elif self.on_gpu:
task.cuda(self.root_gpu)
if not self.testing:
self.optimizers = task.configure_optimizers()
self.fisrt_epoch = True
if checkpoint is not None:
self.restore_opt_state(checkpoint)
del checkpoint
# clear cache after restore
if self.on_gpu:
torch.cuda.empty_cache()
if self.use_ddp:
self.task = self.configure_ddp(self.task)
dist.barrier()
task_ref = self.get_task_ref()
task_ref.trainer = self
task_ref.testing = self.testing
# link up experiment object
if self.proc_rank == 0:
task_ref.build_tensorboard(save_dir=self.work_dir, name='tb_logs')
else:
os.makedirs('tmp', exist_ok=True)
task_ref.build_tensorboard(save_dir='tmp', name='tb_tmp')
self.logger = task_ref.logger
try:
if self.testing:
self.run_evaluation(test=True)
else:
self.train()
except KeyboardInterrupt as e:
traceback.print_exc()
task_ref.on_keyboard_interrupt()
####################
# valid and test
####################
def run_evaluation(self, test=False):
eval_results = self.evaluate(self.task, test, tqdm_desc='Valid' if not test else 'test',
max_batches=hparams['eval_max_batches'])
if eval_results is not None and 'tb_log' in eval_results:
tb_log_output = eval_results['tb_log']
self.log_metrics_to_tb(tb_log_output)
if self.proc_rank == 0 and not test:
self.save_checkpoint(epoch=self.current_epoch, logs=eval_results)
def evaluate(self, task, test=False, tqdm_desc='Valid', max_batches=None):
if max_batches == -1:
max_batches = None
# enable eval mode
task.zero_grad()
task.eval()
torch.set_grad_enabled(False)
task_ref = self.get_task_ref()
if test:
ret = task_ref.test_start()
if ret == 'EXIT':
return
else:
task_ref.validation_start()
outputs = []
dataloader = task_ref.test_dataloader() if test else task_ref.val_dataloader()
pbar = tqdm.tqdm(dataloader, desc=tqdm_desc, total=max_batches, dynamic_ncols=True, unit='step',
disable=self.root_gpu > 0)
# give model a chance to do something with the outputs (and method defined)
for batch_idx, batch in enumerate(pbar):
if batch is None: # pragma: no cover
continue
# stop short when on fast_dev_run (sets max_batch=1)
if max_batches is not None and batch_idx >= max_batches:
break
# make dataloader_idx arg in validation_step optional
if self.on_gpu:
batch = move_to_cuda(batch, self.root_gpu)
args = [batch, batch_idx]
if self.use_ddp:
output = task(*args)
else:
if test:
output = task_ref.test_step(*args)
else:
output = task_ref.validation_step(*args)
# track outputs for collation
outputs.append(output)
# give model a chance to do something with the outputs (and method defined)
if test:
eval_results = task_ref.test_end(outputs)
else:
eval_results = task_ref.validation_end(outputs)
# enable train mode again
task.train()
torch.set_grad_enabled(True)
return eval_results
####################
# train
####################
def train(self):
task_ref = self.get_task_ref()
task_ref.on_train_start()
if self.num_sanity_val_steps > 0:
# run tiny validation (if validation defined) to make sure program won't crash during val
self.evaluate(self.task, False, 'Sanity Val', max_batches=self.num_sanity_val_steps)
# clear cache before training
if self.on_gpu:
torch.cuda.empty_cache()
dataloader = task_ref.train_dataloader()
epoch = self.current_epoch
# run all epochs
while True:
# set seed for distributed sampler (enables shuffling for each epoch)
if self.use_ddp and hasattr(dataloader.sampler, 'set_epoch'):
dataloader.sampler.set_epoch(epoch)
# update training progress in trainer and model
task_ref.current_epoch = epoch
self.current_epoch = epoch
# total batches includes multiple val checks
self.batch_loss_value = 0 # accumulated grads
# before epoch hook
task_ref.on_epoch_start()
# run epoch
train_pbar = tqdm.tqdm(dataloader, initial=self.global_step, total=float('inf'),
dynamic_ncols=True, unit='step', disable=self.root_gpu > 0)
for batch_idx, batch in enumerate(train_pbar):
if self.global_step % self.val_check_interval == 0 and not self.fisrt_epoch:
self.run_evaluation()
pbar_metrics, tb_metrics = self.run_training_batch(batch_idx, batch)
train_pbar.set_postfix(**pbar_metrics)
self.fisrt_epoch = False
# when metrics should be logged
if (self.global_step + 1) % self.tb_log_interval == 0:
# logs user requested information to logger
self.log_metrics_to_tb(tb_metrics)
self.global_step += 1
task_ref.global_step = self.global_step
if self.global_step > self.max_updates:
print("| Training end..")
break
# epoch end hook
task_ref.on_epoch_end()
epoch += 1
if self.global_step > self.max_updates:
break
task_ref.on_train_end()
def run_training_batch(self, batch_idx, batch):
if batch is None:
return {}
all_progress_bar_metrics = []
all_log_metrics = []
task_ref = self.get_task_ref()
for opt_idx, optimizer in enumerate(self.optimizers):
if optimizer is None:
continue
# make sure only the gradients of the current optimizer's paramaters are calculated
# in the training step to prevent dangling gradients in multiple-optimizer setup.
if len(self.optimizers) > 1:
for param in task_ref.parameters():
param.requires_grad = False
for group in optimizer.param_groups:
for param in group['params']:
param.requires_grad = True
# forward pass
with autocast(enabled=self.amp):
if self.on_gpu:
batch = move_to_cuda(copy.copy(batch), self.root_gpu)
args = [batch, batch_idx, opt_idx]
if self.use_ddp:
output = self.task(*args)
else:
output = task_ref.training_step(*args)
loss = output['loss']
if loss is None:
continue
progress_bar_metrics = output['progress_bar']
log_metrics = output['tb_log']
# accumulate loss
loss = loss / self.accumulate_grad_batches
# backward pass
if loss.requires_grad:
if self.amp:
self.amp_scalar.scale(loss).backward()
else:
loss.backward()
# track progress bar metrics
all_log_metrics.append(log_metrics)
all_progress_bar_metrics.append(progress_bar_metrics)
if loss is None:
continue
# nan grads
if self.print_nan_grads:
has_nan_grad = False
for name, param in task_ref.named_parameters():
if (param.grad is not None) and torch.isnan(param.grad.float()).any():
print("| NaN params: ", name, param, param.grad)
has_nan_grad = True
if has_nan_grad:
exit(0)
# gradient update with accumulated gradients
if (self.global_step + 1) % self.accumulate_grad_batches == 0:
task_ref.on_before_optimization(opt_idx)
if self.amp:
self.amp_scalar.step(optimizer)
self.amp_scalar.update()
else:
optimizer.step()
optimizer.zero_grad()
task_ref.on_after_optimization(self.current_epoch, batch_idx, optimizer, opt_idx)
# collapse all metrics into one dict
all_progress_bar_metrics = {k: v for d in all_progress_bar_metrics for k, v in d.items()}
all_log_metrics = {k: v for d in all_log_metrics for k, v in d.items()}
return all_progress_bar_metrics, all_log_metrics
####################
# load and save checkpoint
####################
def restore_weights(self, checkpoint):
# load model state
task_ref = self.get_task_ref()
for k, v in checkpoint['state_dict'].items():
getattr(task_ref, k).load_state_dict(v)
if self.on_gpu:
task_ref.cuda(self.root_gpu)
# load training state (affects trainer only)
self.best_val_results = checkpoint['checkpoint_callback_best']
self.global_step = checkpoint['global_step']
self.current_epoch = checkpoint['epoch']
task_ref.global_step = self.global_step
# wait for all models to restore weights
if self.use_ddp:
# wait for all processes to catch up
dist.barrier()
def restore_opt_state(self, checkpoint):
if self.testing:
return
# restore the optimizers
optimizer_states = checkpoint['optimizer_states']
for optimizer, opt_state in zip(self.optimizers, optimizer_states):
if optimizer is None:
return
try:
optimizer.load_state_dict(opt_state)
# move optimizer to GPU 1 weight at a time
if self.on_gpu:
for state in optimizer.state.values():
for k, v in state.items():
if isinstance(v, torch.Tensor):
state[k] = v.cuda(self.root_gpu)
except ValueError:
print("| WARMING: optimizer parameters not match !!!")
try:
if dist.is_initialized() and dist.get_rank() > 0:
return
except Exception as e:
print(e)
return
did_restore = True
return did_restore
def save_checkpoint(self, epoch, logs=None):
monitor_op = np.less
ckpt_path = f'{self.work_dir}/model_ckpt_steps_{self.global_step}.ckpt'
logging.info(f'Epoch {epoch:05d}@{self.global_step}: saving model to {ckpt_path}')
self._atomic_save(ckpt_path)
for old_ckpt in get_all_ckpts(self.work_dir)[self.num_ckpt_keep:]:
pass
current = None
if logs is not None and self.monitor_key in logs:
current = logs[self.monitor_key]
if current is not None and self.save_best:
if monitor_op(current, self.best_val_results):
best_filepath = f'{self.work_dir}/model_ckpt_best.pt'
self.best_val_results = current
logging.info(
f'Epoch {epoch:05d}@{self.global_step}: {self.monitor_key} reached {current:0.5f}. '
f'Saving model to {best_filepath}')
self._atomic_save(best_filepath)
def _atomic_save(self, filepath):
checkpoint = self.dump_checkpoint()
tmp_path = str(filepath) + ".part"
torch.save(checkpoint, tmp_path, _use_new_zipfile_serialization=False)
os.replace(tmp_path, filepath)
def dump_checkpoint(self):
checkpoint = {'epoch': self.current_epoch, 'global_step': self.global_step,
'checkpoint_callback_best': self.best_val_results}
# save optimizers
optimizer_states = []
for i, optimizer in enumerate(self.optimizers):
if optimizer is not None:
optimizer_states.append(optimizer.state_dict())
checkpoint['optimizer_states'] = optimizer_states
task_ref = self.get_task_ref()
checkpoint['state_dict'] = {
k: v.state_dict() for k, v in task_ref.named_children() if len(list(v.parameters())) > 0}
return checkpoint
####################
# DDP
####################
def configure_ddp(self, task):
task = DDP(task, device_ids=[self.root_gpu], find_unused_parameters=hparams.get('find_unused_parameters', True))
random.seed(self.seed)
np.random.seed(self.seed)
return task
def init_ddp_connection(self, proc_rank, world_size):
root_node = '127.0.0.1'
root_node = self.resolve_root_node_address(root_node)
os.environ['MASTER_ADDR'] = root_node
dist.init_process_group('nccl', rank=proc_rank, world_size=world_size)
def resolve_root_node_address(self, root_node):
if '[' in root_node:
name = root_node.split('[')[0]
number = root_node.split(',')[0]
if '-' in number:
number = number.split('-')[0]
number = re.sub('[^0-9]', '', number)
root_node = name + number
return root_node
####################
# utils
####################
def get_task_ref(self):
from .base_task import BaseTask
task: BaseTask = self.task.module if isinstance(self.task, DDP) else self.task
return task
def log_metrics_to_tb(self, metrics, step=None):
"""Logs the metric dict passed in.
:param metrics:
"""
# turn all tensors to scalars
scalar_metrics = self.metrics_to_scalars(metrics)
step = step if step is not None else self.global_step
# log actual metrics
if self.proc_rank == 0:
self.log_metrics(self.logger, scalar_metrics, step=step)
@staticmethod
def log_metrics(logger, metrics, step=None):
for k, v in metrics.items():
if isinstance(v, torch.Tensor):
v = v.item()
logger.add_scalar(k, v, step)
def metrics_to_scalars(self, metrics):
new_metrics = {}
for k, v in metrics.items():
if isinstance(v, torch.Tensor):
v = v.item()
if type(v) is dict:
v = self.metrics_to_scalars(v)
new_metrics[k] = v
return new_metrics
def save_terminal_logs(self):
t = datetime.now().strftime('%Y%m%d%H%M%S')
os.makedirs(f'{self.work_dir}/terminal_logs', exist_ok=True)
Tee(f'{self.work_dir}/terminal_logs/log_{t}.txt', 'w')
def save_codes(self):
if len(hparams['save_codes']) > 0:
t = datetime.now().strftime('%Y%m%d%H%M%S')
code_dir = f'{self.work_dir}/codes/{t}'
subprocess.check_call(f'mkdir -p "{code_dir}"', shell=True)
for c in hparams['save_codes']:
if os.path.exists(c):
subprocess.check_call(
f'rsync -aR '
f'--include="*.py" '
f'--include="*.yaml" '
f'--exclude="__pycache__" '
f'--include="*/" '
f'--exclude="*" '
f'"./{c}" "{code_dir}/"',
shell=True)
print(f"| Copied codes to {code_dir}.")
@@ -0,0 +1,74 @@
import torch
def get_focus_rate(attn, src_padding_mask=None, tgt_padding_mask=None):
'''
attn: bs x L_t x L_s
'''
if src_padding_mask is not None:
attn = attn * (1 - src_padding_mask.float())[:, None, :]
if tgt_padding_mask is not None:
attn = attn * (1 - tgt_padding_mask.float())[:, :, None]
focus_rate = attn.max(-1).values.sum(-1)
focus_rate = focus_rate / attn.sum(-1).sum(-1)
return focus_rate
def get_phone_coverage_rate(attn, src_padding_mask=None, src_seg_mask=None, tgt_padding_mask=None):
'''
attn: bs x L_t x L_s
'''
src_mask = attn.new(attn.size(0), attn.size(-1)).bool().fill_(False)
if src_padding_mask is not None:
src_mask |= src_padding_mask
if src_seg_mask is not None:
src_mask |= src_seg_mask
attn = attn * (1 - src_mask.float())[:, None, :]
if tgt_padding_mask is not None:
attn = attn * (1 - tgt_padding_mask.float())[:, :, None]
phone_coverage_rate = attn.max(1).values.sum(-1)
# phone_coverage_rate = phone_coverage_rate / attn.sum(-1).sum(-1)
phone_coverage_rate = phone_coverage_rate / (1 - src_mask.float()).sum(-1)
return phone_coverage_rate
def get_diagonal_focus_rate(attn, attn_ks, target_len, src_padding_mask=None, tgt_padding_mask=None,
band_mask_factor=5, band_width=50):
'''
attn: bx x L_t x L_s
attn_ks: shape: tensor with shape [batch_size], input_lens/output_lens
diagonal: y=k*x (k=attn_ks, x:output, y:input)
1 0 0
0 1 0
0 0 1
y>=k*(x-width) and y<=k*(x+width):1
else:0
'''
# width = min(target_len/band_mask_factor, 50)
width1 = target_len / band_mask_factor
width2 = target_len.new(target_len.size()).fill_(band_width)
width = torch.where(width1 < width2, width1, width2).float()
base = torch.ones(attn.size()).to(attn.device)
zero = torch.zeros(attn.size()).to(attn.device)
x = torch.arange(0, attn.size(1)).to(attn.device)[None, :, None].float() * base
y = torch.arange(0, attn.size(2)).to(attn.device)[None, None, :].float() * base
cond = (y - attn_ks[:, None, None] * x)
cond1 = cond + attn_ks[:, None, None] * width[:, None, None]
cond2 = cond - attn_ks[:, None, None] * width[:, None, None]
mask1 = torch.where(cond1 < 0, zero, base)
mask2 = torch.where(cond2 > 0, zero, base)
mask = mask1 * mask2
if src_padding_mask is not None:
attn = attn * (1 - src_padding_mask.float())[:, None, :]
if tgt_padding_mask is not None:
attn = attn * (1 - tgt_padding_mask.float())[:, :, None]
diagonal_attn = attn * mask
diagonal_focus_rate = diagonal_attn.sum(-1).sum(-1) / attn.sum(-1).sum(-1)
return diagonal_focus_rate, mask
@@ -0,0 +1,160 @@
from numpy import array, zeros, full, argmin, inf, ndim
from scipy.spatial.distance import cdist
from math import isinf
def dtw(x, y, dist, warp=1, w=inf, s=1.0):
"""
Computes Dynamic Time Warping (DTW) of two sequences.
:param array x: N1*M array
:param array y: N2*M array
:param func dist: distance used as cost measure
:param int warp: how many shifts are computed.
:param int w: window size limiting the maximal distance between indices of matched entries |i,j|.
:param float s: weight applied on off-diagonal moves of the path. As s gets larger, the warping path is increasingly biased towards the diagonal
Returns the minimum distance, the cost matrix, the accumulated cost matrix, and the wrap path.
"""
assert len(x)
assert len(y)
assert isinf(w) or (w >= abs(len(x) - len(y)))
assert s > 0
r, c = len(x), len(y)
if not isinf(w):
D0 = full((r + 1, c + 1), inf)
for i in range(1, r + 1):
D0[i, max(1, i - w):min(c + 1, i + w + 1)] = 0
D0[0, 0] = 0
else:
D0 = zeros((r + 1, c + 1))
D0[0, 1:] = inf
D0[1:, 0] = inf
D1 = D0[1:, 1:] # view
for i in range(r):
for j in range(c):
if (isinf(w) or (max(0, i - w) <= j <= min(c, i + w))):
D1[i, j] = dist(x[i], y[j])
C = D1.copy()
jrange = range(c)
for i in range(r):
if not isinf(w):
jrange = range(max(0, i - w), min(c, i + w + 1))
for j in jrange:
min_list = [D0[i, j]]
for k in range(1, warp + 1):
i_k = min(i + k, r)
j_k = min(j + k, c)
min_list += [D0[i_k, j] * s, D0[i, j_k] * s]
D1[i, j] += min(min_list)
if len(x) == 1:
path = zeros(len(y)), range(len(y))
elif len(y) == 1:
path = range(len(x)), zeros(len(x))
else:
path = _traceback(D0)
return D1[-1, -1], C, D1, path
def accelerated_dtw(x, y, dist, warp=1):
"""
Computes Dynamic Time Warping (DTW) of two sequences in a faster way.
Instead of iterating through each element and calculating each distance,
this uses the cdist function from scipy (https://docs.scipy.org/doc/scipy/reference/generated/scipy.spatial.distance.cdist.html)
:param array x: N1*M array
:param array y: N2*M array
:param string or func dist: distance parameter for cdist. When string is given, cdist uses optimized functions for the distance metrics.
If a string is passed, the distance function can be 'braycurtis', 'canberra', 'chebyshev', 'cityblock', 'correlation', 'cosine', 'dice', 'euclidean', 'hamming', 'jaccard', 'kulsinski', 'mahalanobis', 'matching', 'minkowski', 'rogerstanimoto', 'russellrao', 'seuclidean', 'sokalmichener', 'sokalsneath', 'sqeuclidean', 'wminkowski', 'yule'.
:param int warp: how many shifts are computed.
Returns the minimum distance, the cost matrix, the accumulated cost matrix, and the wrap path.
"""
assert len(x)
assert len(y)
if ndim(x) == 1:
x = x.reshape(-1, 1)
if ndim(y) == 1:
y = y.reshape(-1, 1)
r, c = len(x), len(y)
D0 = zeros((r + 1, c + 1))
D0[0, 1:] = inf
D0[1:, 0] = inf
D1 = D0[1:, 1:]
D0[1:, 1:] = cdist(x, y, dist)
C = D1.copy()
for i in range(r):
for j in range(c):
min_list = [D0[i, j]]
for k in range(1, warp + 1):
min_list += [D0[min(i + k, r), j],
D0[i, min(j + k, c)]]
D1[i, j] += min(min_list)
if len(x) == 1:
path = zeros(len(y)), range(len(y))
elif len(y) == 1:
path = range(len(x)), zeros(len(x))
else:
path = _traceback(D0)
return D1[-1, -1], C, D1, path
def _traceback(D):
i, j = array(D.shape) - 2
p, q = [i], [j]
while (i > 0) or (j > 0):
tb = argmin((D[i, j], D[i, j + 1], D[i + 1, j]))
if tb == 0:
i -= 1
j -= 1
elif tb == 1:
i -= 1
else: # (tb == 2):
j -= 1
p.insert(0, i)
q.insert(0, j)
return array(p), array(q)
if __name__ == '__main__':
w = inf
s = 1.0
if 1: # 1-D numeric
from sklearn.metrics.pairwise import manhattan_distances
x = [0, 0, 1, 1, 2, 4, 2, 1, 2, 0]
y = [1, 1, 1, 2, 2, 2, 2, 3, 2, 0]
dist_fun = manhattan_distances
w = 1
# s = 1.2
elif 0: # 2-D numeric
from sklearn.metrics.pairwise import euclidean_distances
x = [[0, 0], [0, 1], [1, 1], [1, 2], [2, 2], [4, 3], [2, 3], [1, 1], [2, 2], [0, 1]]
y = [[1, 0], [1, 1], [1, 1], [2, 1], [4, 3], [4, 3], [2, 3], [3, 1], [1, 2], [1, 0]]
dist_fun = euclidean_distances
else: # 1-D list of strings
from nltk.metrics.distance import edit_distance
# x = ['we', 'shelled', 'clams', 'for', 'the', 'chowder']
# y = ['class', 'too']
x = ['i', 'soon', 'found', 'myself', 'muttering', 'to', 'the', 'walls']
y = ['see', 'drown', 'himself']
# x = 'we talked about the situation'.split()
# y = 'we talked about the situation'.split()
dist_fun = edit_distance
dist, cost, acc, path = dtw(x, y, dist_fun, w=w, s=s)
# Vizualize
from matplotlib import pyplot as plt
plt.imshow(cost.T, origin='lower', cmap=plt.cm.Reds, interpolation='nearest')
plt.plot(path[0], path[1], '-o') # relation
plt.xticks(range(len(x)), x)
plt.yticks(range(len(y)), y)
plt.xlabel('x')
plt.ylabel('y')
plt.axis('tight')
if isinf(w):
plt.title('Minimum distance: {}, slope weight: {}'.format(dist, s))
else:
plt.title('Minimum distance: {}, window widht: {}, slope weight: {}'.format(dist, w, s))
plt.show()
@@ -0,0 +1,4 @@
import scipy.ndimage
def laplace_var(x):
return scipy.ndimage.laplace(x).var()
@@ -0,0 +1,102 @@
import numpy as np
import matplotlib.pyplot as plt
from numba import jit
import torch
@jit
def time_warp(costs):
dtw = np.zeros_like(costs)
dtw[0, 1:] = np.inf
dtw[1:, 0] = np.inf
eps = 1e-4
for i in range(1, costs.shape[0]):
for j in range(1, costs.shape[1]):
dtw[i, j] = costs[i, j] + min(dtw[i - 1, j], dtw[i, j - 1], dtw[i - 1, j - 1])
return dtw
def align_from_distances(distance_matrix, debug=False, return_mindist=False):
# for each position in spectrum 1, returns best match position in spectrum2
# using monotonic alignment
dtw = time_warp(distance_matrix)
i = distance_matrix.shape[0] - 1
j = distance_matrix.shape[1] - 1
results = [0] * distance_matrix.shape[0]
while i > 0 and j > 0:
results[i] = j
i, j = min([(i - 1, j), (i, j - 1), (i - 1, j - 1)], key=lambda x: dtw[x[0], x[1]])
if debug:
visual = np.zeros_like(dtw)
visual[range(len(results)), results] = 1
plt.matshow(visual)
plt.show()
if return_mindist:
return results, dtw[-1, -1]
return results
def get_local_context(input_f, max_window=32, scale_factor=1.):
# input_f: [S, 1], support numpy array or torch tensor
# return hist: [S, max_window * 2], list of list
T = input_f.shape[0]
# max_window = int(max_window * scale_factor)
derivative = [[0 for _ in range(max_window * 2)] for _ in range(T)]
for t in range(T): # travel the time series
for feat_idx in range(-max_window, max_window):
if t + feat_idx < 0 or t + feat_idx >= T:
value = 0
else:
value = input_f[t + feat_idx]
derivative[t][feat_idx + max_window] = value
return derivative
def cal_localnorm_dist(src, tgt, src_len, tgt_len):
local_src = torch.tensor(get_local_context(src))
local_tgt = torch.tensor(get_local_context(tgt, scale_factor=tgt_len / src_len))
local_norm_src = (local_src - local_src.mean(-1).unsqueeze(-1)) # / local_src.std(-1).unsqueeze(-1) # [T1, 32]
local_norm_tgt = (local_tgt - local_tgt.mean(-1).unsqueeze(-1)) # / local_tgt.std(-1).unsqueeze(-1) # [T2, 32]
dists = torch.cdist(local_norm_src[None, :, :], local_norm_tgt[None, :, :]) # [1, T1, T2]
return dists
## here is API for one sample
def LoNDTWDistance(src, tgt):
# src: [S]
# tgt: [T]
dists = cal_localnorm_dist(src, tgt, src.shape[0], tgt.shape[0]) # [1, S, T]
costs = dists.squeeze(0) # [S, T]
alignment, min_distance = align_from_distances(costs.T.cpu().detach().numpy(), return_mindist=True) # [T]
return alignment, min_distance
# if __name__ == '__main__':
# # utils from ns
# from utils.pitch_utils import denorm_f0
# from tasks.singing.fsinging import FastSingingDataset
# from utils.hparams import hparams, set_hparams
#
# set_hparams()
#
# train_ds = FastSingingDataset('test')
#
# # Test One sample case
# sample = train_ds[0]
# amateur_f0 = sample['f0']
# prof_f0 = sample['prof_f0']
#
# amateur_uv = sample['uv']
# amateur_padding = sample['mel2ph'] == 0
# prof_uv = sample['prof_uv']
# prof_padding = sample['prof_mel2ph'] == 0
# amateur_f0_denorm = denorm_f0(amateur_f0, amateur_uv, hparams, pitch_padding=amateur_padding)
# prof_f0_denorm = denorm_f0(prof_f0, prof_uv, hparams, pitch_padding=prof_padding)
# alignment, min_distance = LoNDTWDistance(amateur_f0_denorm, prof_f0_denorm)
# print(min_distance)
# python utils/pitch_distance.py --config egs/datasets/audio/molar/svc_ppg.yaml
@@ -0,0 +1,84 @@
"""
Adapted from https://github.com/Po-Hsun-Su/pytorch-ssim
"""
import torch
import torch.nn.functional as F
from torch.autograd import Variable
import numpy as np
from math import exp
def gaussian(window_size, sigma):
gauss = torch.Tensor([exp(-(x - window_size // 2) ** 2 / float(2 * sigma ** 2)) for x in range(window_size)])
return gauss / gauss.sum()
def create_window(window_size, channel):
_1D_window = gaussian(window_size, 1.5).unsqueeze(1)
_2D_window = _1D_window.mm(_1D_window.t()).float().unsqueeze(0).unsqueeze(0)
window = Variable(_2D_window.expand(channel, 1, window_size, window_size).contiguous())
return window
def _ssim(img1, img2, window, window_size, channel, size_average=True):
mu1 = F.conv2d(img1, window, padding=window_size // 2, groups=channel)
mu2 = F.conv2d(img2, window, padding=window_size // 2, groups=channel)
mu1_sq = mu1.pow(2)
mu2_sq = mu2.pow(2)
mu1_mu2 = mu1 * mu2
sigma1_sq = F.conv2d(img1 * img1, window, padding=window_size // 2, groups=channel) - mu1_sq
sigma2_sq = F.conv2d(img2 * img2, window, padding=window_size // 2, groups=channel) - mu2_sq
sigma12 = F.conv2d(img1 * img2, window, padding=window_size // 2, groups=channel) - mu1_mu2
C1 = 0.01 ** 2
C2 = 0.03 ** 2
ssim_map = ((2 * mu1_mu2 + C1) * (2 * sigma12 + C2)) / ((mu1_sq + mu2_sq + C1) * (sigma1_sq + sigma2_sq + C2))
if size_average:
return ssim_map.mean()
else:
return ssim_map.mean(1)
class SSIM(torch.nn.Module):
def __init__(self, window_size=11, size_average=True):
super(SSIM, self).__init__()
self.window_size = window_size
self.size_average = size_average
self.channel = 1
self.window = create_window(window_size, self.channel)
def forward(self, img1, img2):
(_, channel, _, _) = img1.size()
if channel == self.channel and self.window.data.type() == img1.data.type():
window = self.window
else:
window = create_window(self.window_size, channel)
if img1.is_cuda:
window = window.cuda(img1.get_device())
window = window.type_as(img1)
self.window = window
self.channel = channel
return _ssim(img1, img2, window, self.window_size, channel, self.size_average)
window = None
def ssim(img1, img2, window_size=11, size_average=True):
(_, channel, _, _) = img1.size()
global window
if window is None:
window = create_window(window_size, channel)
if img1.is_cuda:
window = window.cuda(img1.get_device())
window = window.type_as(img1)
return _ssim(img1, img2, window, window_size, channel, size_average)
@@ -0,0 +1,14 @@
import numpy as np
def print_arch(model, model_name='model'):
print(f"| {model_name} Arch: ", model)
num_params(model, model_name=model_name)
def num_params(model, print_out=True, model_name="model"):
parameters = filter(lambda p: p.requires_grad, model.parameters())
parameters = sum([np.prod(p.size()) for p in parameters]) / 1_000_000
if print_out:
print(f'| {model_name} Trainable Parameters: %.3fM' % parameters)
return parameters
@@ -0,0 +1,65 @@
class NoneSchedule(object):
def __init__(self, optimizer, lr):
self.optimizer = optimizer
self.constant_lr = lr
self.step(0)
def step(self, num_updates):
self.lr = self.constant_lr
for param_group in self.optimizer.param_groups:
param_group['lr'] = self.lr
return self.lr
def get_lr(self):
return self.optimizer.param_groups[0]['lr']
def get_last_lr(self):
return self.get_lr()
class RSQRTSchedule(NoneSchedule):
def __init__(self, optimizer, lr, warmup_updates, hidden_size, last_step=-1):
self.optimizer = optimizer
self.constant_lr = lr
self.warmup_updates = warmup_updates
self.hidden_size = hidden_size
self.lr = lr
self.last_step = last_step
for param_group in optimizer.param_groups:
param_group['lr'] = self.lr
self.step()
def step(self, num_updates=None):
if num_updates is None:
self.last_step += 1
num_updates = self.last_step
constant_lr = self.constant_lr
warmup = min(num_updates / self.warmup_updates, 1.0)
rsqrt_decay = max(self.warmup_updates, num_updates) ** -0.5
rsqrt_hidden = self.hidden_size ** -0.5
self.lr = max(constant_lr * warmup * rsqrt_decay * rsqrt_hidden, 1e-7)
for param_group in self.optimizer.param_groups:
param_group['lr'] = self.lr
return self.lr
class WarmupSchedule(NoneSchedule):
def __init__(self, optimizer, lr, warmup_updates, last_step=-1):
self.optimizer = optimizer
self.constant_lr = self.lr = lr
self.warmup_updates = warmup_updates
self.last_step = last_step
for param_group in optimizer.param_groups:
param_group['lr'] = self.lr
self.step()
def step(self, num_updates=None):
if num_updates is None:
self.last_step += 1
num_updates = self.last_step
constant_lr = self.constant_lr
warmup = min(num_updates / self.warmup_updates, 1.0)
self.lr = max(constant_lr * warmup, 1e-7)
for param_group in self.optimizer.param_groups:
param_group['lr'] = self.lr
return self.lr
@@ -0,0 +1,305 @@
from collections import defaultdict
import torch
import torch.nn.functional as F
def make_positions(tensor, padding_idx):
"""Replace non-padding symbols with their position numbers.
Position numbers begin at padding_idx+1. Padding symbols are ignored.
"""
# The series of casts and type-conversions here are carefully
# balanced to both work with ONNX export and XLA. In particular XLA
# prefers ints, cumsum defaults to output longs, and ONNX doesn't know
# how to handle the dtype kwarg in cumsum.
mask = tensor.ne(padding_idx).int()
return (
torch.cumsum(mask, dim=1).type_as(mask) * mask
).long() + padding_idx
def softmax(x, dim):
return F.softmax(x, dim=dim, dtype=torch.float32)
def sequence_mask(lengths, maxlen, dtype=torch.bool):
if maxlen is None:
maxlen = lengths.max()
mask = ~(torch.ones((len(lengths), maxlen)).to(lengths.device).cumsum(dim=1).t() > lengths).t()
mask.type(dtype)
return mask
def weights_nonzero_speech(target):
# target : B x T x mel
# Assign weight 1.0 to all labels except for padding (id=0).
dim = target.size(-1)
return target.abs().sum(-1, keepdim=True).ne(0).float().repeat(1, 1, dim)
INCREMENTAL_STATE_INSTANCE_ID = defaultdict(lambda: 0)
def _get_full_incremental_state_key(module_instance, key):
module_name = module_instance.__class__.__name__
# assign a unique ID to each module instance, so that incremental state is
# not shared across module instances
if not hasattr(module_instance, '_instance_id'):
INCREMENTAL_STATE_INSTANCE_ID[module_name] += 1
module_instance._instance_id = INCREMENTAL_STATE_INSTANCE_ID[module_name]
return '{}.{}.{}'.format(module_name, module_instance._instance_id, key)
def get_incremental_state(module, incremental_state, key):
"""Helper for getting incremental state for an nn.Module."""
full_key = _get_full_incremental_state_key(module, key)
if incremental_state is None or full_key not in incremental_state:
return None
return incremental_state[full_key]
def set_incremental_state(module, incremental_state, key, value):
"""Helper for setting incremental state for an nn.Module."""
if incremental_state is not None:
full_key = _get_full_incremental_state_key(module, key)
incremental_state[full_key] = value
def fill_with_neg_inf(t):
"""FP16-compatible function that fills a tensor with -inf."""
return t.float().fill_(float('-inf')).type_as(t)
def fill_with_neg_inf2(t):
"""FP16-compatible function that fills a tensor with -inf."""
return t.float().fill_(-1e8).type_as(t)
def select_attn(attn_logits, type='best'):
"""
:param attn_logits: [n_layers, B, n_head, T_sp, T_txt]
:return:
"""
encdec_attn = torch.stack(attn_logits, 0).transpose(1, 2)
# [n_layers * n_head, B, T_sp, T_txt]
encdec_attn = (encdec_attn.reshape([-1, *encdec_attn.shape[2:]])).softmax(-1)
if type == 'best':
indices = encdec_attn.max(-1).values.sum(-1).argmax(0)
encdec_attn = encdec_attn.gather(
0, indices[None, :, None, None].repeat(1, 1, encdec_attn.size(-2), encdec_attn.size(-1)))[0]
return encdec_attn
elif type == 'mean':
return encdec_attn.mean(0)
def make_pad_mask(lengths, xs=None, length_dim=-1):
"""Make mask tensor containing indices of padded part.
Args:
lengths (LongTensor or List): Batch of lengths (B,).
xs (Tensor, optional): The reference tensor.
If set, masks will be the same shape as this tensor.
length_dim (int, optional): Dimension indicator of the above tensor.
See the example.
Returns:
Tensor: Mask tensor containing indices of padded part.
dtype=torch.uint8 in PyTorch 1.2-
dtype=torch.bool in PyTorch 1.2+ (including 1.2)
Examples:
With only lengths.
>>> lengths = [5, 3, 2]
>>> make_non_pad_mask(lengths)
masks = [[0, 0, 0, 0 ,0],
[0, 0, 0, 1, 1],
[0, 0, 1, 1, 1]]
With the reference tensor.
>>> xs = torch.zeros((3, 2, 4))
>>> make_pad_mask(lengths, xs)
tensor([[[0, 0, 0, 0],
[0, 0, 0, 0]],
[[0, 0, 0, 1],
[0, 0, 0, 1]],
[[0, 0, 1, 1],
[0, 0, 1, 1]]], dtype=torch.uint8)
>>> xs = torch.zeros((3, 2, 6))
>>> make_pad_mask(lengths, xs)
tensor([[[0, 0, 0, 0, 0, 1],
[0, 0, 0, 0, 0, 1]],
[[0, 0, 0, 1, 1, 1],
[0, 0, 0, 1, 1, 1]],
[[0, 0, 1, 1, 1, 1],
[0, 0, 1, 1, 1, 1]]], dtype=torch.uint8)
With the reference tensor and dimension indicator.
>>> xs = torch.zeros((3, 6, 6))
>>> make_pad_mask(lengths, xs, 1)
tensor([[[0, 0, 0, 0, 0, 0],
[0, 0, 0, 0, 0, 0],
[0, 0, 0, 0, 0, 0],
[0, 0, 0, 0, 0, 0],
[0, 0, 0, 0, 0, 0],
[1, 1, 1, 1, 1, 1]],
[[0, 0, 0, 0, 0, 0],
[0, 0, 0, 0, 0, 0],
[0, 0, 0, 0, 0, 0],
[1, 1, 1, 1, 1, 1],
[1, 1, 1, 1, 1, 1],
[1, 1, 1, 1, 1, 1]],
[[0, 0, 0, 0, 0, 0],
[0, 0, 0, 0, 0, 0],
[1, 1, 1, 1, 1, 1],
[1, 1, 1, 1, 1, 1],
[1, 1, 1, 1, 1, 1],
[1, 1, 1, 1, 1, 1]]], dtype=torch.uint8)
>>> make_pad_mask(lengths, xs, 2)
tensor([[[0, 0, 0, 0, 0, 1],
[0, 0, 0, 0, 0, 1],
[0, 0, 0, 0, 0, 1],
[0, 0, 0, 0, 0, 1],
[0, 0, 0, 0, 0, 1],
[0, 0, 0, 0, 0, 1]],
[[0, 0, 0, 1, 1, 1],
[0, 0, 0, 1, 1, 1],
[0, 0, 0, 1, 1, 1],
[0, 0, 0, 1, 1, 1],
[0, 0, 0, 1, 1, 1],
[0, 0, 0, 1, 1, 1]],
[[0, 0, 1, 1, 1, 1],
[0, 0, 1, 1, 1, 1],
[0, 0, 1, 1, 1, 1],
[0, 0, 1, 1, 1, 1],
[0, 0, 1, 1, 1, 1],
[0, 0, 1, 1, 1, 1]]], dtype=torch.uint8)
"""
if length_dim == 0:
raise ValueError("length_dim cannot be 0: {}".format(length_dim))
if not isinstance(lengths, list):
lengths = lengths.tolist()
bs = int(len(lengths))
if xs is None:
maxlen = int(max(lengths))
else:
maxlen = xs.size(length_dim)
seq_range = torch.arange(0, maxlen, dtype=torch.int64)
seq_range_expand = seq_range.unsqueeze(0).expand(bs, maxlen)
seq_length_expand = seq_range_expand.new(lengths).unsqueeze(-1)
mask = seq_range_expand >= seq_length_expand
if xs is not None:
assert xs.size(0) == bs, (xs.size(0), bs)
if length_dim < 0:
length_dim = xs.dim() + length_dim
# ind = (:, None, ..., None, :, , None, ..., None)
ind = tuple(
slice(None) if i in (0, length_dim) else None for i in range(xs.dim())
)
mask = mask[ind].expand_as(xs).to(xs.device)
return mask
def make_non_pad_mask(lengths, xs=None, length_dim=-1):
"""Make mask tensor containing indices of non-padded part.
Args:
lengths (LongTensor or List): Batch of lengths (B,).
xs (Tensor, optional): The reference tensor.
If set, masks will be the same shape as this tensor.
length_dim (int, optional): Dimension indicator of the above tensor.
See the example.
Returns:
ByteTensor: mask tensor containing indices of padded part.
dtype=torch.uint8 in PyTorch 1.2-
dtype=torch.bool in PyTorch 1.2+ (including 1.2)
Examples:
With only lengths.
>>> lengths = [5, 3, 2]
>>> make_non_pad_mask(lengths)
masks = [[1, 1, 1, 1 ,1],
[1, 1, 1, 0, 0],
[1, 1, 0, 0, 0]]
With the reference tensor.
>>> xs = torch.zeros((3, 2, 4))
>>> make_non_pad_mask(lengths, xs)
tensor([[[1, 1, 1, 1],
[1, 1, 1, 1]],
[[1, 1, 1, 0],
[1, 1, 1, 0]],
[[1, 1, 0, 0],
[1, 1, 0, 0]]], dtype=torch.uint8)
>>> xs = torch.zeros((3, 2, 6))
>>> make_non_pad_mask(lengths, xs)
tensor([[[1, 1, 1, 1, 1, 0],
[1, 1, 1, 1, 1, 0]],
[[1, 1, 1, 0, 0, 0],
[1, 1, 1, 0, 0, 0]],
[[1, 1, 0, 0, 0, 0],
[1, 1, 0, 0, 0, 0]]], dtype=torch.uint8)
With the reference tensor and dimension indicator.
>>> xs = torch.zeros((3, 6, 6))
>>> make_non_pad_mask(lengths, xs, 1)
tensor([[[1, 1, 1, 1, 1, 1],
[1, 1, 1, 1, 1, 1],
[1, 1, 1, 1, 1, 1],
[1, 1, 1, 1, 1, 1],
[1, 1, 1, 1, 1, 1],
[0, 0, 0, 0, 0, 0]],
[[1, 1, 1, 1, 1, 1],
[1, 1, 1, 1, 1, 1],
[1, 1, 1, 1, 1, 1],
[0, 0, 0, 0, 0, 0],
[0, 0, 0, 0, 0, 0],
[0, 0, 0, 0, 0, 0]],
[[1, 1, 1, 1, 1, 1],
[1, 1, 1, 1, 1, 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]]], dtype=torch.uint8)
>>> make_non_pad_mask(lengths, xs, 2)
tensor([[[1, 1, 1, 1, 1, 0],
[1, 1, 1, 1, 1, 0],
[1, 1, 1, 1, 1, 0],
[1, 1, 1, 1, 1, 0],
[1, 1, 1, 1, 1, 0],
[1, 1, 1, 1, 1, 0]],
[[1, 1, 1, 0, 0, 0],
[1, 1, 1, 0, 0, 0],
[1, 1, 1, 0, 0, 0],
[1, 1, 1, 0, 0, 0],
[1, 1, 1, 0, 0, 0],
[1, 1, 1, 0, 0, 0]],
[[1, 1, 0, 0, 0, 0],
[1, 1, 0, 0, 0, 0],
[1, 1, 0, 0, 0, 0],
[1, 1, 0, 0, 0, 0],
[1, 1, 0, 0, 0, 0],
[1, 1, 0, 0, 0, 0]]], dtype=torch.uint8)
"""
return ~make_pad_mask(lengths, xs, length_dim)
def get_mask_from_lengths(lengths):
max_len = torch.max(lengths).item()
ids = torch.arange(0, max_len).to(lengths.device)
mask = (ids < lengths.unsqueeze(1)).bool()
return mask
def group_hidden_by_segs(h, seg_ids, max_len):
"""
:param h: [B, T, H]
:param seg_ids: [B, T]
:return: h_ph: [B, T_ph, H]
"""
B, T, H = h.shape
h_gby_segs = h.new_zeros([B, max_len + 1, H]).scatter_add_(1, seg_ids[:, :, None].repeat([1, 1, H]), h)
all_ones = h.new_ones(h.shape[:2])
cnt_gby_segs = h.new_zeros([B, max_len + 1]).scatter_add_(1, seg_ids, all_ones).contiguous()
h_gby_segs = h_gby_segs[:, 1:]
cnt_gby_segs = cnt_gby_segs[:, 1:]
h_gby_segs = h_gby_segs / torch.clamp(cnt_gby_segs[:, :, None], min=1)
return h_gby_segs, cnt_gby_segs
@@ -0,0 +1,21 @@
import os
import subprocess
from pathlib import Path
def link_file(from_file, to_file):
subprocess.check_call(
f'ln -s "`realpath --relative-to="{os.path.dirname(to_file)}" "{from_file}"`" "{to_file}"', shell=True)
def move_file(from_file, to_file):
subprocess.check_call(f'mv "{from_file}" "{to_file}"', shell=True)
def copy_file(from_file, to_file):
subprocess.check_call(f'cp -r "{from_file}" "{to_file}"', shell=True)
def safe_path(path):
os.makedirs(Path(path).parent, exist_ok=True)
return path
@@ -0,0 +1,51 @@
import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt
import numpy as np
import torch
LINE_COLORS = ['w', 'r', 'orange', 'k', 'cyan', 'm', 'b', 'lime', 'g', 'brown', 'navy']
def spec_to_figure(spec, vmin=None, vmax=None, title='', f0s=None, dur_info=None):
if isinstance(spec, torch.Tensor):
spec = spec.cpu().numpy()
H = spec.shape[1] // 2
fig = plt.figure(figsize=(12, 6))
plt.title(title)
plt.pcolor(spec.T, vmin=vmin, vmax=vmax)
if dur_info is not None:
assert isinstance(dur_info, dict)
txt = dur_info['txt']
dur_gt = dur_info['dur_gt']
if isinstance(dur_gt, torch.Tensor):
dur_gt = dur_gt.cpu().numpy()
dur_gt = np.cumsum(dur_gt).astype(int)
for i in range(len(dur_gt)):
shift = (i % 8) + 1
plt.text(dur_gt[i], shift * 4, txt[i])
plt.vlines(dur_gt[i], 0, H // 2, colors='b') # blue is gt
plt.xlim(0, dur_gt[-1])
if 'dur_pred' in dur_info:
dur_pred = dur_info['dur_pred']
if isinstance(dur_pred, torch.Tensor):
dur_pred = dur_pred.cpu().numpy()
dur_pred = np.cumsum(dur_pred).astype(int)
for i in range(len(dur_pred)):
shift = (i % 8) + 1
plt.text(dur_pred[i], H + shift * 4, txt[i])
plt.vlines(dur_pred[i], H, H * 1.5, colors='r') # red is pred
plt.xlim(0, max(dur_gt[-1], dur_pred[-1]))
if f0s is not None:
ax = plt.gca()
ax2 = ax.twinx()
if not isinstance(f0s, dict):
f0s = {'f0': f0s}
for i, (k, f0) in enumerate(f0s.items()):
if isinstance(f0, torch.Tensor):
f0 = f0.cpu().numpy()
ax2.plot(f0, label=k, c=LINE_COLORS[i], linewidth=1, alpha=0.5)
ax2.set_ylim(0, 1000)
ax2.legend()
return fig
@@ -0,0 +1,111 @@
import numpy as np
import torch
def regulate_real_note_itv(note_itv, note_bd, word_bd, word_durs, hop_size, audio_sample_rate):
# regulate note_itv in seconds according to the correspondence between note_bd and word_bd
assert note_itv.shape[0] == np.sum(note_bd) + 1
assert np.sum(word_bd) <= np.sum(note_bd)
assert word_durs.shape[0] == np.sum(word_bd) + 1, f"{word_durs.shape[0]} {np.sum(word_bd) + 1}"
word_bd = np.cumsum(word_bd) * word_bd # [0,1,0,0,1,0,0,0] -> [0,1,0,0,2,0,0,0]
word_itv = np.zeros((word_durs.shape[0], 2))
word_offsets = np.cumsum(word_durs)
note2words = np.zeros(note_itv.shape[0], dtype=int)
for idx in range(len(word_offsets) - 1):
word_itv[idx, 1] = word_itv[idx + 1, 0] = word_offsets[idx]
word_itv[-1, 1] = word_offsets[-1]
note_itv_secs = note_itv * hop_size / audio_sample_rate
for idx, itv in enumerate(note_itv):
start_idx, end_idx = itv
if word_bd[start_idx] > 0:
word_dur_idx = word_bd[start_idx]
note_itv_secs[idx, 0] = word_itv[word_dur_idx, 0]
note2words[idx] = word_dur_idx
if word_bd[end_idx] > 0:
word_dur_idx = word_bd[end_idx] - 1
note_itv_secs[idx, 1] = word_itv[word_dur_idx, 1]
note2words[idx] = word_dur_idx
note2words += 1 # mel2ph fashion: start from 1
return note_itv_secs, note2words
def regulate_ill_slur(notes, note_itv, note2words):
res_note2words = []
res_note_itv = []
res_notes = []
note_idx = 0
note_idx_end = 0
while True:
if note_idx > len(notes) - 1:
break
while note_idx <= note_idx_end < len(notes) and note2words[note_idx] == note2words[note_idx_end]:
note_idx_end += 1
res_note2words.append(note2words[note_idx])
res_note_itv.append(note_itv[note_idx].tolist())
res_notes.append(notes[note_idx])
for idx in range(note_idx+1, note_idx_end):
if notes[idx] == notes[idx-1]:
res_note_itv[-1][1] = note_itv[idx][1]
else:
res_note_itv.append(note_itv[idx].tolist())
res_note2words.append(note2words[idx])
res_notes.append(notes[idx])
note_idx = note_idx_end
res_notes = np.array(res_notes, dtype=notes.dtype)
res_note_itv = np.array(res_note_itv, dtype=note_itv.dtype)
res_note2words = np.array(res_note2words, dtype=note2words.dtype)
return res_notes, res_note_itv, res_note2words
def bd_to_idxs(bd):
# bd [T]
idxs = []
for idx in range(len(bd)):
if bd[idx] == 1:
idxs.append(idx)
return idxs
def bd_to_durs(bd):
# bd [T]
last_idx = 0
durs = []
for idx in range(len(bd)):
if bd[idx] == 1:
durs.append(idx - last_idx)
last_idx = idx
durs.append(len(bd) - last_idx)
return durs
def get_mel_len(wav_len, hop_size):
return (wav_len + hop_size - 1) // hop_size
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
def align_word(word_durs, mel_len, hop_size, audio_sample_rate):
mel2word = np.zeros([mel_len], int)
start_time = 0
for i_word in range(len(word_durs)):
start_frame = int(start_time * audio_sample_rate / hop_size + 0.5)
end_frame = int((start_time + word_durs[i_word]) * audio_sample_rate / hop_size + 0.5)
mel2word[start_frame:end_frame] = i_word + 1
start_time = start_time + word_durs[i_word]
dur_word = mel2token_to_dur(mel2word)
return mel2word, dur_word.tolist()
@@ -0,0 +1,9 @@
import chardet
def get_encoding(file):
with open(file, 'rb') as f:
encoding = chardet.detect(f.read())['encoding']
if encoding == 'GB2312':
encoding = 'GB18030'
return encoding
@@ -0,0 +1,263 @@
import json
import re
import six
from six.moves import range # pylint: disable=redefined-builtin
PAD = "<pad>"
EOS = "<EOS>"
UNK = "<UNK>"
SEG = "|"
PUNCS = '!,.?;:'
RESERVED_TOKENS = [PAD, EOS, UNK]
NUM_RESERVED_TOKENS = len(RESERVED_TOKENS)
PAD_ID = RESERVED_TOKENS.index(PAD) # Normally 0
EOS_ID = RESERVED_TOKENS.index(EOS) # Normally 1
UNK_ID = RESERVED_TOKENS.index(UNK) # Normally 2
if six.PY2:
RESERVED_TOKENS_BYTES = RESERVED_TOKENS
else:
RESERVED_TOKENS_BYTES = [bytes(PAD, "ascii"), bytes(EOS, "ascii")]
# Regular expression for unescaping token strings.
# '\u' is converted to '_'
# '\\' is converted to '\'
# '\213;' is converted to unichr(213)
_UNESCAPE_REGEX = re.compile(r"\\u|\\\\|\\([0-9]+);")
_ESCAPE_CHARS = set(u"\\_u;0123456789")
def strip_ids(ids, ids_to_strip):
"""Strip ids_to_strip from the end ids."""
ids = list(ids)
while ids and ids[-1] in ids_to_strip:
ids.pop()
return ids
class TextEncoder(object):
"""Base class for converting from ints to/from human readable strings."""
def __init__(self, num_reserved_ids=NUM_RESERVED_TOKENS):
self._num_reserved_ids = num_reserved_ids
@property
def num_reserved_ids(self):
return self._num_reserved_ids
def encode(self, s):
"""Transform a human-readable string into a sequence of int ids.
The ids should be in the range [num_reserved_ids, vocab_size). Ids [0,
num_reserved_ids) are reserved.
EOS is not appended.
Args:
s: human-readable string to be converted.
Returns:
ids: list of integers
"""
return [int(w) + self._num_reserved_ids for w in s.split()]
def decode(self, ids, strip_extraneous=False):
"""Transform a sequence of int ids into a human-readable string.
EOS is not expected in ids.
Args:
ids: list of integers to be converted.
strip_extraneous: bool, whether to strip off extraneous tokens
(EOS and PAD).
Returns:
s: human-readable string.
"""
if strip_extraneous:
ids = strip_ids(ids, list(range(self._num_reserved_ids or 0)))
return " ".join(self.decode_list(ids))
def decode_list(self, ids):
"""Transform a sequence of int ids into a their string versions.
This method supports transforming individual input/output ids to their
string versions so that sequence to/from text conversions can be visualized
in a human readable format.
Args:
ids: list of integers to be converted.
Returns:
strs: list of human-readable string.
"""
decoded_ids = []
for id_ in ids:
if 0 <= id_ < self._num_reserved_ids:
decoded_ids.append(RESERVED_TOKENS[int(id_)])
else:
decoded_ids.append(id_ - self._num_reserved_ids)
return [str(d) for d in decoded_ids]
@property
def vocab_size(self):
raise NotImplementedError()
class TokenTextEncoder(TextEncoder):
"""Encoder based on a user-supplied vocabulary (file or list)."""
def __init__(self,
vocab_filename,
reverse=False,
vocab_list=None,
replace_oov=None,
num_reserved_ids=NUM_RESERVED_TOKENS):
"""Initialize from a file or list, one token per line.
Handling of reserved tokens works as follows:
- When initializing from a list, we add reserved tokens to the vocab.
- When initializing from a file, we do not add reserved tokens to the vocab.
- When saving vocab files, we save reserved tokens to the file.
Args:
vocab_filename: If not None, the full filename to read vocab from. If this
is not None, then vocab_list should be None.
reverse: Boolean indicating if tokens should be reversed during encoding
and decoding.
vocab_list: If not None, a list of elements of the vocabulary. If this is
not None, then vocab_filename should be None.
replace_oov: If not None, every out-of-vocabulary token seen when
encoding will be replaced by this string (which must be in vocab).
num_reserved_ids: Number of IDs to save for reserved tokens like <EOS>.
"""
super(TokenTextEncoder, self).__init__(num_reserved_ids=num_reserved_ids)
self._reverse = reverse
self._replace_oov = replace_oov
if vocab_filename:
self._init_vocab_from_file(vocab_filename)
else:
assert vocab_list is not None
self._init_vocab_from_list(vocab_list)
self.pad_index = self.token_to_id[PAD]
self.eos_index = self.token_to_id[EOS]
self.unk_index = self.token_to_id[UNK]
self.seg_index = self.token_to_id[SEG] if SEG in self.token_to_id else self.eos_index
def encode(self, s):
"""Converts a space-separated string of tokens to a list of ids."""
sentence = s
tokens = sentence.strip().split()
if self._replace_oov is not None:
tokens = [t if t in self.token_to_id else self._replace_oov
for t in tokens]
ret = [self.token_to_id[tok] for tok in tokens]
return ret[::-1] if self._reverse else ret
def decode(self, ids, strip_eos=False, strip_padding=False):
if strip_padding and self.pad() in list(ids):
pad_pos = list(ids).index(self.pad())
ids = ids[:pad_pos]
if strip_eos and self.eos() in list(ids):
eos_pos = list(ids).index(self.eos())
ids = ids[:eos_pos]
return " ".join(self.decode_list(ids))
def decode_list(self, ids):
seq = reversed(ids) if self._reverse else ids
return [self._safe_id_to_token(i) for i in seq]
@property
def vocab_size(self):
return len(self.id_to_token)
def __len__(self):
return self.vocab_size
def _safe_id_to_token(self, idx):
return self.id_to_token.get(idx, "ID_%d" % idx)
def _init_vocab_from_file(self, filename):
"""Load vocab from a file.
Args:
filename: The file to load vocabulary from.
"""
with open(filename) as f:
tokens = [token.strip() for token in f.readlines()]
def token_gen():
for token in tokens:
yield token
self._init_vocab(token_gen(), add_reserved_tokens=False)
def _init_vocab_from_list(self, vocab_list):
"""Initialize tokens from a list of tokens.
It is ok if reserved tokens appear in the vocab list. They will be
removed. The set of tokens in vocab_list should be unique.
Args:
vocab_list: A list of tokens.
"""
def token_gen():
for token in vocab_list:
if token not in RESERVED_TOKENS:
yield token
self._init_vocab(token_gen())
def _init_vocab(self, token_generator, add_reserved_tokens=True):
"""Initialize vocabulary with tokens from token_generator."""
self.id_to_token = {}
non_reserved_start_index = 0
if add_reserved_tokens:
self.id_to_token.update(enumerate(RESERVED_TOKENS))
non_reserved_start_index = len(RESERVED_TOKENS)
self.id_to_token.update(
enumerate(token_generator, start=non_reserved_start_index))
# _token_to_id is the reverse of _id_to_token
self.token_to_id = dict((v, k) for k, v in six.iteritems(self.id_to_token))
def pad(self):
return self.pad_index
def eos(self):
return self.eos_index
def unk(self):
return self.unk_index
def seg(self):
return self.seg_index
def store_to_file(self, filename):
"""Write vocab file to disk.
Vocab files have one token per line. The file ends in a newline. Reserved
tokens are written to the vocab file as well.
Args:
filename: Full path of the file to store the vocab to.
"""
with open(filename, "w") as f:
for i in range(len(self.id_to_token)):
f.write(self.id_to_token[i] + "\n")
def sil_phonemes(self):
return [p for p in self.id_to_token.values() if is_sil_phoneme(p)]
def build_token_encoder(token_list_file):
token_list = json.load(open(token_list_file))
return TokenTextEncoder(None, vocab_list=token_list, replace_oov='<UNK>')
def is_sil_phoneme(p):
return p == '' or not p[0].isalpha()
@@ -0,0 +1,90 @@
from collections import OrderedDict
import re
import json
def remove_empty_lines(text):
"""remove empty lines"""
assert (len(text) > 0)
assert (isinstance(text, list))
text = [t.strip() for t in text]
if "" in text:
text.remove("")
return text
class TextGrid(object):
def __init__(self, text):
text = remove_empty_lines(text)
self.text = text
self.line_count = 0
self._get_type()
self._get_time_intval()
self._get_size()
self.tier_list = []
self._get_item_list()
def _extract_pattern(self, pattern, inc):
"""
Parameters
----------
pattern : regex to extract pattern
inc : increment of line count after extraction
Returns
-------
group : extracted info
"""
try:
group = re.match(pattern, self.text[self.line_count]).group(1)
self.line_count += inc
except AttributeError:
raise ValueError("File format error at line %d:%s" % (self.line_count, self.text[self.line_count]))
return group
def _get_type(self):
self.file_type = self._extract_pattern(r"File type = \"(.*)\"", 2)
def _get_time_intval(self):
self.xmin = self._extract_pattern(r"xmin = (.*)", 1)
self.xmax = self._extract_pattern(r"xmax = (.*)", 2)
def _get_size(self):
self.size = int(self._extract_pattern(r"size = (.*)", 2))
def _get_item_list(self):
"""Only supports IntervalTier currently"""
for itemIdx in range(1, self.size + 1):
tier = OrderedDict()
item_list = []
tier_idx = self._extract_pattern(r"item \[(.*)\]:", 1)
tier_class = self._extract_pattern(r"class = \"(.*)\"", 1)
if tier_class != "IntervalTier":
raise NotImplementedError("Only IntervalTier class is supported currently")
tier_name = self._extract_pattern(r"name = \"(.*)\"", 1)
tier_xmin = self._extract_pattern(r"xmin = (.*)", 1)
tier_xmax = self._extract_pattern(r"xmax = (.*)", 1)
tier_size = self._extract_pattern(r"intervals: size = (.*)", 1)
for i in range(int(tier_size)):
item = OrderedDict()
item["idx"] = self._extract_pattern(r"intervals \[(.*)\]", 1)
item["xmin"] = self._extract_pattern(r"xmin = (.*)", 1)
item["xmax"] = self._extract_pattern(r"xmax = (.*)", 1)
item["text"] = self._extract_pattern(r"text = \"(.*)\"", 1)
item_list.append(item)
tier["idx"] = tier_idx
tier["class"] = tier_class
tier["name"] = tier_name
tier["xmin"] = tier_xmin
tier["xmax"] = tier_xmax
tier["size"] = tier_size
tier["items"] = item_list
self.tier_list.append(tier)
def toJson(self):
_json = OrderedDict()
_json["file_type"] = self.file_type
_json["xmin"] = self.xmin
_json["xmax"] = self.xmax
_json["size"] = self.size
_json["tiers"] = self.tier_list
return json.dumps(_json, ensure_ascii=False, indent=2)
+320
View File
@@ -0,0 +1,320 @@
import os
import time
from dataclasses import dataclass
from typing import List, Optional
import librosa
import numpy as np
from soundfile import write
@dataclass(frozen=True)
class VocalDetectionConfig:
hop_ms: int = 20
smooth_ms: int = 200
start_ms: int = 120
end_ms: int = 200
prepad_ms: int = 80
postpad_ms: int = 120
min_len_ms: int = 1000
max_len_ms: int = 20000
short_seg_merge_gap_ms: int = 8000
small_gap_ms: int = 500
lookback_ms: int = 200
lookahead_ms: int = 100
def _moving_average(x: np.ndarray, win: int) -> np.ndarray:
if win <= 1:
return x
kernel = np.ones(win, dtype=np.float32) / float(win)
return np.convolve(x, kernel, mode="same")
def _merge_short_segments(
segments_ms: List[List[int]],
*,
min_len_ms: int,
max_len_ms: int,
short_seg_merge_gap_ms: int,
small_gap_ms: int,
) -> List[List[int]]:
if not segments_ms:
return []
merged: List[List[int]] = []
cur_start, cur_end = segments_ms[0]
for next_start, next_end in segments_ms[1:]:
cur_len = cur_end - cur_start
gap_ms = next_start - cur_end
merged_len = next_end - cur_start
should_merge = (
(cur_len < min_len_ms and gap_ms < short_seg_merge_gap_ms)
or (gap_ms < small_gap_ms and merged_len < max_len_ms)
)
if should_merge:
cur_end = next_end
continue
if (cur_end - cur_start) >= min_len_ms:
merged.append([cur_start, cur_end])
cur_start, cur_end = next_start, next_end
if (cur_end - cur_start) >= min_len_ms:
merged.append([cur_start, cur_end])
if not merged:
return segments_ms
return merged
def _voiced_to_segments(
voiced: np.ndarray,
*,
hop_ms: int,
smooth_ms: int,
start_ms: int,
end_ms: int,
prepad_ms: int,
postpad_ms: int,
max_len_ms: int,
) -> List[List[int]]:
smooth_frames = max(1, int(round(smooth_ms / hop_ms)))
smooth_voiced = _moving_average(voiced.astype(np.float32), smooth_frames)
active = smooth_voiced >= 0.5
segments: List[List[int]] = []
start_idx = None
start_frames = max(1, int(round(start_ms / hop_ms)))
end_frames = max(1, int(round(end_ms / hop_ms)))
prepad_frames = max(0, int(round(prepad_ms / hop_ms)))
postpad_frames = max(0, int(round(postpad_ms / hop_ms)))
active_count = 0
inactive_count = 0
for i, flag in enumerate(active):
if flag:
active_count += 1
inactive_count = 0
else:
inactive_count += 1
active_count = 0
if start_idx is None:
if active_count >= start_frames:
start_idx = max(0, i - start_frames + 1 - prepad_frames)
else:
if inactive_count >= end_frames:
end_idx = min(len(active) - 1, i - end_frames + 1 + postpad_frames)
start_ms_val = start_idx * hop_ms
end_ms_val = end_idx * hop_ms + hop_ms
if end_ms_val > start_ms_val:
segments.append([int(start_ms_val), int(end_ms_val)])
start_idx = None
if start_idx is not None:
start_ms_val = start_idx * hop_ms
end_idx = min(len(active) - 1, len(active) - 1 + postpad_frames)
end_ms_val = end_idx * hop_ms + hop_ms
if end_ms_val > start_ms_val:
segments.append([int(start_ms_val), int(end_ms_val)])
def _split_segment(seg: List[int]) -> List[List[int]]:
start_ms_val, end_ms_val = seg
start_frame = int(start_ms_val // hop_ms)
end_frame = int((end_ms_val - 1) // hop_ms)
end_frame = max(start_frame, min(end_frame, len(active) - 1))
best_start = None
best_len = 0
cur_start = None
cur_len = 0
for idx in range(start_frame, end_frame + 1):
if not active[idx]:
if cur_start is None:
cur_start = idx
cur_len = 1
else:
cur_len += 1
else:
if cur_start is not None and cur_len > best_len:
best_start, best_len = cur_start, cur_len
cur_start = None
cur_len = 0
if cur_start is not None and cur_len > best_len:
best_start, best_len = cur_start, cur_len
if best_start is None:
split_frame = (start_frame + end_frame) // 2
else:
split_frame = best_start + best_len // 2
split_ms = split_frame * hop_ms
if split_ms <= start_ms_val:
split_ms = start_ms_val + hop_ms
if split_ms >= end_ms_val:
split_ms = end_ms_val - hop_ms
if split_ms <= start_ms_val or split_ms >= end_ms_val:
return [seg]
return [[start_ms_val, int(split_ms)], [int(split_ms), end_ms_val]]
queue = segments[:]
segments = []
while queue:
seg = queue.pop(0)
if (seg[1] - seg[0]) <= max_len_ms:
segments.append(seg)
continue
parts = _split_segment(seg)
if len(parts) == 1:
segments.append(seg)
else:
queue = parts + queue
return segments
class VocalDetector:
"""Detect vocal segments based on f0 voiced decisions.
This component consumes a precomputed ``*_f0.npy`` track and
produces vocal segments (and cuts wav files) for downstream
transcription or singing voice tasks.
"""
def __init__(
self,
cut_wavs_output_dir: str = "cut_wavs",
config: VocalDetectionConfig | None = None,
*,
verbose: bool = True,
):
"""Initialize the vocal detector.
Args:
cut_wavs_output_dir: Directory to save cut wav segments.
config: Detection configuration; uses :class:`VocalDetectionConfig` by default.
verbose: Whether to print verbose logs.
"""
self.cut_wavs_output_dir = cut_wavs_output_dir
self.config = config or VocalDetectionConfig()
self.verbose = verbose
if self.verbose:
print(
"[vocal detection] init success:",
f"cut_wavs_output_dir={self.cut_wavs_output_dir}",
f"hop_ms={self.config.hop_ms}",
)
def process(self, audio_path: str, f0: np.ndarray, *, verbose: Optional[bool] = None) -> List[dict]:
"""Run vocal detection on a single wav.
Args:
audio_path: Path to the input wav file.
f0: The f0 contour to use for vocal detection.
verbose: Override instance-level verbose flag for this call.
Returns:
A list of segment metadata dicts with fields like
``item_name``, ``wav_fn``, ``start_time_ms``, ``end_time_ms``.
"""
verbose = self.verbose if verbose is None else verbose
if verbose:
print(f"[vocal detection] process: start: {audio_path}")
t0 = time.time()
os.makedirs(self.cut_wavs_output_dir, exist_ok=True)
base_name = os.path.basename(audio_path)
base_name_no_ext = os.path.splitext(base_name)[0]
voiced = f0 > 0
segments_ms = _voiced_to_segments(
voiced,
hop_ms=self.config.hop_ms,
smooth_ms=self.config.smooth_ms,
start_ms=self.config.start_ms,
end_ms=self.config.end_ms,
prepad_ms=self.config.prepad_ms,
postpad_ms=self.config.postpad_ms,
max_len_ms=self.config.max_len_ms,
)
if verbose:
print(f"[vocal detection] segments(before_merge)={len(segments_ms)}")
segments_ms = _merge_short_segments(
segments_ms,
min_len_ms=self.config.min_len_ms,
max_len_ms=self.config.max_len_ms,
short_seg_merge_gap_ms=self.config.short_seg_merge_gap_ms,
small_gap_ms=self.config.small_gap_ms,
)
if verbose:
print(f"[vocal detection] segments(after_merge)={len(segments_ms)}")
y, sr = librosa.load(audio_path, sr=None, mono=True)
# Apply global lookback/lookahead in milliseconds
lookback_ms = self.config.lookback_ms
lookahead_ms = self.config.lookahead_ms
adjusted_segments: List[List[int]] = []
prev_end = 0
for start_ms, end_ms in segments_ms:
start_ms = max(0, start_ms - lookback_ms)
end_ms = min(end_ms + lookahead_ms, int(y.shape[0] / sr * 1000))
# Enforce non-overlap with previous segment, move backward the previous one.
if start_ms < prev_end and len(adjusted_segments) > 0:
adjusted_segments[-1][1] = start_ms
adjusted_segments.append([start_ms, end_ms])
prev_end = end_ms
segment_infos = []
for idx, (start_ms, end_ms) in enumerate(adjusted_segments):
if end_ms - start_ms > self.config.max_len_ms:
start_ms = end_ms - self.config.max_len_ms
key = f"{base_name_no_ext}_{idx}"
start_sample = librosa.time_to_samples(start_ms / 1000, sr=sr)
end_sample = librosa.time_to_samples(end_ms / 1000, sr=sr)
segment = y[start_sample:end_sample]
write(f"{self.cut_wavs_output_dir}/{key}.wav", segment, sr)
segment_infos.append(
{
"item_name": key,
"wav_fn": f"{self.cut_wavs_output_dir}/{key}.wav",
"start_time_ms": int(start_sample * 1000 / sr),
"end_time_ms": int(end_sample * 1000 / sr),
"origin_wav_fn": audio_path,
"duration": int((end_sample - start_sample) * 1000 / sr),
}
)
if verbose:
dt = time.time() - t0
print(
"[vocal detection] process: done:",
f"n_segments={len(segment_infos)}",
f"time={dt:.3f}s",
)
return segment_infos
if __name__ == "__main__":
m = VocalDetector(cut_wavs_output_dir="outputs/transcription/cut_wavs")
segment_infos = m.process("./outputs/transcription/test.wav")
print(segment_infos)
+225
View File
@@ -0,0 +1,225 @@
# https://github.com/ZFTurbo/Music-Source-Separation-Training
# https://huggingface.co/becruily/mel-band-roformer-karaoke/blob/main/mel_band_roformer_karaoke_becruily.ckpt
# https://huggingface.co/anvuew/dereverb_mel_band_roformer/blob/main/dereverb_mel_band_roformer_anvuew_sdr_19.1729.ckpt
from __future__ import annotations
from dataclasses import dataclass
from typing import Any, Dict, Optional, Tuple
import librosa
import sys
import os
import time
import torch
import numpy as np
from .utils.audio_utils import normalize_audio, denormalize_audio
from .utils.settings import get_model_from_config, parse_args_inference
from .utils.model_utils import demix
from .utils.model_utils import prefer_target_instrument, apply_tta, load_start_checkpoint
def process(mix, model, args, config, device):
instruments = prefer_target_instrument(config)[:]
# If mono audio we must adjust it depending on model
if len(mix.shape) == 1:
mix = np.expand_dims(mix, axis=0)
if 'num_channels' in config.audio:
if config.audio['num_channels'] == 2:
# print(f'Convert mono track to stereo...')
mix = np.concatenate([mix, mix], axis=0)
if 'normalize' in config.inference:
if config.inference['normalize'] is True:
mix, norm_params = normalize_audio(mix)
waveforms_orig = demix(config, model, mix, device, model_type=args.model_type, pbar=not args.disable_detailed_pbar)
instr = 'vocals' if 'vocals' in instruments else instruments[0]
estimates = waveforms_orig[instr]
if 'normalize' in config.inference:
if config.inference['normalize'] is True:
estimates = denormalize_audio(estimates, norm_params)
return estimates
def build_model(args):
model, config = get_model_from_config(args.model_type, args.config_path)
load_start_checkpoint(args, model, None, type_='inference')
return model, config
def build_models(dict_args):
args = parse_args_inference(dict_args)
########## load model ##########
torch.backends.cudnn.benchmark = True
args.config_path = args.sep_config_path
args.start_check_point = args.sep_start_check_point
sep_model, sep_config = build_model(args)
args.config_path = args.der_config_path
args.start_check_point = args.der_start_check_point
dereverb_model, dereverb_config = build_model(args)
sep_model = sep_model
dereverb_model = dereverb_model
return sep_model, sep_config, dereverb_model, dereverb_config, args
def main(args, sep_model=None, sep_config=None, dereverb_model=None, dereverb_config=None, device=None):
######## process data ##########
sample_rate = getattr(sep_config.audio, 'sample_rate', 44100)
path = args.input_path
mix, _ = librosa.load(path, sr=sample_rate, mono=False)
vocals = process(mix, sep_model, args, sep_config, device)
dereverbed_vocals = process(vocals.mean(0), dereverb_model, args, dereverb_config, device)
accompaniment = mix - dereverbed_vocals
return mix, vocals, dereverbed_vocals, accompaniment, sample_rate
@dataclass
class VocalSeparationOutputs:
"""Vocal extraction output container."""
mix: np.ndarray
vocals: np.ndarray
vocals_dereverbed: np.ndarray
accompaniment: np.ndarray
sample_rate: int
class VocalSeparator:
"""Vocal separation and dereverb wrapper.
Wraps the karaoke separation and dereverb models from the
ZFTurbo Music Source Separation project and exposes a simple
:py:meth:`process` API that returns mix/vocals/dereverbed/accompaniment.
"""
def __init__(
self,
sep_model_path: str,
sep_config_path: str,
der_model_path: str,
der_config_path: str,
*,
model_type: str = "mel_band_roformer",
disable_detailed_pbar: bool = True,
device: str = "cuda",
verbose: bool = True,
):
"""Initialize the vocal separator.
Args:
device: Torch device string, e.g. ``"cuda:0"``.
model_type: Separation model type key.
sep_config_path: Config path for separation model.
sep_start_check_point: Checkpoint path for separation model.
der_config_path: Config path for dereverb model.
der_start_check_point: Checkpoint path for dereverb model.
disable_detailed_pbar: Disable detailed progress bars in underlying utils.
verbose: Whether to print verbose logs.
"""
# Match original script args schema
args_dict: Dict[str, Any] = {
"model_type": model_type,
"disable_detailed_pbar": disable_detailed_pbar,
"sep_config_path": sep_config_path,
"sep_start_check_point": sep_model_path,
"der_config_path": der_config_path,
"der_start_check_point": der_model_path,
}
if verbose:
print("[vocal extraction] init: start")
sep_model, sep_config, dereverb_model, dereverb_config, args = build_models(args_dict)
sep_model = sep_model.to(device)
dereverb_model = dereverb_model.to(device)
self.sep_model = sep_model
self.sep_config = sep_config
self.dereverb_model = dereverb_model
self.dereverb_config = dereverb_config
self.device = device
self.args = args
self.verbose = verbose
if verbose:
print(
"[vocal extraction] init success: sep=loaded, dereverb=loaded, device=",
device,
)
def process(self, input_path: str, *, verbose: Optional[bool] = None) -> VocalSeparationOutputs:
"""Separate a single audio file into sources.
Args:
input_path: Path to the mixture wav.
verbose: Override instance-level verbose flag for this call.
Returns:
:class:`VocalSeparationOutputs` containing mix, vocals,
dereverbed vocals, accompaniment and sample rate.
"""
verbose = self.verbose if verbose is None else verbose
if verbose:
print(f"[vocal extraction] process_file: start: {input_path}")
t0 = time.time()
self.args.input_path = input_path
mix, vocals, dereverbed, accompaniment, sample_rate = main(
self.args,
self.sep_model,
self.sep_config,
self.dereverb_model,
self.dereverb_config,
torch.device(self.device) if not isinstance(self.device, torch.device) else self.device,
)
if verbose:
dt = time.time() - t0
print(
"[vocal extraction] process_file: done:",
f"sr={sample_rate}",
f"mix={getattr(mix, 'shape', None)}",
f"vocals={getattr(vocals, 'shape', None)}",
f"dereverbed={getattr(dereverbed, 'shape', None)}",
f"acc={getattr(accompaniment, 'shape', None)}",
f"time={dt:.3f}s",
)
return VocalSeparationOutputs(
mix=mix,
vocals=vocals,
vocals_dereverbed=dereverbed,
accompaniment=accompaniment,
sample_rate=sample_rate,
)
if __name__ == "__main__":
m = VocalSeparator(
sep_model_path="pretrained_models/mel-band-roformer-karaoke/mel_band_roformer_karaoke_becruily.ckpt",
sep_config_path="pretrained_models/mel-band-roformer-karaoke/config_karaoke_becruily.yaml",
der_model_path="pretrained_models/dereverb_mel_band_roformer/dereverb_mel_band_roformer_anvuew_sdr_19.1729.ckpt",
der_config_path="pretrained_models/dereverb_mel_band_roformer/dereverb_mel_band_roformer_anvuew.yaml",
device="cuda"
)
out = m.process("example/test/separation_test.mp3")
print(out.vocals_dereverbed.shape)
@@ -0,0 +1,2 @@
from .bs_roformer import BSRoformer
from .mel_band_roformer import MelBandRoformer
@@ -0,0 +1,126 @@
from functools import wraps
from packaging import version
from collections import namedtuple
import os
import torch
from torch import nn, einsum
import torch.nn.functional as F
from einops import rearrange, reduce
# constants
FlashAttentionConfig = namedtuple('FlashAttentionConfig', ['enable_flash', 'enable_math', 'enable_mem_efficient'])
# helpers
def exists(val):
return val is not None
def default(v, d):
return v if exists(v) else d
def once(fn):
called = False
@wraps(fn)
def inner(x):
nonlocal called
if called:
return
called = True
return fn(x)
return inner
print_once = once(print)
# main class
class Attend(nn.Module):
def __init__(
self,
dropout = 0.,
flash = False,
scale = None
):
super().__init__()
self.scale = scale
self.dropout = dropout
self.attn_dropout = nn.Dropout(dropout)
self.flash = flash
assert not (flash and version.parse(torch.__version__) < version.parse('2.0.0')), 'in order to use flash attention, you must be using pytorch 2.0 or above'
# determine efficient attention configs for cuda and cpu
self.cpu_config = FlashAttentionConfig(True, True, True)
self.cuda_config = None
if not torch.cuda.is_available() or not flash:
return
device_properties = torch.cuda.get_device_properties(torch.device('cuda'))
device_version = version.parse(f'{device_properties.major}.{device_properties.minor}')
if device_version >= version.parse('8.0'):
if os.name == 'nt':
print_once('Windows OS detected, using math or mem efficient attention if input tensor is on cuda')
self.cuda_config = FlashAttentionConfig(False, True, True)
else:
print_once('GPU Compute Capability equal or above 8.0, using flash attention if input tensor is on cuda')
self.cuda_config = FlashAttentionConfig(True, False, False)
else:
print_once('GPU Compute Capability below 8.0, using math or mem efficient attention if input tensor is on cuda')
self.cuda_config = FlashAttentionConfig(False, True, True)
def flash_attn(self, q, k, v):
_, heads, q_len, _, k_len, is_cuda, device = *q.shape, k.shape[-2], q.is_cuda, q.device
if exists(self.scale):
default_scale = q.shape[-1] ** -0.5
q = q * (self.scale / default_scale)
# Check if there is a compatible device for flash attention
config = self.cuda_config if is_cuda else self.cpu_config
# pytorch 2.0 flash attn: q, k, v, mask, dropout, softmax_scale
with torch.backends.cuda.sdp_kernel(**config._asdict()):
out = F.scaled_dot_product_attention(
q, k, v,
dropout_p = self.dropout if self.training else 0.
)
return out
def forward(self, q, k, v):
"""
einstein notation
b - batch
h - heads
n, i, j - sequence length (base sequence length, source, target)
d - feature dimension
"""
q_len, k_len, device = q.shape[-2], k.shape[-2], q.device
scale = default(self.scale, q.shape[-1] ** -0.5)
if self.flash:
return self.flash_attn(q, k, v)
# similarity
sim = einsum(f"b h i d, b h j d -> b h i j", q, k) * scale
# attention
attn = sim.softmax(dim=-1)
attn = self.attn_dropout(attn)
# aggregate values
out = einsum(f"b h i j, b h j d -> b h i d", attn, v)
return out

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