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
@@ -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)