Initial commit
This commit is contained in:
@@ -0,0 +1,522 @@
|
||||
# https://github.com/RickyL-2000/ROSVOT
|
||||
import math
|
||||
import sys
|
||||
import traceback
|
||||
import json
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
import librosa
|
||||
import numpy as np
|
||||
import torch
|
||||
import matplotlib.pyplot as plt
|
||||
|
||||
from .utils.os_utils import safe_path
|
||||
from .utils.commons.hparams import set_hparams
|
||||
from .utils.commons.ckpt_utils import load_ckpt
|
||||
from .utils.commons.dataset_utils import pad_or_cut_xd
|
||||
from .utils.audio.mel import MelNet
|
||||
from .utils.audio.pitch_utils import (
|
||||
norm_interp_f0,
|
||||
denorm_f0,
|
||||
f0_to_coarse,
|
||||
boundary2Interval,
|
||||
save_midi,
|
||||
midi_to_hz,
|
||||
)
|
||||
from .utils.rosvot_utils import (
|
||||
get_mel_len,
|
||||
align_word,
|
||||
regulate_real_note_itv,
|
||||
regulate_ill_slur,
|
||||
bd_to_durs,
|
||||
)
|
||||
from .modules.pe.rmvpe import RMVPE
|
||||
from .modules.rosvot.rosvot import MidiExtractor, WordbdExtractor
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def infer_sample(
|
||||
item: Dict[str, Any],
|
||||
hparams: Dict[str, Any],
|
||||
models: Dict[str, Any],
|
||||
device: torch.device,
|
||||
*,
|
||||
save_dir: Optional[str] = None,
|
||||
apply_rwbd: Optional[bool] = None,
|
||||
# outputs
|
||||
save_plot: bool = False,
|
||||
no_save_midi: bool = True,
|
||||
no_save_npy: bool = True,
|
||||
verbose: bool = False,
|
||||
) -> Dict[str, Any]:
|
||||
if "item_name" not in item or "wav_fn" not in item:
|
||||
raise ValueError('item must contain keys: "item_name" and "wav_fn"')
|
||||
|
||||
item_name = item["item_name"]
|
||||
wav_src = item["wav_fn"]
|
||||
|
||||
# Decide RWBD usage
|
||||
if apply_rwbd is None:
|
||||
apply_rwbd_ = ("word_durs" not in item)
|
||||
else:
|
||||
apply_rwbd_ = bool(apply_rwbd)
|
||||
|
||||
# Models
|
||||
model = models["model"]
|
||||
mel_net = models["mel_net"]
|
||||
pe = models.get("pe")
|
||||
wbd_predictor = models.get("wbd_predictor")
|
||||
|
||||
if wbd_predictor is None and apply_rwbd_:
|
||||
raise ValueError("apply_rwbd is True but wbd_predictor model is not provided in models")
|
||||
|
||||
# ---- Prepare Data ----
|
||||
if isinstance(wav_src, str):
|
||||
wav, _ = librosa.core.load(wav_src, sr=hparams["audio_sample_rate"])
|
||||
else:
|
||||
wav = wav_src
|
||||
if not isinstance(wav, np.ndarray):
|
||||
wav = np.asarray(wav)
|
||||
wav = wav.astype(np.float32)
|
||||
|
||||
# Calculate timestamps and alignment lengths
|
||||
wav_len_samples = wav.shape[-1]
|
||||
mel_len = get_mel_len(wav_len_samples, hparams["hop_size"])
|
||||
|
||||
# Word boundary preparation
|
||||
mel2word = None
|
||||
word_durs_filtered = None
|
||||
|
||||
if not apply_rwbd_:
|
||||
if "word_durs" not in item:
|
||||
raise ValueError('apply_rwbd=False but item has no "word_durs"')
|
||||
|
||||
wd_raw = list(item["word_durs"])
|
||||
min_word_dur = hparams.get("min_word_dur", 20) / 1000
|
||||
word_durs_filtered = []
|
||||
|
||||
for i, wd in enumerate(wd_raw):
|
||||
if wd < min_word_dur:
|
||||
if i == 0 and len(wd_raw) > 1:
|
||||
wd_raw[i + 1] += wd
|
||||
elif len(word_durs_filtered) > 0:
|
||||
word_durs_filtered[-1] += wd
|
||||
else:
|
||||
word_durs_filtered.append(wd)
|
||||
|
||||
mel2word, _ = align_word(word_durs_filtered, mel_len, hparams["hop_size"], hparams["audio_sample_rate"])
|
||||
mel2word = np.asarray(mel2word)
|
||||
if mel2word.size > 0 and mel2word[0] == 0:
|
||||
mel2word = mel2word + 1
|
||||
|
||||
mel2word_len = int(np.sum(mel2word > 0))
|
||||
real_len = min(mel_len, mel2word_len)
|
||||
else:
|
||||
real_len = min(mel_len, hparams["max_frames"])
|
||||
|
||||
T = math.ceil(min(real_len, hparams["max_frames"]) / hparams["frames_multiple"]) * hparams["frames_multiple"]
|
||||
|
||||
# ---- Input Tensors & Padding ----
|
||||
target_samples = T * hparams["hop_size"]
|
||||
wav_t = torch.from_numpy(wav).float().to(device).unsqueeze(0) # [1, L]
|
||||
if wav_t.shape[-1] < target_samples:
|
||||
wav_t = pad_or_cut_xd(wav_t, target_samples, 1)
|
||||
|
||||
# ---- Pitch Extraction ----
|
||||
if pe is not None:
|
||||
f0s, uvs = pe.get_pitch_batch(
|
||||
wav_t,
|
||||
sample_rate=hparams["audio_sample_rate"],
|
||||
hop_size=hparams["hop_size"],
|
||||
lengths=[real_len],
|
||||
fmax=hparams["f0_max"],
|
||||
fmin=hparams["f0_min"],
|
||||
)
|
||||
f0_1d, uv_1d = norm_interp_f0(f0s[0][:T])
|
||||
f0_t = pad_or_cut_xd(torch.FloatTensor(f0_1d).to(device), T, 0).unsqueeze(0)
|
||||
uv_t = pad_or_cut_xd(torch.FloatTensor(uv_1d).to(device), T, 0).long().unsqueeze(0)
|
||||
pitch_coarse = f0_to_coarse(denorm_f0(f0_t, uv_t)).to(device)
|
||||
f0_np = denorm_f0(f0_t, uv_t)[0].detach().cpu().numpy()[:real_len]
|
||||
else:
|
||||
f0_t = uv_t = pitch_coarse = None
|
||||
f0_np = None
|
||||
|
||||
# ---- Mel Extraction ----
|
||||
mel = mel_net(wav_t) # [1, T_padded, C]
|
||||
mel = pad_or_cut_xd(mel, T, 1)
|
||||
|
||||
# Construct non-padding mask
|
||||
mel_nonpadding_mask = torch.zeros(1, T, device=device)
|
||||
mel_nonpadding_mask[:, :real_len] = 1.0
|
||||
|
||||
# Apply mask to mel (zero out padding)
|
||||
mel = (mel.transpose(1, 2) * mel_nonpadding_mask.unsqueeze(1)).transpose(1, 2)
|
||||
# Re-calculate non_padding bool mask
|
||||
mel_nonpadding = mel.abs().sum(-1) > 0
|
||||
|
||||
# ---- Word Boundary ----
|
||||
word_durs_used = None
|
||||
if apply_rwbd_:
|
||||
mel_input = mel[:, :, : hparams.get("wbd_use_mel_bins", 80)]
|
||||
wbd_outputs = wbd_predictor(
|
||||
mel=mel_input,
|
||||
pitch=pitch_coarse,
|
||||
uv=uv_t,
|
||||
non_padding=mel_nonpadding,
|
||||
train=False,
|
||||
)
|
||||
word_bd = wbd_outputs["word_bd_pred"] # [1, T]
|
||||
else:
|
||||
# Construct word_bd from provided durs
|
||||
mel2word_t = pad_or_cut_xd(torch.LongTensor(mel2word).to(device), T, 0)
|
||||
word_bd = torch.zeros_like(mel2word_t)
|
||||
# Vectorized check
|
||||
word_bd[1:] = (mel2word_t[1:] != mel2word_t[:-1]).long()
|
||||
word_bd[real_len:] = 0
|
||||
word_bd = word_bd.unsqueeze(0) # [1, T]
|
||||
|
||||
word_durs_used = np.array(word_durs_filtered)
|
||||
|
||||
# ---- Main Inference ----
|
||||
mel_input = mel[:, :, : hparams.get("use_mel_bins", 80)]
|
||||
outputs = model(
|
||||
mel=mel_input,
|
||||
word_bd=word_bd,
|
||||
pitch=pitch_coarse,
|
||||
uv=uv_t,
|
||||
non_padding=mel_nonpadding,
|
||||
train=False,
|
||||
)
|
||||
|
||||
note_lengths = outputs["note_lengths"].detach().cpu().numpy()
|
||||
note_bd_pred = outputs["note_bd_pred"][0].detach().cpu().numpy()[:real_len]
|
||||
note_pred = outputs["note_pred"][0].detach().cpu().numpy()[: note_lengths[0]]
|
||||
note_bd_logits = torch.sigmoid(outputs["note_bd_logits"])[0].detach().cpu().numpy()[:real_len]
|
||||
|
||||
if note_pred.shape == (0,):
|
||||
if verbose:
|
||||
print(f"skip {item_name}: no notes detected")
|
||||
return {
|
||||
"item_name": item_name,
|
||||
"pitches": [],
|
||||
"note_durs": [],
|
||||
"note2words": None,
|
||||
}
|
||||
|
||||
# ---- Post-Processing & Regulation ----
|
||||
note_itv_pred = boundary2Interval(note_bd_pred)
|
||||
note2words = None
|
||||
|
||||
if apply_rwbd_:
|
||||
word_bd_np = outputs['word_bd_pred'][0].detach().cpu().numpy()[:real_len]
|
||||
word_durs_derived = np.array(bd_to_durs(word_bd_np)) * hparams['hop_size'] / hparams['audio_sample_rate']
|
||||
word_durs_for_reg = word_durs_derived
|
||||
word_bd_for_reg = word_bd_np
|
||||
else:
|
||||
word_bd_for_reg = word_bd[0].detach().cpu().numpy()[:real_len]
|
||||
word_durs_for_reg = word_durs_used
|
||||
|
||||
should_regulate = hparams.get("infer_regulate_real_note_itv", True) and (not apply_rwbd_)
|
||||
|
||||
if should_regulate and (word_durs_for_reg is not None):
|
||||
try:
|
||||
note_itv_pred_secs, note2words = regulate_real_note_itv(
|
||||
note_itv_pred,
|
||||
note_bd_pred,
|
||||
word_bd_for_reg,
|
||||
word_durs_for_reg,
|
||||
hparams["hop_size"],
|
||||
hparams["audio_sample_rate"],
|
||||
)
|
||||
note_pred, note_itv_pred_secs, note2words = regulate_ill_slur(note_pred, note_itv_pred_secs, note2words)
|
||||
except Exception as err:
|
||||
if verbose:
|
||||
_, exc_value, exc_tb = sys.exc_info()
|
||||
tb = traceback.extract_tb(exc_tb)[-1]
|
||||
print(f"postprocess failed: {err}: {exc_value} in {tb[0]}:{tb[1]} '{tb[2]}' in {tb[3]}")
|
||||
# Fallback
|
||||
note_itv_pred_secs = note_itv_pred * hparams["hop_size"] / hparams["audio_sample_rate"]
|
||||
note2words = None
|
||||
else:
|
||||
note_itv_pred_secs = note_itv_pred * hparams["hop_size"] / hparams["audio_sample_rate"]
|
||||
|
||||
# ---- Output ----
|
||||
note_durs = [float((itv[1] - itv[0])) for itv in note_itv_pred_secs]
|
||||
|
||||
out = {
|
||||
"item_name": item_name,
|
||||
"pitches": note_pred.tolist(),
|
||||
"note_durs": note_durs,
|
||||
"note2words": note2words.tolist() if note2words is not None else None,
|
||||
}
|
||||
|
||||
# ---- Saving ----
|
||||
if save_dir is not None:
|
||||
save_dir_path = Path(save_dir)
|
||||
save_dir_path.mkdir(parents=True, exist_ok=True)
|
||||
fn = str(item_name)
|
||||
|
||||
if not no_save_midi:
|
||||
save_midi(note_pred, note_itv_pred_secs, safe_path(save_dir_path / "midi" / f"{fn}.mid"))
|
||||
|
||||
if not no_save_npy:
|
||||
np.save(safe_path(save_dir_path / "npy" / f"[note]{fn}.npy"), out, allow_pickle=True)
|
||||
|
||||
if save_plot:
|
||||
fig = plt.figure()
|
||||
if f0_np is not None:
|
||||
plt.plot(f0_np, color="red", label="f0")
|
||||
|
||||
midi_pred = np.zeros(note_bd_pred.shape[0], dtype=np.float32)
|
||||
itvs = np.round(note_itv_pred_secs * hparams["audio_sample_rate"] / hparams["hop_size"]).astype(int)
|
||||
for i, itv in enumerate(itvs):
|
||||
midi_pred[itv[0] : itv[1]] = note_pred[i]
|
||||
plt.plot(midi_to_hz(midi_pred), color="blue", label="pred midi")
|
||||
plt.plot(note_bd_logits * 100, color="green", label="note bd logits x100")
|
||||
plt.legend()
|
||||
plt.tight_layout()
|
||||
plt.savefig(safe_path(save_dir_path / "plot" / f"[MIDI]{fn}.png"), format="png")
|
||||
plt.close(fig)
|
||||
|
||||
return out
|
||||
|
||||
|
||||
def load_rosvot_models(ckpt, config="", wbd_ckpt="", wbd_config="", device="cuda:0", verbose=False, thr=0.85):
|
||||
"""
|
||||
Load models once to reuse across multiple items.
|
||||
"""
|
||||
dev = torch.device(device)
|
||||
|
||||
# 1. Hparams
|
||||
config_path = Path(ckpt).with_name("config.yaml") if config == "" else config
|
||||
pe_ckpt = Path(ckpt).parent.parent / "rmvpe/model.pt"
|
||||
hparams = set_hparams(
|
||||
config=config_path,
|
||||
print_hparams=verbose,
|
||||
hparams_str=f"note_bd_threshold={thr}",
|
||||
)
|
||||
|
||||
# 2. Main Model
|
||||
model = MidiExtractor(hparams)
|
||||
load_ckpt(model, ckpt, verbose=verbose)
|
||||
model.eval().to(dev)
|
||||
|
||||
# 3. MelNet
|
||||
mel_net = MelNet(hparams)
|
||||
mel_net.to(dev)
|
||||
|
||||
# 4. Pitch Extractor
|
||||
pe = None
|
||||
if hparams.get("use_pitch_embed", False):
|
||||
pe = RMVPE(pe_ckpt, device=dev)
|
||||
|
||||
# 5. Word Boundary Predictor (optional but we load if ckpt provided or needed)
|
||||
wbd_predictor = None
|
||||
if wbd_ckpt:
|
||||
wbd_config_path = Path(wbd_ckpt).with_name("config.yaml") if wbd_config == "" else wbd_config
|
||||
wbd_hparams = set_hparams(
|
||||
config=wbd_config_path,
|
||||
print_hparams=False,
|
||||
hparams_str="",
|
||||
)
|
||||
hparams.update({
|
||||
"wbd_use_mel_bins": wbd_hparams["use_mel_bins"],
|
||||
"min_word_dur": wbd_hparams["min_word_dur"],
|
||||
})
|
||||
wbd_predictor = WordbdExtractor(wbd_hparams)
|
||||
load_ckpt(wbd_predictor, wbd_ckpt, verbose=verbose)
|
||||
wbd_predictor.eval().to(dev)
|
||||
|
||||
models = {
|
||||
"model": model,
|
||||
"mel_net": mel_net,
|
||||
"pe": pe,
|
||||
"wbd_predictor": wbd_predictor
|
||||
}
|
||||
return hparams, models
|
||||
|
||||
|
||||
class NoteTranscriber:
|
||||
"""Note transcription wrapper based on ROSVOT.
|
||||
|
||||
Loads ROSVOT and optional RWBD models once in ``__init__`` and
|
||||
exposes a :py:meth:`process` API that turns an item dict into
|
||||
aligned note metadata for downstream SVS.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
rosvot_model_path: str,
|
||||
rwbd_model_path: str,
|
||||
*,
|
||||
rosvot_config_path: str = "",
|
||||
rwbd_config_path: str = "",
|
||||
device: str = "cuda:0",
|
||||
thr: float = 0.85,
|
||||
verbose: bool = True,
|
||||
):
|
||||
"""Initialize the note transcriber.
|
||||
|
||||
Args:
|
||||
ckpt: Path to the main ROSVOT checkpoint.
|
||||
config: Optional config YAML path for ROSVOT.
|
||||
wbd_ckpt: Optional word-boundary checkpoint path.
|
||||
wbd_config: Optional config YAML path for RWBD.
|
||||
device: Torch device string, e.g. ``"cuda:0"`` / ``"cpu"``.
|
||||
thr: Note boundary threshold.
|
||||
verbose: Whether to print verbose logs.
|
||||
"""
|
||||
self.verbose = verbose
|
||||
self.device = torch.device(device)
|
||||
self.hparams, self.models = load_rosvot_models(
|
||||
ckpt=rosvot_model_path,
|
||||
config=rosvot_config_path,
|
||||
wbd_ckpt=rwbd_model_path,
|
||||
wbd_config=rwbd_config_path,
|
||||
device=device,
|
||||
verbose=verbose,
|
||||
thr=thr,
|
||||
)
|
||||
|
||||
if self.verbose:
|
||||
print(
|
||||
"[note transcription] init success:",
|
||||
f"device={self.device}",
|
||||
f"rosvot_model_path={rosvot_model_path}",
|
||||
f"rwbd_model_path={rwbd_model_path if rwbd_model_path else 'None'}",
|
||||
f"thr={thr}",
|
||||
)
|
||||
|
||||
def process(
|
||||
self,
|
||||
item: Dict[str, Any],
|
||||
*,
|
||||
segment_info: Optional[Dict[str, Any]] = None,
|
||||
save_dir: Optional[str] = None,
|
||||
apply_rwbd: Optional[bool] = None,
|
||||
save_plot: bool = False,
|
||||
no_save_midi: bool = True,
|
||||
no_save_npy: bool = True,
|
||||
verbose: Optional[bool] = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""Run ROSVOT on a single item and post-process outputs.
|
||||
|
||||
Args:
|
||||
item: Input metadata dict with at least ``item_name`` and ``wav_fn``.
|
||||
segment_info: Optional segment metadata for sliced audio.
|
||||
save_dir: Optional directory for debug artifacts (plots, midis).
|
||||
apply_rwbd: Whether to run RWBD-based word boundary refinement.
|
||||
save_plot: Whether to save diagnostic plots.
|
||||
no_save_midi: If True, skip saving midi.
|
||||
no_save_npy: If True, skip saving numpy intermediates.
|
||||
verbose: Override instance-level verbose flag for this call.
|
||||
|
||||
Returns:
|
||||
Dict with aligned note information for downstream SVS.
|
||||
"""
|
||||
v = self.verbose if verbose is None else verbose
|
||||
if v:
|
||||
item_name = item.get("item_name", "")
|
||||
wav_fn = item.get("wav_fn", "")
|
||||
print(f"[note transcription] process: start: item_name={item_name} wav_fn={wav_fn}")
|
||||
t0 = time.time()
|
||||
|
||||
rosvot_out = infer_sample(
|
||||
item,
|
||||
self.hparams,
|
||||
self.models,
|
||||
device=self.device,
|
||||
save_dir=save_dir,
|
||||
apply_rwbd=apply_rwbd,
|
||||
save_plot=save_plot,
|
||||
no_save_midi=no_save_midi,
|
||||
no_save_npy=no_save_npy,
|
||||
verbose=v,
|
||||
)
|
||||
|
||||
out = self.post_process(
|
||||
metadata=item,
|
||||
segment_info=segment_info,
|
||||
rosvot_out=rosvot_out,
|
||||
)
|
||||
|
||||
if v:
|
||||
dt = time.time() - t0
|
||||
print(
|
||||
"[note transcription] process: done:",
|
||||
f"item_name={out.get('item_name','')}",
|
||||
f"n_notes={len(out.get('note_pitch', []) or [])}",
|
||||
f"time={dt:.3f}s",
|
||||
)
|
||||
|
||||
return out
|
||||
|
||||
@staticmethod
|
||||
def _normalize_note2words(note2words: list[int]) -> list[int]:
|
||||
if not note2words:
|
||||
return []
|
||||
normalized = [note2words[0]]
|
||||
for idx in range(1, len(note2words)):
|
||||
if note2words[idx] < normalized[-1]:
|
||||
normalized.append(normalized[-1])
|
||||
else:
|
||||
normalized.append(note2words[idx])
|
||||
return normalized
|
||||
|
||||
@staticmethod
|
||||
def _build_ep_types(note2words: list[int], align_words: list[str]) -> list[int]:
|
||||
ep_types: list[int] = []
|
||||
prev = -1
|
||||
for i, w in zip(note2words, align_words):
|
||||
if w == "<SP>":
|
||||
ep_types.append(1)
|
||||
else:
|
||||
ep_types.append(2 if i != prev else 3)
|
||||
prev = i
|
||||
return ep_types
|
||||
|
||||
def post_process(
|
||||
self,
|
||||
*,
|
||||
metadata: Dict[str, Any],
|
||||
segment_info: Dict[str, Any],
|
||||
rosvot_out: Dict[str, Any],
|
||||
) -> Dict[str, Any]:
|
||||
"""Build aligned note metadata using ROSVOT outputs."""
|
||||
note2words_raw = rosvot_out.get("note2words") or []
|
||||
note2words = self._normalize_note2words(note2words_raw)
|
||||
align_words = [
|
||||
metadata["words"][idx - 1]
|
||||
for idx in note2words_raw
|
||||
if 0 < idx <= len(metadata["words"])
|
||||
]
|
||||
ep_types = self._build_ep_types(note2words, align_words) if align_words else []
|
||||
|
||||
return {
|
||||
"item_name": rosvot_out.get("item_name", "") if not segment_info else segment_info["item_name"],
|
||||
"wav_fn": metadata.get("wav_fn", "") if not segment_info else segment_info["wav_fn"],
|
||||
"origin_wav_fn": metadata.get("origin_wav_fn", "") if not segment_info else segment_info["origin_wav_fn"],
|
||||
"start_time_ms": "" if not segment_info else segment_info["start_time_ms"],
|
||||
"end_time_ms": "" if not segment_info else segment_info["end_time_ms"],
|
||||
"language": metadata.get("language", ""),
|
||||
"note_text": align_words,
|
||||
"note_dur": rosvot_out.get("note_durs", []),
|
||||
"note_type": ep_types,
|
||||
"note_pitch": rosvot_out.get("pitches", []),
|
||||
}
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
items = json.load(open("example/test/rosvot_input.json", "r"))
|
||||
item = items[0]
|
||||
|
||||
m = NoteTranscriber(
|
||||
rosvot_model_path="pretrained_models/rosvot/rosvot/model.pt",
|
||||
rwbd_model_path="pretrained_models/rosvot/rwbd/model.pt",
|
||||
device="cuda"
|
||||
)
|
||||
out = m.process(item)
|
||||
|
||||
print(out)
|
||||
@@ -0,0 +1 @@
|
||||
"""ROSVOT model submodules."""
|
||||
@@ -0,0 +1 @@
|
||||
"""Common ROSVOT layers and utilities."""
|
||||
@@ -0,0 +1 @@
|
||||
"""Conformer layers for ROSVOT."""
|
||||
@@ -0,0 +1,96 @@
|
||||
from torch import nn
|
||||
from .espnet_positional_embedding import RelPositionalEncoding, ScaledPositionalEncoding, PositionalEncoding
|
||||
from .espnet_transformer_attn import RelPositionMultiHeadedAttention, MultiHeadedAttention
|
||||
from .layers import Swish, ConvolutionModule, EncoderLayer, MultiLayeredConv1d
|
||||
from ..layers import Embedding
|
||||
|
||||
|
||||
class ConformerLayers(nn.Module):
|
||||
def __init__(self, hidden_size, num_layers, kernel_size=9, dropout=0.0, num_heads=4,
|
||||
use_last_norm=True, save_hidden=False):
|
||||
super().__init__()
|
||||
self.use_last_norm = use_last_norm
|
||||
self.layers = nn.ModuleList()
|
||||
positionwise_layer = MultiLayeredConv1d
|
||||
positionwise_layer_args = (hidden_size, hidden_size * 4, 1, dropout)
|
||||
self.pos_embed = RelPositionalEncoding(hidden_size, dropout)
|
||||
self.encoder_layers = nn.ModuleList([EncoderLayer(
|
||||
hidden_size,
|
||||
RelPositionMultiHeadedAttention(num_heads, hidden_size, 0.0),
|
||||
positionwise_layer(*positionwise_layer_args),
|
||||
positionwise_layer(*positionwise_layer_args),
|
||||
ConvolutionModule(hidden_size, kernel_size, Swish()),
|
||||
dropout,
|
||||
) for _ in range(num_layers)])
|
||||
if self.use_last_norm:
|
||||
self.layer_norm = nn.LayerNorm(hidden_size)
|
||||
else:
|
||||
self.layer_norm = nn.Linear(hidden_size, hidden_size)
|
||||
self.save_hidden = save_hidden
|
||||
if save_hidden:
|
||||
self.hiddens = []
|
||||
|
||||
def forward(self, x, padding_mask=None):
|
||||
"""
|
||||
|
||||
:param x: [B, T, H]
|
||||
:param padding_mask: [B, T]
|
||||
:return: [B, T, H]
|
||||
"""
|
||||
self.hiddens = []
|
||||
nonpadding_mask = x.abs().sum(-1) > 0
|
||||
x = self.pos_embed(x)
|
||||
for l in self.encoder_layers:
|
||||
x, mask = l(x, nonpadding_mask[:, None, :])
|
||||
if self.save_hidden:
|
||||
self.hiddens.append(x[0])
|
||||
x = x[0]
|
||||
x = self.layer_norm(x) * nonpadding_mask.float()[:, :, None]
|
||||
return x
|
||||
|
||||
class FastConformerLayers(ConformerLayers):
|
||||
def __init__(self, hidden_size, num_layers, kernel_size=9, dropout=0.0, num_heads=4,
|
||||
use_last_norm=True, save_hidden=False):
|
||||
super(ConformerLayers, self).__init__()
|
||||
self.use_last_norm = use_last_norm
|
||||
self.layers = nn.ModuleList()
|
||||
positionwise_layer = MultiLayeredConv1d
|
||||
positionwise_layer_args = (hidden_size, hidden_size * 4, 1, dropout)
|
||||
self.pos_embed = PositionalEncoding(hidden_size, dropout)
|
||||
self.encoder_layers = nn.ModuleList([EncoderLayer(
|
||||
hidden_size,
|
||||
MultiHeadedAttention(num_heads, hidden_size, 0.0, flash=True),
|
||||
positionwise_layer(*positionwise_layer_args),
|
||||
positionwise_layer(*positionwise_layer_args),
|
||||
ConvolutionModule(hidden_size, kernel_size, Swish()),
|
||||
dropout,
|
||||
) for _ in range(num_layers)])
|
||||
if self.use_last_norm:
|
||||
self.layer_norm = nn.LayerNorm(hidden_size)
|
||||
else:
|
||||
self.layer_norm = nn.Linear(hidden_size, hidden_size)
|
||||
self.save_hidden = save_hidden
|
||||
if save_hidden:
|
||||
self.hiddens = []
|
||||
|
||||
class ConformerEncoder(ConformerLayers):
|
||||
def __init__(self, hidden_size, dict_size, num_layers=None):
|
||||
conformer_enc_kernel_size = 9
|
||||
super().__init__(hidden_size, num_layers, conformer_enc_kernel_size)
|
||||
self.embed = Embedding(dict_size, hidden_size, padding_idx=0)
|
||||
|
||||
def forward(self, x):
|
||||
"""
|
||||
|
||||
:param src_tokens: [B, T]
|
||||
:return: [B x T x C]
|
||||
"""
|
||||
x = self.embed(x) # [B, T, H]
|
||||
x = super(ConformerEncoder, self).forward(x)
|
||||
return x
|
||||
|
||||
|
||||
class ConformerDecoder(ConformerLayers):
|
||||
def __init__(self, hidden_size, num_layers):
|
||||
conformer_dec_kernel_size = 9
|
||||
super().__init__(hidden_size, num_layers, conformer_dec_kernel_size)
|
||||
+113
@@ -0,0 +1,113 @@
|
||||
import math
|
||||
import torch
|
||||
|
||||
|
||||
class PositionalEncoding(torch.nn.Module):
|
||||
"""Positional encoding.
|
||||
Args:
|
||||
d_model (int): Embedding dimension.
|
||||
dropout_rate (float): Dropout rate.
|
||||
max_len (int): Maximum input length.
|
||||
reverse (bool): Whether to reverse the input position.
|
||||
"""
|
||||
|
||||
def __init__(self, d_model, dropout_rate, max_len=5000, reverse=False):
|
||||
"""Construct an PositionalEncoding object."""
|
||||
super(PositionalEncoding, self).__init__()
|
||||
self.d_model = d_model
|
||||
self.reverse = reverse
|
||||
self.xscale = math.sqrt(self.d_model)
|
||||
self.dropout = torch.nn.Dropout(p=dropout_rate)
|
||||
self.pe = None
|
||||
self.extend_pe(torch.tensor(0.0).expand(1, max_len))
|
||||
|
||||
def extend_pe(self, x):
|
||||
"""Reset the positional encodings."""
|
||||
if self.pe is not None:
|
||||
if self.pe.size(1) >= x.size(1):
|
||||
if self.pe.dtype != x.dtype or self.pe.device != x.device:
|
||||
self.pe = self.pe.to(dtype=x.dtype, device=x.device)
|
||||
return
|
||||
pe = torch.zeros(x.size(1), self.d_model)
|
||||
if self.reverse:
|
||||
position = torch.arange(
|
||||
x.size(1) - 1, -1, -1.0, dtype=torch.float32
|
||||
).unsqueeze(1)
|
||||
else:
|
||||
position = torch.arange(0, x.size(1), dtype=torch.float32).unsqueeze(1)
|
||||
div_term = torch.exp(
|
||||
torch.arange(0, self.d_model, 2, dtype=torch.float32)
|
||||
* -(math.log(10000.0) / self.d_model)
|
||||
)
|
||||
pe[:, 0::2] = torch.sin(position * div_term)
|
||||
pe[:, 1::2] = torch.cos(position * div_term)
|
||||
pe = pe.unsqueeze(0)
|
||||
self.pe = pe.to(device=x.device, dtype=x.dtype)
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
"""Add positional encoding.
|
||||
Args:
|
||||
x (torch.Tensor): Input tensor (batch, time, `*`).
|
||||
Returns:
|
||||
torch.Tensor: Encoded tensor (batch, time, `*`).
|
||||
"""
|
||||
self.extend_pe(x)
|
||||
x = x * self.xscale + self.pe[:, : x.size(1)]
|
||||
return self.dropout(x)
|
||||
|
||||
|
||||
class ScaledPositionalEncoding(PositionalEncoding):
|
||||
"""Scaled positional encoding module.
|
||||
See Sec. 3.2 https://arxiv.org/abs/1809.08895
|
||||
Args:
|
||||
d_model (int): Embedding dimension.
|
||||
dropout_rate (float): Dropout rate.
|
||||
max_len (int): Maximum input length.
|
||||
"""
|
||||
|
||||
def __init__(self, d_model, dropout_rate, max_len=5000):
|
||||
"""Initialize class."""
|
||||
super().__init__(d_model=d_model, dropout_rate=dropout_rate, max_len=max_len)
|
||||
self.alpha = torch.nn.Parameter(torch.tensor(1.0))
|
||||
|
||||
def reset_parameters(self):
|
||||
"""Reset parameters."""
|
||||
self.alpha.data = torch.tensor(1.0)
|
||||
|
||||
def forward(self, x):
|
||||
"""Add positional encoding.
|
||||
Args:
|
||||
x (torch.Tensor): Input tensor (batch, time, `*`).
|
||||
Returns:
|
||||
torch.Tensor: Encoded tensor (batch, time, `*`).
|
||||
"""
|
||||
self.extend_pe(x)
|
||||
x = x + self.alpha * self.pe[:, : x.size(1)]
|
||||
return self.dropout(x)
|
||||
|
||||
|
||||
class RelPositionalEncoding(PositionalEncoding):
|
||||
"""Relative positional encoding module.
|
||||
See : Appendix B in https://arxiv.org/abs/1901.02860
|
||||
Args:
|
||||
d_model (int): Embedding dimension.
|
||||
dropout_rate (float): Dropout rate.
|
||||
max_len (int): Maximum input length.
|
||||
"""
|
||||
|
||||
def __init__(self, d_model, dropout_rate, max_len=5000):
|
||||
"""Initialize class."""
|
||||
super().__init__(d_model, dropout_rate, max_len, reverse=True)
|
||||
|
||||
def forward(self, x):
|
||||
"""Compute positional encoding.
|
||||
Args:
|
||||
x (torch.Tensor): Input tensor (batch, time, `*`).
|
||||
Returns:
|
||||
torch.Tensor: Encoded tensor (batch, time, `*`).
|
||||
torch.Tensor: Positional embedding tensor (1, time, `*`).
|
||||
"""
|
||||
self.extend_pe(x)
|
||||
x = x * self.xscale
|
||||
pos_emb = self.pe[:, : x.size(1)]
|
||||
return self.dropout(x), self.dropout(pos_emb)
|
||||
+198
@@ -0,0 +1,198 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
# Copyright 2019 Shigeki Karita
|
||||
# Apache 2.0 (http://www.apache.org/licenses/LICENSE-2.0)
|
||||
|
||||
"""Multi-Head Attention layer definition."""
|
||||
|
||||
from packaging import version
|
||||
import math
|
||||
|
||||
import numpy
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
|
||||
class MultiHeadedAttention(nn.Module):
|
||||
"""Multi-Head Attention layer.
|
||||
Args:
|
||||
n_head (int): The number of heads.
|
||||
n_feat (int): The number of features.
|
||||
dropout_rate (float): Dropout rate.
|
||||
"""
|
||||
|
||||
def __init__(self, n_head, n_feat, dropout_rate, flash=False):
|
||||
"""Construct an MultiHeadedAttention object."""
|
||||
super(MultiHeadedAttention, self).__init__()
|
||||
assert n_feat % n_head == 0
|
||||
# We assume d_v always equals d_k
|
||||
self.d_k = n_feat // n_head
|
||||
self.h = n_head
|
||||
self.linear_q = nn.Linear(n_feat, n_feat)
|
||||
self.linear_k = nn.Linear(n_feat, n_feat)
|
||||
self.linear_v = nn.Linear(n_feat, n_feat)
|
||||
self.linear_out = nn.Linear(n_feat, n_feat)
|
||||
self.attn = None
|
||||
self.dropout = nn.Dropout(p=dropout_rate)
|
||||
self.dropout_rate = dropout_rate
|
||||
self.flash = flash
|
||||
|
||||
def forward_qkv(self, query, key, value):
|
||||
"""Transform query, key and value.
|
||||
Args:
|
||||
query (torch.Tensor): Query tensor (#batch, time1, size).
|
||||
key (torch.Tensor): Key tensor (#batch, time2, size).
|
||||
value (torch.Tensor): Value tensor (#batch, time2, size).
|
||||
Returns:
|
||||
torch.Tensor: Transformed query tensor (#batch, n_head, time1, d_k).
|
||||
torch.Tensor: Transformed key tensor (#batch, n_head, time2, d_k).
|
||||
torch.Tensor: Transformed value tensor (#batch, n_head, time2, d_k).
|
||||
"""
|
||||
n_batch = query.size(0)
|
||||
q = self.linear_q(query).view(n_batch, -1, self.h, self.d_k)
|
||||
k = self.linear_k(key).view(n_batch, -1, self.h, self.d_k)
|
||||
v = self.linear_v(value).view(n_batch, -1, self.h, self.d_k)
|
||||
q = q.transpose(1, 2) # (batch, head, time1, d_k)
|
||||
k = k.transpose(1, 2) # (batch, head, time2, d_k)
|
||||
v = v.transpose(1, 2) # (batch, head, time2, d_k)
|
||||
|
||||
return q, k, v
|
||||
|
||||
def forward_attention(self, value, scores, mask):
|
||||
"""Compute attention context vector.
|
||||
Args:
|
||||
value (torch.Tensor): Transformed value (#batch, n_head, time2, d_k).
|
||||
scores (torch.Tensor): Attention score (#batch, n_head, time1, time2).
|
||||
mask (torch.Tensor): Mask (#batch, 1, time2) or (#batch, time1, time2).
|
||||
Returns:
|
||||
torch.Tensor: Transformed value (#batch, time1, d_model)
|
||||
weighted by the attention score (#batch, time1, time2).
|
||||
"""
|
||||
n_batch = value.size(0)
|
||||
if mask is not None:
|
||||
mask = mask.unsqueeze(1).eq(0) # (batch, 1, *, time2)
|
||||
min_value = float(
|
||||
numpy.finfo(torch.tensor(0, dtype=scores.dtype).numpy().dtype).min
|
||||
)
|
||||
scores = scores.masked_fill(mask, min_value)
|
||||
self.attn = torch.softmax(scores, dim=-1).masked_fill(
|
||||
mask, 0.0
|
||||
) # (batch, head, time1, time2)
|
||||
else:
|
||||
self.attn = torch.softmax(scores, dim=-1) # (batch, head, time1, time2)
|
||||
|
||||
p_attn = self.dropout(self.attn)
|
||||
x = torch.matmul(p_attn, value) # (batch, head, time1, d_k)
|
||||
x = (
|
||||
x.transpose(1, 2).contiguous().view(n_batch, -1, self.h * self.d_k)
|
||||
) # (batch, time1, d_model)
|
||||
|
||||
return self.linear_out(x) # (batch, time1, d_model)
|
||||
|
||||
def forward(self, query, key, value, mask):
|
||||
"""Compute scaled dot product attention.
|
||||
Args:
|
||||
query (torch.Tensor): Query tensor (#batch, time1, size).
|
||||
key (torch.Tensor): Key tensor (#batch, time2, size).
|
||||
value (torch.Tensor): Value tensor (#batch, time2, size).
|
||||
mask (torch.Tensor): Mask tensor (#batch, 1, time2) or
|
||||
(#batch, time1, time2).
|
||||
Returns:
|
||||
torch.Tensor: Output tensor (#batch, time1, d_model).
|
||||
"""
|
||||
q, k, v = self.forward_qkv(query, key, value)
|
||||
if version.parse(torch.__version__) >= version.parse("2.0") and self.flash:
|
||||
n_batch = value.size(0)
|
||||
x = torch.nn.functional.scaled_dot_product_attention(
|
||||
q, k, v, attn_mask=mask.unsqueeze(1) if mask is not None else None, dropout_p=self.dropout_rate)
|
||||
x = (
|
||||
x.transpose(1, 2).contiguous().view(n_batch, -1, self.h * self.d_k)
|
||||
) # (batch, time1, d_model)
|
||||
return self.linear_out(x)
|
||||
else:
|
||||
scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.d_k)
|
||||
return self.forward_attention(v, scores, mask)
|
||||
|
||||
|
||||
class RelPositionMultiHeadedAttention(MultiHeadedAttention):
|
||||
"""Multi-Head Attention layer with relative position encoding.
|
||||
Paper: https://arxiv.org/abs/1901.02860
|
||||
Args:
|
||||
n_head (int): The number of heads.
|
||||
n_feat (int): The number of features.
|
||||
dropout_rate (float): Dropout rate.
|
||||
"""
|
||||
|
||||
def __init__(self, n_head, n_feat, dropout_rate):
|
||||
"""Construct an RelPositionMultiHeadedAttention object."""
|
||||
super().__init__(n_head, n_feat, dropout_rate)
|
||||
# linear transformation for positional ecoding
|
||||
self.linear_pos = nn.Linear(n_feat, n_feat, bias=False)
|
||||
# these two learnable bias are used in matrix c and matrix d
|
||||
# as described in https://arxiv.org/abs/1901.02860 Section 3.3
|
||||
self.pos_bias_u = nn.Parameter(torch.Tensor(self.h, self.d_k))
|
||||
self.pos_bias_v = nn.Parameter(torch.Tensor(self.h, self.d_k))
|
||||
torch.nn.init.xavier_uniform_(self.pos_bias_u)
|
||||
torch.nn.init.xavier_uniform_(self.pos_bias_v)
|
||||
|
||||
def rel_shift(self, x, zero_triu=False):
|
||||
"""Compute relative positinal encoding.
|
||||
Args:
|
||||
x (torch.Tensor): Input tensor (batch, time, size).
|
||||
zero_triu (bool): If true, return the lower triangular part of the matrix.
|
||||
Returns:
|
||||
torch.Tensor: Output tensor.
|
||||
"""
|
||||
zero_pad = torch.zeros((*x.size()[:3], 1), device=x.device, dtype=x.dtype)
|
||||
x_padded = torch.cat([zero_pad, x], dim=-1)
|
||||
|
||||
x_padded = x_padded.view(*x.size()[:2], x.size(3) + 1, x.size(2))
|
||||
x = x_padded[:, :, 1:].view_as(x)
|
||||
|
||||
if zero_triu:
|
||||
ones = torch.ones((x.size(2), x.size(3)))
|
||||
x = x * torch.tril(ones, x.size(3) - x.size(2))[None, None, :, :]
|
||||
|
||||
return x
|
||||
|
||||
def forward(self, query, key, value, pos_emb, mask):
|
||||
"""Compute 'Scaled Dot Product Attention' with rel. positional encoding.
|
||||
Args:
|
||||
query (torch.Tensor): Query tensor (#batch, time1, size).
|
||||
key (torch.Tensor): Key tensor (#batch, time2, size).
|
||||
value (torch.Tensor): Value tensor (#batch, time2, size).
|
||||
pos_emb (torch.Tensor): Positional embedding tensor (#batch, time2, size).
|
||||
mask (torch.Tensor): Mask tensor (#batch, 1, time2) or
|
||||
(#batch, time1, time2).
|
||||
Returns:
|
||||
torch.Tensor: Output tensor (#batch, time1, d_model).
|
||||
"""
|
||||
q, k, v = self.forward_qkv(query, key, value)
|
||||
q = q.transpose(1, 2) # (batch, time1, head, d_k)
|
||||
|
||||
n_batch_pos = pos_emb.size(0)
|
||||
p = self.linear_pos(pos_emb).view(n_batch_pos, -1, self.h, self.d_k)
|
||||
p = p.transpose(1, 2) # (batch, head, time1, d_k)
|
||||
|
||||
# (batch, head, time1, d_k)
|
||||
q_with_bias_u = (q + self.pos_bias_u).transpose(1, 2)
|
||||
# (batch, head, time1, d_k)
|
||||
q_with_bias_v = (q + self.pos_bias_v).transpose(1, 2)
|
||||
|
||||
# compute attention score
|
||||
# first compute matrix a and matrix c
|
||||
# as described in https://arxiv.org/abs/1901.02860 Section 3.3
|
||||
# (batch, head, time1, time2)
|
||||
matrix_ac = torch.matmul(q_with_bias_u, k.transpose(-2, -1))
|
||||
|
||||
# compute matrix b and matrix d
|
||||
# (batch, head, time1, time2)
|
||||
matrix_bd = torch.matmul(q_with_bias_v, p.transpose(-2, -1))
|
||||
matrix_bd = self.rel_shift(matrix_bd)
|
||||
|
||||
scores = (matrix_ac + matrix_bd) / math.sqrt(
|
||||
self.d_k
|
||||
) # (batch, head, time1, time2)
|
||||
|
||||
return self.forward_attention(v, scores, mask)
|
||||
@@ -0,0 +1,260 @@
|
||||
from torch import nn
|
||||
import torch
|
||||
|
||||
from ..layers import LayerNorm
|
||||
|
||||
|
||||
class ConvolutionModule(nn.Module):
|
||||
"""ConvolutionModule in Conformer model.
|
||||
Args:
|
||||
channels (int): The number of channels of conv layers.
|
||||
kernel_size (int): Kernerl size of conv layers.
|
||||
"""
|
||||
|
||||
def __init__(self, channels, kernel_size, activation=nn.ReLU(), bias=True):
|
||||
"""Construct an ConvolutionModule object."""
|
||||
super(ConvolutionModule, self).__init__()
|
||||
# kernerl_size should be a odd number for 'SAME' padding
|
||||
assert (kernel_size - 1) % 2 == 0
|
||||
|
||||
self.pointwise_conv1 = nn.Conv1d(
|
||||
channels,
|
||||
2 * channels,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0,
|
||||
bias=bias,
|
||||
)
|
||||
self.depthwise_conv = nn.Conv1d(
|
||||
channels,
|
||||
channels,
|
||||
kernel_size,
|
||||
stride=1,
|
||||
padding=(kernel_size - 1) // 2,
|
||||
groups=channels,
|
||||
bias=bias,
|
||||
)
|
||||
self.norm = nn.BatchNorm1d(channels)
|
||||
self.pointwise_conv2 = nn.Conv1d(
|
||||
channels,
|
||||
channels,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0,
|
||||
bias=bias,
|
||||
)
|
||||
self.activation = activation
|
||||
|
||||
def forward(self, x):
|
||||
"""Compute convolution module.
|
||||
Args:
|
||||
x (torch.Tensor): Input tensor (#batch, time, channels).
|
||||
Returns:
|
||||
torch.Tensor: Output tensor (#batch, time, channels).
|
||||
"""
|
||||
# exchange the temporal dimension and the feature dimension
|
||||
x = x.transpose(1, 2)
|
||||
|
||||
# GLU mechanism
|
||||
x = self.pointwise_conv1(x) # (batch, 2*channel, dim)
|
||||
x = nn.functional.glu(x, dim=1) # (batch, channel, dim)
|
||||
|
||||
# 1D Depthwise Conv
|
||||
x = self.depthwise_conv(x)
|
||||
x = self.activation(self.norm(x))
|
||||
|
||||
x = self.pointwise_conv2(x)
|
||||
|
||||
return x.transpose(1, 2)
|
||||
|
||||
|
||||
class MultiLayeredConv1d(torch.nn.Module):
|
||||
"""Multi-layered conv1d for Transformer block.
|
||||
This is a module of multi-leyered conv1d designed
|
||||
to replace positionwise feed-forward network
|
||||
in Transforner block, which is introduced in
|
||||
`FastSpeech: Fast, Robust and Controllable Text to Speech`_.
|
||||
.. _`FastSpeech: Fast, Robust and Controllable Text to Speech`:
|
||||
https://arxiv.org/pdf/1905.09263.pdf
|
||||
"""
|
||||
|
||||
def __init__(self, in_chans, hidden_chans, kernel_size, dropout_rate):
|
||||
"""Initialize MultiLayeredConv1d module.
|
||||
Args:
|
||||
in_chans (int): Number of input channels.
|
||||
hidden_chans (int): Number of hidden channels.
|
||||
kernel_size (int): Kernel size of conv1d.
|
||||
dropout_rate (float): Dropout rate.
|
||||
"""
|
||||
super(MultiLayeredConv1d, self).__init__()
|
||||
self.w_1 = torch.nn.Conv1d(
|
||||
in_chans,
|
||||
hidden_chans,
|
||||
kernel_size,
|
||||
stride=1,
|
||||
padding=(kernel_size - 1) // 2,
|
||||
)
|
||||
self.w_2 = torch.nn.Conv1d(
|
||||
hidden_chans,
|
||||
in_chans,
|
||||
kernel_size,
|
||||
stride=1,
|
||||
padding=(kernel_size - 1) // 2,
|
||||
)
|
||||
self.dropout = torch.nn.Dropout(dropout_rate)
|
||||
|
||||
def forward(self, x):
|
||||
"""Calculate forward propagation.
|
||||
Args:
|
||||
x (torch.Tensor): Batch of input tensors (B, T, in_chans).
|
||||
Returns:
|
||||
torch.Tensor: Batch of output tensors (B, T, hidden_chans).
|
||||
"""
|
||||
x = torch.relu(self.w_1(x.transpose(-1, 1))).transpose(-1, 1)
|
||||
return self.w_2(self.dropout(x).transpose(-1, 1)).transpose(-1, 1)
|
||||
|
||||
|
||||
class Swish(torch.nn.Module):
|
||||
"""Construct an Swish object."""
|
||||
|
||||
def forward(self, x):
|
||||
"""Return Swich activation function."""
|
||||
return x * torch.sigmoid(x)
|
||||
|
||||
|
||||
class EncoderLayer(nn.Module):
|
||||
"""Encoder layer module.
|
||||
Args:
|
||||
size (int): Input dimension.
|
||||
self_attn (torch.nn.Module): Self-attention module instance.
|
||||
`MultiHeadedAttention` or `RelPositionMultiHeadedAttention` instance
|
||||
can be used as the argument.
|
||||
feed_forward (torch.nn.Module): Feed-forward module instance.
|
||||
`PositionwiseFeedForward`, `MultiLayeredConv1d`, or `Conv1dLinear` instance
|
||||
can be used as the argument.
|
||||
feed_forward_macaron (torch.nn.Module): Additional feed-forward module instance.
|
||||
`PositionwiseFeedForward`, `MultiLayeredConv1d`, or `Conv1dLinear` instance
|
||||
can be used as the argument.
|
||||
conv_module (torch.nn.Module): Convolution module instance.
|
||||
`ConvlutionModule` instance can be used as the argument.
|
||||
dropout_rate (float): Dropout rate.
|
||||
normalize_before (bool): Whether to use layer_norm before the first block.
|
||||
concat_after (bool): Whether to concat attention layer's input and output.
|
||||
if True, additional linear will be applied.
|
||||
i.e. x -> x + linear(concat(x, att(x)))
|
||||
if False, no additional linear will be applied. i.e. x -> x + att(x)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size,
|
||||
self_attn,
|
||||
feed_forward,
|
||||
feed_forward_macaron,
|
||||
conv_module,
|
||||
dropout_rate,
|
||||
normalize_before=True,
|
||||
concat_after=False,
|
||||
):
|
||||
"""Construct an EncoderLayer object."""
|
||||
super(EncoderLayer, self).__init__()
|
||||
self.self_attn = self_attn
|
||||
self.feed_forward = feed_forward
|
||||
self.feed_forward_macaron = feed_forward_macaron
|
||||
self.conv_module = conv_module
|
||||
self.norm_ff = LayerNorm(size) # for the FNN module
|
||||
self.norm_mha = LayerNorm(size) # for the MHA module
|
||||
if feed_forward_macaron is not None:
|
||||
self.norm_ff_macaron = LayerNorm(size)
|
||||
self.ff_scale = 0.5
|
||||
else:
|
||||
self.ff_scale = 1.0
|
||||
if self.conv_module is not None:
|
||||
self.norm_conv = LayerNorm(size) # for the CNN module
|
||||
self.norm_final = LayerNorm(size) # for the final output of the block
|
||||
self.dropout = nn.Dropout(dropout_rate)
|
||||
self.size = size
|
||||
self.normalize_before = normalize_before
|
||||
self.concat_after = concat_after
|
||||
if self.concat_after:
|
||||
self.concat_linear = nn.Linear(size + size, size)
|
||||
|
||||
def forward(self, x_input, mask, cache=None):
|
||||
"""Compute encoded features.
|
||||
Args:
|
||||
x_input (Union[Tuple, torch.Tensor]): Input tensor w/ or w/o pos emb.
|
||||
- w/ pos emb: Tuple of tensors [(#batch, time, size), (1, time, size)].
|
||||
- w/o pos emb: Tensor (#batch, time, size).
|
||||
mask (torch.Tensor): Mask tensor for the input (#batch, time).
|
||||
cache (torch.Tensor): Cache tensor of the input (#batch, time - 1, size).
|
||||
Returns:
|
||||
torch.Tensor: Output tensor (#batch, time, size).
|
||||
torch.Tensor: Mask tensor (#batch, time).
|
||||
"""
|
||||
if isinstance(x_input, tuple):
|
||||
x, pos_emb = x_input[0], x_input[1]
|
||||
else:
|
||||
x, pos_emb = x_input, None
|
||||
|
||||
# whether to use macaron style
|
||||
if self.feed_forward_macaron is not None:
|
||||
residual = x
|
||||
if self.normalize_before:
|
||||
x = self.norm_ff_macaron(x)
|
||||
x = residual + self.ff_scale * self.dropout(self.feed_forward_macaron(x))
|
||||
if not self.normalize_before:
|
||||
x = self.norm_ff_macaron(x)
|
||||
|
||||
# multi-headed self-attention module
|
||||
residual = x
|
||||
if self.normalize_before:
|
||||
x = self.norm_mha(x)
|
||||
|
||||
if cache is None:
|
||||
x_q = x
|
||||
else:
|
||||
assert cache.shape == (x.shape[0], x.shape[1] - 1, self.size)
|
||||
x_q = x[:, -1:, :]
|
||||
residual = residual[:, -1:, :]
|
||||
mask = None if mask is None else mask[:, -1:, :]
|
||||
|
||||
if pos_emb is not None:
|
||||
x_att = self.self_attn(x_q, x, x, pos_emb, mask)
|
||||
else:
|
||||
x_att = self.self_attn(x_q, x, x, mask)
|
||||
|
||||
if self.concat_after:
|
||||
x_concat = torch.cat((x, x_att), dim=-1)
|
||||
x = residual + self.concat_linear(x_concat)
|
||||
else:
|
||||
x = residual + self.dropout(x_att)
|
||||
if not self.normalize_before:
|
||||
x = self.norm_mha(x)
|
||||
|
||||
# convolution module
|
||||
if self.conv_module is not None:
|
||||
residual = x
|
||||
if self.normalize_before:
|
||||
x = self.norm_conv(x)
|
||||
x = residual + self.dropout(self.conv_module(x))
|
||||
if not self.normalize_before:
|
||||
x = self.norm_conv(x)
|
||||
|
||||
# feed forward module
|
||||
residual = x
|
||||
if self.normalize_before:
|
||||
x = self.norm_ff(x)
|
||||
x = residual + self.ff_scale * self.dropout(self.feed_forward(x))
|
||||
if not self.normalize_before:
|
||||
x = self.norm_ff(x)
|
||||
|
||||
if self.conv_module is not None:
|
||||
x = self.norm_final(x)
|
||||
|
||||
if cache is not None:
|
||||
x = torch.cat([cache, x], dim=1)
|
||||
|
||||
if pos_emb is not None:
|
||||
return (x, pos_emb), mask
|
||||
|
||||
return x, mask
|
||||
@@ -0,0 +1,175 @@
|
||||
import math
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from .layers import LayerNorm, Embedding
|
||||
|
||||
class LambdaLayer(nn.Module):
|
||||
def __init__(self, lambd):
|
||||
super(LambdaLayer, self).__init__()
|
||||
self.lambd = lambd
|
||||
|
||||
def forward(self, x):
|
||||
return self.lambd(x)
|
||||
|
||||
def init_weights_func(m):
|
||||
classname = m.__class__.__name__
|
||||
if classname.find("Conv1d") != -1:
|
||||
torch.nn.init.xavier_uniform_(m.weight)
|
||||
|
||||
def get_norm_builder(norm_type, channels, ln_eps=1e-6):
|
||||
if norm_type == 'bn':
|
||||
norm_builder = lambda: nn.BatchNorm1d(channels)
|
||||
elif norm_type == 'in':
|
||||
norm_builder = lambda: nn.InstanceNorm1d(channels, affine=True)
|
||||
elif norm_type == 'gn':
|
||||
norm_builder = lambda: nn.GroupNorm(8, channels)
|
||||
elif norm_type == 'ln':
|
||||
norm_builder = lambda: LayerNorm(channels, dim=1, eps=ln_eps)
|
||||
else:
|
||||
norm_builder = lambda: nn.Identity()
|
||||
return norm_builder
|
||||
|
||||
def get_act_builder(act_type):
|
||||
if act_type == 'gelu':
|
||||
act_builder = lambda: nn.GELU()
|
||||
elif act_type == 'relu':
|
||||
act_builder = lambda: nn.ReLU(inplace=True)
|
||||
elif act_type == 'leakyrelu':
|
||||
act_builder = lambda: nn.LeakyReLU(negative_slope=0.01, inplace=True)
|
||||
elif act_type == 'swish':
|
||||
act_builder = lambda: nn.SiLU(inplace=True)
|
||||
else:
|
||||
act_builder = lambda: nn.Identity()
|
||||
return act_builder
|
||||
|
||||
class ResidualBlock(nn.Module):
|
||||
"""Implements conv->PReLU->norm n-times"""
|
||||
|
||||
def __init__(self, channels, kernel_size, dilation, n=2, norm_type='bn', dropout=0.0,
|
||||
c_multiple=2, ln_eps=1e-12, act_type='gelu'):
|
||||
super(ResidualBlock, self).__init__()
|
||||
|
||||
norm_builder = get_norm_builder(norm_type, channels, ln_eps)
|
||||
act_builder = get_act_builder(act_type)
|
||||
|
||||
self.blocks = [
|
||||
nn.Sequential(
|
||||
norm_builder(),
|
||||
nn.Conv1d(channels, c_multiple * channels, kernel_size, dilation=dilation,
|
||||
padding=(dilation * (kernel_size - 1)) // 2),
|
||||
LambdaLayer(lambda x: x * kernel_size ** -0.5),
|
||||
act_builder(),
|
||||
nn.Conv1d(c_multiple * channels, channels, 1, dilation=dilation),
|
||||
)
|
||||
for i in range(n)
|
||||
]
|
||||
|
||||
self.blocks = nn.ModuleList(self.blocks)
|
||||
self.dropout = dropout
|
||||
|
||||
def forward(self, x):
|
||||
nonpadding = (x.abs().sum(1) > 0).float()[:, None, :]
|
||||
for b in self.blocks:
|
||||
x_ = b(x)
|
||||
if self.dropout > 0 and self.training:
|
||||
x_ = F.dropout(x_, self.dropout, training=self.training)
|
||||
x = x + x_
|
||||
x = x * nonpadding
|
||||
return x
|
||||
|
||||
|
||||
class ConvBlocks(nn.Module):
|
||||
"""Decodes the expanded phoneme encoding into spectrograms"""
|
||||
|
||||
def __init__(self, hidden_size, out_dims, dilations, kernel_size,
|
||||
norm_type='ln', layers_in_block=2, c_multiple=2,
|
||||
dropout=0.0, ln_eps=1e-5,
|
||||
init_weights=True, is_BTC=True, num_layers=None, post_net_kernel=3, act_type='gelu'):
|
||||
super(ConvBlocks, self).__init__()
|
||||
self.is_BTC = is_BTC
|
||||
if num_layers is not None:
|
||||
dilations = [1] * num_layers
|
||||
self.res_blocks = nn.Sequential(
|
||||
*[ResidualBlock(hidden_size, kernel_size, d,
|
||||
n=layers_in_block, norm_type=norm_type, c_multiple=c_multiple,
|
||||
dropout=dropout, ln_eps=ln_eps, act_type=act_type)
|
||||
for d in dilations],
|
||||
)
|
||||
norm = get_norm_builder(norm_type, hidden_size, ln_eps)()
|
||||
self.last_norm = norm
|
||||
self.post_net1 = nn.Conv1d(hidden_size, out_dims, kernel_size=post_net_kernel,
|
||||
padding=post_net_kernel // 2)
|
||||
if init_weights:
|
||||
self.apply(init_weights_func)
|
||||
|
||||
def forward(self, x, nonpadding=None):
|
||||
"""
|
||||
|
||||
:param x: [B, T, H]
|
||||
:return: [B, T, H]
|
||||
"""
|
||||
if self.is_BTC:
|
||||
x = x.transpose(1, 2)
|
||||
if nonpadding is None:
|
||||
nonpadding = (x.abs().sum(1) > 0).float()[:, None, :]
|
||||
elif self.is_BTC:
|
||||
nonpadding = nonpadding.transpose(1, 2)
|
||||
x = self.res_blocks(x) * nonpadding
|
||||
x = self.last_norm(x) * nonpadding
|
||||
x = self.post_net1(x) * nonpadding
|
||||
if self.is_BTC:
|
||||
x = x.transpose(1, 2)
|
||||
return x
|
||||
|
||||
|
||||
class TextConvEncoder(ConvBlocks):
|
||||
def __init__(self, dict_size, hidden_size, out_dims, dilations, kernel_size,
|
||||
norm_type='ln', layers_in_block=2, c_multiple=2,
|
||||
dropout=0.0, ln_eps=1e-5, init_weights=True, num_layers=None, post_net_kernel=3):
|
||||
super().__init__(hidden_size, out_dims, dilations, kernel_size,
|
||||
norm_type, layers_in_block, c_multiple,
|
||||
dropout, ln_eps, init_weights, num_layers=num_layers,
|
||||
post_net_kernel=post_net_kernel)
|
||||
self.embed_tokens = Embedding(dict_size, hidden_size, 0)
|
||||
self.embed_scale = math.sqrt(hidden_size)
|
||||
|
||||
def forward(self, txt_tokens):
|
||||
"""
|
||||
|
||||
:param txt_tokens: [B, T]
|
||||
:return: {
|
||||
'encoder_out': [B x T x C]
|
||||
}
|
||||
"""
|
||||
x = self.embed_scale * self.embed_tokens(txt_tokens)
|
||||
return super().forward(x)
|
||||
|
||||
|
||||
class ConditionalConvBlocks(ConvBlocks):
|
||||
def __init__(self, hidden_size, c_cond, c_out, dilations, kernel_size,
|
||||
norm_type='ln', layers_in_block=2, c_multiple=2,
|
||||
dropout=0.0, ln_eps=1e-5, init_weights=True, is_BTC=True, num_layers=None):
|
||||
super().__init__(hidden_size, c_out, dilations, kernel_size,
|
||||
norm_type, layers_in_block, c_multiple,
|
||||
dropout, ln_eps, init_weights, is_BTC=False, num_layers=num_layers)
|
||||
self.g_prenet = nn.Conv1d(c_cond, hidden_size, 3, padding=1)
|
||||
self.is_BTC_ = is_BTC
|
||||
if init_weights:
|
||||
self.g_prenet.apply(init_weights_func)
|
||||
|
||||
def forward(self, x, cond, nonpadding=None):
|
||||
if self.is_BTC_:
|
||||
x = x.transpose(1, 2)
|
||||
cond = cond.transpose(1, 2)
|
||||
if nonpadding is not None:
|
||||
nonpadding = nonpadding.transpose(1, 2)
|
||||
if nonpadding is None:
|
||||
nonpadding = x.abs().sum(1)[:, None]
|
||||
x = x + self.g_prenet(cond)
|
||||
x = x * nonpadding
|
||||
x = super(ConditionalConvBlocks, self).forward(x) # input needs to be BTC
|
||||
if self.is_BTC_:
|
||||
x = x.transpose(1, 2)
|
||||
return x
|
||||
@@ -0,0 +1,85 @@
|
||||
import torch
|
||||
from torch import nn
|
||||
from torch.autograd import Function
|
||||
|
||||
class LayerNorm(torch.nn.LayerNorm):
|
||||
"""Layer normalization module.
|
||||
:param int nout: output dim size
|
||||
:param int dim: dimension to be normalized
|
||||
"""
|
||||
|
||||
def __init__(self, nout, dim=-1, eps=1e-5):
|
||||
"""Construct an LayerNorm object."""
|
||||
super(LayerNorm, self).__init__(nout, eps=eps)
|
||||
self.dim = dim
|
||||
|
||||
def forward(self, x):
|
||||
"""Apply layer normalization.
|
||||
:param torch.Tensor x: input tensor
|
||||
:return: layer normalized tensor
|
||||
:rtype torch.Tensor
|
||||
"""
|
||||
if self.dim == -1:
|
||||
return super(LayerNorm, self).forward(x)
|
||||
return super(LayerNorm, self).forward(x.transpose(1, -1)).transpose(1, -1)
|
||||
|
||||
|
||||
class Reshape(nn.Module):
|
||||
def __init__(self, *args):
|
||||
super(Reshape, self).__init__()
|
||||
self.shape = args
|
||||
|
||||
def forward(self, x):
|
||||
return x.view(self.shape)
|
||||
|
||||
|
||||
class Permute(nn.Module):
|
||||
def __init__(self, *args):
|
||||
super(Permute, self).__init__()
|
||||
self.args = args
|
||||
|
||||
def forward(self, x):
|
||||
return x.permute(self.args)
|
||||
|
||||
|
||||
def Linear(in_features, out_features, bias=True, init_type='xavier'):
|
||||
m = nn.Linear(in_features, out_features, bias)
|
||||
if init_type == 'xavier':
|
||||
nn.init.xavier_uniform_(m.weight)
|
||||
elif init_type == 'kaiming':
|
||||
nn.init.kaiming_normal_(m.weight, mode='fan_in')
|
||||
if bias:
|
||||
nn.init.constant_(m.bias, 0.)
|
||||
return m
|
||||
|
||||
|
||||
def Embedding(num_embeddings, embedding_dim, padding_idx=None, init_type='normal'):
|
||||
m = nn.Embedding(num_embeddings, embedding_dim, padding_idx=padding_idx)
|
||||
if init_type == 'normal':
|
||||
nn.init.normal_(m.weight, mean=0, std=embedding_dim ** -0.5)
|
||||
elif init_type == 'kaiming':
|
||||
nn.init.kaiming_normal_(m.weight, mode='fan_in')
|
||||
if padding_idx is not None:
|
||||
nn.init.constant_(m.weight[padding_idx], 0)
|
||||
return m
|
||||
|
||||
|
||||
class GradientReverseFunction(Function):
|
||||
@staticmethod
|
||||
def forward(ctx, input, coeff=1.):
|
||||
ctx.coeff = coeff
|
||||
output = input * 1.0
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
return grad_output.neg() * ctx.coeff, None
|
||||
|
||||
|
||||
class GRL(nn.Module):
|
||||
def __init__(self):
|
||||
super(GRL, self).__init__()
|
||||
|
||||
def forward(self, *input):
|
||||
return GradientReverseFunction.apply(*input)
|
||||
|
||||
@@ -0,0 +1,378 @@
|
||||
import math
|
||||
import torch
|
||||
from torch import nn
|
||||
from torch.nn import functional as F
|
||||
|
||||
from .layers import Embedding
|
||||
|
||||
|
||||
def convert_pad_shape(pad_shape):
|
||||
l = pad_shape[::-1]
|
||||
pad_shape = [item for sublist in l for item in sublist]
|
||||
return pad_shape
|
||||
|
||||
|
||||
def shift_1d(x):
|
||||
x = F.pad(x, convert_pad_shape([[0, 0], [0, 0], [1, 0]]))[:, :, :-1]
|
||||
return x
|
||||
|
||||
|
||||
def sequence_mask(length, max_length=None):
|
||||
if max_length is None:
|
||||
max_length = length.max()
|
||||
x = torch.arange(max_length, dtype=length.dtype, device=length.device)
|
||||
return x.unsqueeze(0) < length.unsqueeze(1)
|
||||
|
||||
|
||||
class Encoder(nn.Module):
|
||||
def __init__(self, hidden_channels, filter_channels, n_heads, n_layers, kernel_size=1, p_dropout=0.,
|
||||
window_size=None, block_length=None, pre_ln=False, **kwargs):
|
||||
super().__init__()
|
||||
self.hidden_channels = hidden_channels
|
||||
self.filter_channels = filter_channels
|
||||
self.n_heads = n_heads
|
||||
self.n_layers = n_layers
|
||||
self.kernel_size = kernel_size
|
||||
self.p_dropout = p_dropout
|
||||
self.window_size = window_size
|
||||
self.block_length = block_length
|
||||
self.pre_ln = pre_ln
|
||||
|
||||
self.drop = nn.Dropout(p_dropout)
|
||||
self.attn_layers = nn.ModuleList()
|
||||
self.norm_layers_1 = nn.ModuleList()
|
||||
self.ffn_layers = nn.ModuleList()
|
||||
self.norm_layers_2 = nn.ModuleList()
|
||||
for i in range(self.n_layers):
|
||||
self.attn_layers.append(
|
||||
MultiHeadAttention(hidden_channels, hidden_channels, n_heads, window_size=window_size,
|
||||
p_dropout=p_dropout, block_length=block_length))
|
||||
self.norm_layers_1.append(LayerNorm(hidden_channels))
|
||||
self.ffn_layers.append(
|
||||
FFN(hidden_channels, hidden_channels, filter_channels, kernel_size, p_dropout=p_dropout))
|
||||
self.norm_layers_2.append(LayerNorm(hidden_channels))
|
||||
if pre_ln:
|
||||
self.last_ln = LayerNorm(hidden_channels)
|
||||
|
||||
def forward(self, x, x_mask):
|
||||
attn_mask = x_mask.unsqueeze(2) * x_mask.unsqueeze(-1)
|
||||
for i in range(self.n_layers):
|
||||
x = x * x_mask
|
||||
x_ = x
|
||||
if self.pre_ln:
|
||||
x = self.norm_layers_1[i](x)
|
||||
y = self.attn_layers[i](x, x, attn_mask)
|
||||
y = self.drop(y)
|
||||
x = x_ + y
|
||||
if not self.pre_ln:
|
||||
x = self.norm_layers_1[i](x)
|
||||
|
||||
x_ = x
|
||||
if self.pre_ln:
|
||||
x = self.norm_layers_2[i](x)
|
||||
y = self.ffn_layers[i](x, x_mask)
|
||||
y = self.drop(y)
|
||||
x = x_ + y
|
||||
if not self.pre_ln:
|
||||
x = self.norm_layers_2[i](x)
|
||||
if self.pre_ln:
|
||||
x = self.last_ln(x)
|
||||
x = x * x_mask
|
||||
return x
|
||||
|
||||
|
||||
class MultiHeadAttention(nn.Module):
|
||||
def __init__(self, channels, out_channels, n_heads, window_size=None, heads_share=True, p_dropout=0.,
|
||||
block_length=None, proximal_bias=False, proximal_init=False):
|
||||
super().__init__()
|
||||
assert channels % n_heads == 0
|
||||
|
||||
self.channels = channels
|
||||
self.out_channels = out_channels
|
||||
self.n_heads = n_heads
|
||||
self.window_size = window_size
|
||||
self.heads_share = heads_share
|
||||
self.block_length = block_length
|
||||
self.proximal_bias = proximal_bias
|
||||
self.p_dropout = p_dropout
|
||||
self.attn = None
|
||||
|
||||
self.k_channels = channels // n_heads
|
||||
self.conv_q = nn.Conv1d(channels, channels, 1)
|
||||
self.conv_k = nn.Conv1d(channels, channels, 1)
|
||||
self.conv_v = nn.Conv1d(channels, channels, 1)
|
||||
if window_size is not None:
|
||||
n_heads_rel = 1 if heads_share else n_heads
|
||||
rel_stddev = self.k_channels ** -0.5
|
||||
self.emb_rel_k = nn.Parameter(torch.randn(n_heads_rel, window_size * 2 + 1, self.k_channels) * rel_stddev)
|
||||
self.emb_rel_v = nn.Parameter(torch.randn(n_heads_rel, window_size * 2 + 1, self.k_channels) * rel_stddev)
|
||||
self.conv_o = nn.Conv1d(channels, out_channels, 1)
|
||||
self.drop = nn.Dropout(p_dropout)
|
||||
|
||||
nn.init.xavier_uniform_(self.conv_q.weight)
|
||||
nn.init.xavier_uniform_(self.conv_k.weight)
|
||||
if proximal_init:
|
||||
self.conv_k.weight.data.copy_(self.conv_q.weight.data)
|
||||
self.conv_k.bias.data.copy_(self.conv_q.bias.data)
|
||||
nn.init.xavier_uniform_(self.conv_v.weight)
|
||||
|
||||
def forward(self, x, c, attn_mask=None):
|
||||
q = self.conv_q(x)
|
||||
k = self.conv_k(c)
|
||||
v = self.conv_v(c)
|
||||
|
||||
x, self.attn = self.attention(q, k, v, mask=attn_mask)
|
||||
|
||||
x = self.conv_o(x)
|
||||
return x
|
||||
|
||||
def attention(self, query, key, value, mask=None):
|
||||
# reshape [b, d, t] -> [b, n_h, t, d_k]
|
||||
b, d, t_s, t_t = (*key.size(), query.size(2))
|
||||
query = query.view(b, self.n_heads, self.k_channels, t_t).transpose(2, 3)
|
||||
key = key.view(b, self.n_heads, self.k_channels, t_s).transpose(2, 3)
|
||||
value = value.view(b, self.n_heads, self.k_channels, t_s).transpose(2, 3)
|
||||
|
||||
scores = torch.matmul(query, key.transpose(-2, -1)) / math.sqrt(self.k_channels)
|
||||
if self.window_size is not None:
|
||||
assert t_s == t_t, "Relative attention is only available for self-attention."
|
||||
key_relative_embeddings = self._get_relative_embeddings(self.emb_rel_k, t_s)
|
||||
rel_logits = self._matmul_with_relative_keys(query, key_relative_embeddings)
|
||||
rel_logits = self._relative_position_to_absolute_position(rel_logits)
|
||||
scores_local = rel_logits / math.sqrt(self.k_channels)
|
||||
scores = scores + scores_local
|
||||
if self.proximal_bias:
|
||||
assert t_s == t_t, "Proximal bias is only available for self-attention."
|
||||
scores = scores + self._attention_bias_proximal(t_s).to(device=scores.device, dtype=scores.dtype)
|
||||
if mask is not None:
|
||||
scores = scores.masked_fill(mask == 0, -1e4)
|
||||
if self.block_length is not None:
|
||||
block_mask = torch.ones_like(scores).triu(-self.block_length).tril(self.block_length)
|
||||
scores = scores * block_mask + -1e4 * (1 - block_mask)
|
||||
p_attn = F.softmax(scores, dim=-1) # [b, n_h, t_t, t_s]
|
||||
p_attn = self.drop(p_attn)
|
||||
output = torch.matmul(p_attn, value)
|
||||
if self.window_size is not None:
|
||||
relative_weights = self._absolute_position_to_relative_position(p_attn)
|
||||
value_relative_embeddings = self._get_relative_embeddings(self.emb_rel_v, t_s)
|
||||
output = output + self._matmul_with_relative_values(relative_weights, value_relative_embeddings)
|
||||
output = output.transpose(2, 3).contiguous().view(b, d, t_t) # [b, n_h, t_t, d_k] -> [b, d, t_t]
|
||||
return output, p_attn
|
||||
|
||||
def _matmul_with_relative_values(self, x, y):
|
||||
"""
|
||||
x: [b, h, l, m]
|
||||
y: [h or 1, m, d]
|
||||
ret: [b, h, l, d]
|
||||
"""
|
||||
ret = torch.matmul(x, y.unsqueeze(0))
|
||||
return ret
|
||||
|
||||
def _matmul_with_relative_keys(self, x, y):
|
||||
"""
|
||||
x: [b, h, l, d]
|
||||
y: [h or 1, m, d]
|
||||
ret: [b, h, l, m]
|
||||
"""
|
||||
ret = torch.matmul(x, y.unsqueeze(0).transpose(-2, -1))
|
||||
return ret
|
||||
|
||||
def _get_relative_embeddings(self, relative_embeddings, length):
|
||||
max_relative_position = 2 * self.window_size + 1
|
||||
# Pad first before slice to avoid using cond ops.
|
||||
pad_length = max(length - (self.window_size + 1), 0)
|
||||
slice_start_position = max((self.window_size + 1) - length, 0)
|
||||
slice_end_position = slice_start_position + 2 * length - 1
|
||||
if pad_length > 0:
|
||||
padded_relative_embeddings = F.pad(
|
||||
relative_embeddings,
|
||||
convert_pad_shape([[0, 0], [pad_length, pad_length], [0, 0]]))
|
||||
else:
|
||||
padded_relative_embeddings = relative_embeddings
|
||||
used_relative_embeddings = padded_relative_embeddings[:, slice_start_position:slice_end_position]
|
||||
return used_relative_embeddings
|
||||
|
||||
def _relative_position_to_absolute_position(self, x):
|
||||
"""
|
||||
x: [b, h, l, 2*l-1]
|
||||
ret: [b, h, l, l]
|
||||
"""
|
||||
batch, heads, length, _ = x.size()
|
||||
# Concat columns of pad to shift from relative to absolute indexing.
|
||||
x = F.pad(x, convert_pad_shape([[0, 0], [0, 0], [0, 0], [0, 1]]))
|
||||
|
||||
# Concat extra elements so to add up to shape (len+1, 2*len-1).
|
||||
x_flat = x.view([batch, heads, length * 2 * length])
|
||||
x_flat = F.pad(x_flat, convert_pad_shape([[0, 0], [0, 0], [0, length - 1]]))
|
||||
|
||||
# Reshape and slice out the padded elements.
|
||||
x_final = x_flat.view([batch, heads, length + 1, 2 * length - 1])[:, :, :length, length - 1:]
|
||||
return x_final
|
||||
|
||||
def _absolute_position_to_relative_position(self, x):
|
||||
"""
|
||||
x: [b, h, l, l]
|
||||
ret: [b, h, l, 2*l-1]
|
||||
"""
|
||||
batch, heads, length, _ = x.size()
|
||||
# padd along column
|
||||
x = F.pad(x, convert_pad_shape([[0, 0], [0, 0], [0, 0], [0, length - 1]]))
|
||||
x_flat = x.view([batch, heads, length ** 2 + length * (length - 1)])
|
||||
# add 0's in the beginning that will skew the elements after reshape
|
||||
x_flat = F.pad(x_flat, convert_pad_shape([[0, 0], [0, 0], [length, 0]]))
|
||||
x_final = x_flat.view([batch, heads, length, 2 * length])[:, :, :, 1:]
|
||||
return x_final
|
||||
|
||||
def _attention_bias_proximal(self, length):
|
||||
"""Bias for self-attention to encourage attention to close positions.
|
||||
Args:
|
||||
length: an integer scalar.
|
||||
Returns:
|
||||
a Tensor with shape [1, 1, length, length]
|
||||
"""
|
||||
r = torch.arange(length, dtype=torch.float32)
|
||||
diff = torch.unsqueeze(r, 0) - torch.unsqueeze(r, 1)
|
||||
return torch.unsqueeze(torch.unsqueeze(-torch.log1p(torch.abs(diff)), 0), 0)
|
||||
|
||||
|
||||
class FFN(nn.Module):
|
||||
def __init__(self, in_channels, out_channels, filter_channels, kernel_size, p_dropout=0., activation=None):
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
self.out_channels = out_channels
|
||||
self.filter_channels = filter_channels
|
||||
self.kernel_size = kernel_size
|
||||
self.p_dropout = p_dropout
|
||||
self.activation = activation
|
||||
|
||||
self.conv_1 = nn.Conv1d(in_channels, filter_channels, kernel_size, padding=kernel_size // 2)
|
||||
self.conv_2 = nn.Conv1d(filter_channels, out_channels, 1)
|
||||
self.drop = nn.Dropout(p_dropout)
|
||||
|
||||
def forward(self, x, x_mask):
|
||||
x = self.conv_1(x * x_mask)
|
||||
if self.activation == "gelu":
|
||||
x = x * torch.sigmoid(1.702 * x)
|
||||
else:
|
||||
x = torch.relu(x)
|
||||
x = self.drop(x)
|
||||
x = self.conv_2(x * x_mask)
|
||||
return x * x_mask
|
||||
|
||||
|
||||
class LayerNorm(nn.Module):
|
||||
def __init__(self, channels, eps=1e-4):
|
||||
super().__init__()
|
||||
self.channels = channels
|
||||
self.eps = eps
|
||||
|
||||
self.gamma = nn.Parameter(torch.ones(channels))
|
||||
self.beta = nn.Parameter(torch.zeros(channels))
|
||||
|
||||
def forward(self, x):
|
||||
n_dims = len(x.shape)
|
||||
mean = torch.mean(x, 1, keepdim=True)
|
||||
variance = torch.mean((x - mean) ** 2, 1, keepdim=True)
|
||||
|
||||
x = (x - mean) * torch.rsqrt(variance + self.eps)
|
||||
|
||||
shape = [1, -1] + [1] * (n_dims - 2)
|
||||
x = x * self.gamma.view(*shape) + self.beta.view(*shape)
|
||||
return x
|
||||
|
||||
|
||||
class ConvReluNorm(nn.Module):
|
||||
def __init__(self, in_channels, hidden_channels, out_channels, kernel_size, n_layers, p_dropout):
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
self.hidden_channels = hidden_channels
|
||||
self.out_channels = out_channels
|
||||
self.kernel_size = kernel_size
|
||||
self.n_layers = n_layers
|
||||
self.p_dropout = p_dropout
|
||||
assert n_layers > 1, "Number of layers should be larger than 0."
|
||||
|
||||
self.conv_layers = nn.ModuleList()
|
||||
self.norm_layers = nn.ModuleList()
|
||||
self.conv_layers.append(nn.Conv1d(in_channels, hidden_channels, kernel_size, padding=kernel_size // 2))
|
||||
self.norm_layers.append(LayerNorm(hidden_channels))
|
||||
self.relu_drop = nn.Sequential(
|
||||
nn.ReLU(),
|
||||
nn.Dropout(p_dropout))
|
||||
for _ in range(n_layers - 1):
|
||||
self.conv_layers.append(nn.Conv1d(hidden_channels, hidden_channels, kernel_size, padding=kernel_size // 2))
|
||||
self.norm_layers.append(LayerNorm(hidden_channels))
|
||||
self.proj = nn.Conv1d(hidden_channels, out_channels, 1)
|
||||
self.proj.weight.data.zero_()
|
||||
self.proj.bias.data.zero_()
|
||||
|
||||
def forward(self, x, x_mask):
|
||||
x_org = x
|
||||
for i in range(self.n_layers):
|
||||
x = self.conv_layers[i](x * x_mask)
|
||||
x = self.norm_layers[i](x)
|
||||
x = self.relu_drop(x)
|
||||
x = x_org + self.proj(x)
|
||||
return x * x_mask
|
||||
|
||||
|
||||
class RelTransformerEncoder(nn.Module):
|
||||
def __init__(self,
|
||||
n_vocab,
|
||||
out_channels,
|
||||
hidden_channels,
|
||||
filter_channels,
|
||||
n_heads,
|
||||
n_layers,
|
||||
kernel_size,
|
||||
p_dropout=0.0,
|
||||
window_size=4,
|
||||
block_length=None,
|
||||
prenet=True,
|
||||
pre_ln=True,
|
||||
):
|
||||
|
||||
super().__init__()
|
||||
|
||||
self.n_vocab = n_vocab
|
||||
self.out_channels = out_channels
|
||||
self.hidden_channels = hidden_channels
|
||||
self.filter_channels = filter_channels
|
||||
self.n_heads = n_heads
|
||||
self.n_layers = n_layers
|
||||
self.kernel_size = kernel_size
|
||||
self.p_dropout = p_dropout
|
||||
self.window_size = window_size
|
||||
self.block_length = block_length
|
||||
self.prenet = prenet
|
||||
if n_vocab > 0:
|
||||
self.emb = Embedding(n_vocab, hidden_channels, padding_idx=0)
|
||||
|
||||
if prenet:
|
||||
self.pre = ConvReluNorm(hidden_channels, hidden_channels, hidden_channels,
|
||||
kernel_size=5, n_layers=3, p_dropout=0)
|
||||
self.encoder = Encoder(
|
||||
hidden_channels,
|
||||
filter_channels,
|
||||
n_heads,
|
||||
n_layers,
|
||||
kernel_size,
|
||||
p_dropout,
|
||||
window_size=window_size,
|
||||
block_length=block_length,
|
||||
pre_ln=pre_ln,
|
||||
)
|
||||
|
||||
def forward(self, x, x_mask=None):
|
||||
if self.n_vocab > 0:
|
||||
x_lengths = (x > 0).long().sum(-1)
|
||||
x = self.emb(x) * math.sqrt(self.hidden_channels) # [b, t, h]
|
||||
else:
|
||||
x_lengths = (x.abs().sum(-1) > 0).long().sum(-1)
|
||||
x = torch.transpose(x, 1, -1) # [b, h, t]
|
||||
x_mask = torch.unsqueeze(sequence_mask(x_lengths, x.size(2)), 1).to(x.dtype)
|
||||
|
||||
if self.prenet:
|
||||
x = self.pre(x, x_mask)
|
||||
x = self.encoder(x, x_mask)
|
||||
return x.transpose(1, 2)
|
||||
@@ -0,0 +1,261 @@
|
||||
import torch
|
||||
from torch import nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
class PreNet(nn.Module):
|
||||
def __init__(self, in_dims, fc1_dims=256, fc2_dims=128, dropout=0.5):
|
||||
super().__init__()
|
||||
self.fc1 = nn.Linear(in_dims, fc1_dims)
|
||||
self.fc2 = nn.Linear(fc1_dims, fc2_dims)
|
||||
self.p = dropout
|
||||
|
||||
def forward(self, x):
|
||||
x = self.fc1(x)
|
||||
x = F.relu(x)
|
||||
x = F.dropout(x, self.p, training=self.training)
|
||||
x = self.fc2(x)
|
||||
x = F.relu(x)
|
||||
x = F.dropout(x, self.p, training=self.training)
|
||||
return x
|
||||
|
||||
|
||||
class HighwayNetwork(nn.Module):
|
||||
def __init__(self, size):
|
||||
super().__init__()
|
||||
self.W1 = nn.Linear(size, size)
|
||||
self.W2 = nn.Linear(size, size)
|
||||
self.W1.bias.data.fill_(0.)
|
||||
|
||||
def forward(self, x):
|
||||
x1 = self.W1(x)
|
||||
x2 = self.W2(x)
|
||||
g = torch.sigmoid(x2)
|
||||
y = g * F.relu(x1) + (1. - g) * x
|
||||
return y
|
||||
|
||||
|
||||
class BatchNormConv(nn.Module):
|
||||
def __init__(self, in_channels, out_channels, kernel, relu=True):
|
||||
super().__init__()
|
||||
self.conv = nn.Conv1d(in_channels, out_channels, kernel, stride=1, padding=kernel // 2, bias=False)
|
||||
self.bnorm = nn.BatchNorm1d(out_channels)
|
||||
self.relu = relu
|
||||
|
||||
def forward(self, x):
|
||||
x = self.conv(x)
|
||||
x = F.relu(x) if self.relu is True else x
|
||||
return self.bnorm(x)
|
||||
|
||||
|
||||
class ConvNorm(torch.nn.Module):
|
||||
def __init__(self, in_channels, out_channels, kernel_size=1, stride=1,
|
||||
padding=None, dilation=1, bias=True, w_init_gain='linear'):
|
||||
super(ConvNorm, self).__init__()
|
||||
if padding is None:
|
||||
assert (kernel_size % 2 == 1)
|
||||
padding = int(dilation * (kernel_size - 1) / 2)
|
||||
|
||||
self.conv = torch.nn.Conv1d(in_channels, out_channels,
|
||||
kernel_size=kernel_size, stride=stride,
|
||||
padding=padding, dilation=dilation,
|
||||
bias=bias)
|
||||
|
||||
torch.nn.init.xavier_uniform_(
|
||||
self.conv.weight, gain=torch.nn.init.calculate_gain(w_init_gain))
|
||||
|
||||
def forward(self, signal):
|
||||
conv_signal = self.conv(signal)
|
||||
return conv_signal
|
||||
|
||||
|
||||
class CBHG(nn.Module):
|
||||
def __init__(self, K, in_channels, channels, proj_channels, num_highways):
|
||||
super().__init__()
|
||||
|
||||
# List of all rnns to call `flatten_parameters()` on
|
||||
self._to_flatten = []
|
||||
|
||||
self.bank_kernels = [i for i in range(1, K + 1)]
|
||||
self.conv1d_bank = nn.ModuleList()
|
||||
for k in self.bank_kernels:
|
||||
conv = BatchNormConv(in_channels, channels, k)
|
||||
self.conv1d_bank.append(conv)
|
||||
|
||||
self.maxpool = nn.MaxPool1d(kernel_size=2, stride=1, padding=1)
|
||||
|
||||
self.conv_project1 = BatchNormConv(len(self.bank_kernels) * channels, proj_channels[0], 3)
|
||||
self.conv_project2 = BatchNormConv(proj_channels[0], proj_channels[1], 3, relu=False)
|
||||
|
||||
# Fix the highway input if necessary
|
||||
if proj_channels[-1] != channels:
|
||||
self.highway_mismatch = True
|
||||
self.pre_highway = nn.Linear(proj_channels[-1], channels, bias=False)
|
||||
else:
|
||||
self.highway_mismatch = False
|
||||
|
||||
self.highways = nn.ModuleList()
|
||||
for i in range(num_highways):
|
||||
hn = HighwayNetwork(channels)
|
||||
self.highways.append(hn)
|
||||
|
||||
self.rnn = nn.GRU(channels, channels, batch_first=True, bidirectional=True)
|
||||
self._to_flatten.append(self.rnn)
|
||||
|
||||
# Avoid fragmentation of RNN parameters and associated warning
|
||||
self._flatten_parameters()
|
||||
|
||||
def forward(self, x):
|
||||
# Although we `_flatten_parameters()` on init, when using DataParallel
|
||||
# the model gets replicated, making it no longer guaranteed that the
|
||||
# weights are contiguous in GPU memory. Hence, we must call it again
|
||||
self._flatten_parameters()
|
||||
|
||||
# Save these for later
|
||||
residual = x
|
||||
seq_len = x.size(-1)
|
||||
conv_bank = []
|
||||
|
||||
# Convolution Bank
|
||||
for conv in self.conv1d_bank:
|
||||
c = conv(x) # Convolution
|
||||
conv_bank.append(c[:, :, :seq_len])
|
||||
|
||||
# Stack along the channel axis
|
||||
conv_bank = torch.cat(conv_bank, dim=1)
|
||||
|
||||
# dump the last padding to fit residual
|
||||
x = self.maxpool(conv_bank)[:, :, :seq_len]
|
||||
|
||||
# Conv1d projections
|
||||
x = self.conv_project1(x)
|
||||
x = self.conv_project2(x)
|
||||
|
||||
# Residual Connect
|
||||
x = x + residual
|
||||
|
||||
# Through the highways
|
||||
x = x.transpose(1, 2)
|
||||
if self.highway_mismatch is True:
|
||||
x = self.pre_highway(x)
|
||||
for h in self.highways:
|
||||
x = h(x)
|
||||
|
||||
# And then the RNN
|
||||
x, _ = self.rnn(x)
|
||||
return x
|
||||
|
||||
def _flatten_parameters(self):
|
||||
"""Calls `flatten_parameters` on all the rnns used by the WaveRNN. Used
|
||||
to improve efficiency and avoid PyTorch yelling at us."""
|
||||
[m.flatten_parameters() for m in self._to_flatten]
|
||||
|
||||
|
||||
class TacotronEncoder(nn.Module):
|
||||
def __init__(self, embed_dims, num_chars, cbhg_channels, K, num_highways, dropout):
|
||||
super().__init__()
|
||||
self.embedding = nn.Embedding(num_chars, embed_dims)
|
||||
self.pre_net = PreNet(embed_dims, embed_dims, embed_dims, dropout=dropout)
|
||||
self.cbhg = CBHG(K=K, in_channels=cbhg_channels, channels=cbhg_channels,
|
||||
proj_channels=[cbhg_channels, cbhg_channels],
|
||||
num_highways=num_highways)
|
||||
self.proj_out = nn.Linear(cbhg_channels * 2, cbhg_channels)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.embedding(x)
|
||||
x = self.pre_net(x)
|
||||
x.transpose_(1, 2)
|
||||
x = self.cbhg(x)
|
||||
x = self.proj_out(x)
|
||||
return x
|
||||
|
||||
|
||||
class RNNEncoder(nn.Module):
|
||||
def __init__(self, num_chars, embedding_dim, n_convolutions=3, kernel_size=5):
|
||||
super(RNNEncoder, self).__init__()
|
||||
self.embedding = nn.Embedding(num_chars, embedding_dim, padding_idx=0)
|
||||
convolutions = []
|
||||
for _ in range(n_convolutions):
|
||||
conv_layer = nn.Sequential(
|
||||
ConvNorm(embedding_dim,
|
||||
embedding_dim,
|
||||
kernel_size=kernel_size, stride=1,
|
||||
padding=int((kernel_size - 1) / 2),
|
||||
dilation=1, w_init_gain='relu'),
|
||||
nn.BatchNorm1d(embedding_dim))
|
||||
convolutions.append(conv_layer)
|
||||
self.convolutions = nn.ModuleList(convolutions)
|
||||
|
||||
self.lstm = nn.LSTM(embedding_dim, int(embedding_dim / 2), 1,
|
||||
batch_first=True, bidirectional=True)
|
||||
|
||||
def forward(self, x):
|
||||
input_lengths = (x > 0).sum(-1)
|
||||
input_lengths = input_lengths.cpu().numpy()
|
||||
|
||||
x = self.embedding(x)
|
||||
x = x.transpose(1, 2) # [B, H, T]
|
||||
for conv in self.convolutions:
|
||||
x = F.dropout(F.relu(conv(x)), 0.5, self.training) + x
|
||||
x = x.transpose(1, 2) # [B, T, H]
|
||||
|
||||
# pytorch tensor are not reversible, hence the conversion
|
||||
x = nn.utils.rnn.pack_padded_sequence(x, input_lengths, batch_first=True, enforce_sorted=False)
|
||||
|
||||
self.lstm.flatten_parameters()
|
||||
outputs, _ = self.lstm(x)
|
||||
outputs, _ = nn.utils.rnn.pad_packed_sequence(outputs, batch_first=True)
|
||||
|
||||
return outputs
|
||||
|
||||
|
||||
class DecoderRNN(torch.nn.Module):
|
||||
def __init__(self, hidden_size, decoder_rnn_dim, dropout):
|
||||
super(DecoderRNN, self).__init__()
|
||||
self.in_conv1d = nn.Sequential(
|
||||
torch.nn.Conv1d(
|
||||
in_channels=hidden_size,
|
||||
out_channels=hidden_size,
|
||||
kernel_size=9, padding=4,
|
||||
),
|
||||
torch.nn.ReLU(),
|
||||
torch.nn.Conv1d(
|
||||
in_channels=hidden_size,
|
||||
out_channels=hidden_size,
|
||||
kernel_size=9, padding=4,
|
||||
),
|
||||
)
|
||||
self.ln = nn.LayerNorm(hidden_size)
|
||||
if decoder_rnn_dim == 0:
|
||||
decoder_rnn_dim = hidden_size * 2
|
||||
self.rnn = torch.nn.LSTM(
|
||||
input_size=hidden_size,
|
||||
hidden_size=decoder_rnn_dim,
|
||||
num_layers=1,
|
||||
batch_first=True,
|
||||
bidirectional=True,
|
||||
dropout=dropout
|
||||
)
|
||||
self.rnn.flatten_parameters()
|
||||
self.conv1d = torch.nn.Conv1d(
|
||||
in_channels=decoder_rnn_dim * 2,
|
||||
out_channels=hidden_size,
|
||||
kernel_size=3,
|
||||
padding=1,
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
input_masks = x.abs().sum(-1).ne(0).data[:, :, None]
|
||||
input_lengths = input_masks.sum([-1, -2])
|
||||
input_lengths = input_lengths.cpu().numpy()
|
||||
|
||||
x = self.in_conv1d(x.transpose(1, 2)).transpose(1, 2)
|
||||
x = self.ln(x)
|
||||
x = nn.utils.rnn.pack_padded_sequence(x, input_lengths, batch_first=True, enforce_sorted=False)
|
||||
self.rnn.flatten_parameters()
|
||||
x, _ = self.rnn(x) # [B, T, C]
|
||||
x, _ = nn.utils.rnn.pad_packed_sequence(x, batch_first=True)
|
||||
x = x * input_masks
|
||||
pre_mel = self.conv1d(x.transpose(1, 2)).transpose(1, 2) # [B, T, C]
|
||||
pre_mel = pre_mel * input_masks
|
||||
return pre_mel
|
||||
@@ -0,0 +1,751 @@
|
||||
import math
|
||||
import torch
|
||||
from torch import nn
|
||||
from torch.nn import Parameter, Linear
|
||||
from .layers import LayerNorm, Embedding
|
||||
from ...utils.nn.seq_utils import (
|
||||
get_incremental_state,
|
||||
set_incremental_state,
|
||||
softmax,
|
||||
make_positions,
|
||||
)
|
||||
import torch.nn.functional as F
|
||||
|
||||
DEFAULT_MAX_SOURCE_POSITIONS = 2000
|
||||
DEFAULT_MAX_TARGET_POSITIONS = 2000
|
||||
|
||||
|
||||
class SinusoidalPositionalEmbedding(nn.Module):
|
||||
"""This module produces sinusoidal positional embeddings of any length.
|
||||
|
||||
Padding symbols are ignored.
|
||||
"""
|
||||
|
||||
def __init__(self, embedding_dim, padding_idx, init_size=1024):
|
||||
super().__init__()
|
||||
self.embedding_dim = embedding_dim
|
||||
self.padding_idx = padding_idx
|
||||
self.weights = SinusoidalPositionalEmbedding.get_embedding(
|
||||
init_size,
|
||||
embedding_dim,
|
||||
padding_idx,
|
||||
)
|
||||
self.register_buffer('_float_tensor', torch.FloatTensor(1))
|
||||
|
||||
@staticmethod
|
||||
def get_embedding(num_embeddings, embedding_dim, padding_idx=None):
|
||||
"""Build sinusoidal embeddings.
|
||||
|
||||
This matches the implementation in tensor2tensor, but differs slightly
|
||||
from the description in Section 3.5 of "Attention Is All You Need".
|
||||
"""
|
||||
half_dim = embedding_dim // 2
|
||||
emb = math.log(10000) / (half_dim - 1)
|
||||
emb = torch.exp(torch.arange(half_dim, dtype=torch.float) * -emb)
|
||||
emb = torch.arange(num_embeddings, dtype=torch.float).unsqueeze(1) * emb.unsqueeze(0)
|
||||
emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=1).view(num_embeddings, -1)
|
||||
if embedding_dim % 2 == 1:
|
||||
# zero pad
|
||||
emb = torch.cat([emb, torch.zeros(num_embeddings, 1)], dim=1)
|
||||
if padding_idx is not None:
|
||||
emb[padding_idx, :] = 0
|
||||
return emb
|
||||
|
||||
def forward(self, input, incremental_state=None, timestep=None, positions=None, **kwargs):
|
||||
"""Input is expected to be of size [bsz x seqlen]."""
|
||||
bsz, seq_len = input.shape[:2]
|
||||
max_pos = self.padding_idx + 1 + seq_len
|
||||
if self.weights is None or max_pos > self.weights.size(0):
|
||||
# recompute/expand embeddings if needed
|
||||
self.weights = SinusoidalPositionalEmbedding.get_embedding(
|
||||
max_pos,
|
||||
self.embedding_dim,
|
||||
self.padding_idx,
|
||||
)
|
||||
self.weights = self.weights.to(self._float_tensor)
|
||||
|
||||
if incremental_state is not None:
|
||||
# positions is the same for every token when decoding a single step
|
||||
pos = timestep.view(-1)[0] + 1 if timestep is not None else seq_len
|
||||
return self.weights[self.padding_idx + pos, :].expand(bsz, 1, -1)
|
||||
|
||||
positions = make_positions(input, self.padding_idx) if positions is None else positions
|
||||
return self.weights.index_select(0, positions.view(-1)).view(bsz, seq_len, -1).detach()
|
||||
|
||||
def max_positions(self):
|
||||
"""Maximum number of supported positions."""
|
||||
return int(1e5) # an arbitrary large number
|
||||
|
||||
|
||||
class TransformerFFNLayer(nn.Module):
|
||||
def __init__(self, hidden_size, filter_size, padding="SAME", kernel_size=1, dropout=0., act='gelu'):
|
||||
super().__init__()
|
||||
self.kernel_size = kernel_size
|
||||
self.dropout = dropout
|
||||
self.act = act
|
||||
if padding == 'SAME':
|
||||
self.ffn_1 = nn.Conv1d(hidden_size, filter_size, kernel_size, padding=kernel_size // 2)
|
||||
elif padding == 'LEFT':
|
||||
self.ffn_1 = nn.Sequential(
|
||||
nn.ConstantPad1d((kernel_size - 1, 0), 0.0),
|
||||
nn.Conv1d(hidden_size, filter_size, kernel_size)
|
||||
)
|
||||
self.ffn_2 = Linear(filter_size, hidden_size)
|
||||
|
||||
def forward(self, x, incremental_state=None):
|
||||
# x: T x B x C
|
||||
if incremental_state is not None:
|
||||
saved_state = self._get_input_buffer(incremental_state)
|
||||
if 'prev_input' in saved_state:
|
||||
prev_input = saved_state['prev_input']
|
||||
x = torch.cat((prev_input, x), dim=0)
|
||||
x = x[-self.kernel_size:]
|
||||
saved_state['prev_input'] = x
|
||||
self._set_input_buffer(incremental_state, saved_state)
|
||||
|
||||
x = self.ffn_1(x.permute(1, 2, 0)).permute(2, 0, 1)
|
||||
x = x * self.kernel_size ** -0.5
|
||||
|
||||
if incremental_state is not None:
|
||||
x = x[-1:]
|
||||
if self.act == 'gelu':
|
||||
x = F.gelu(x)
|
||||
if self.act == 'relu':
|
||||
x = F.relu(x)
|
||||
x = F.dropout(x, self.dropout, training=self.training)
|
||||
x = self.ffn_2(x)
|
||||
return x
|
||||
|
||||
def _get_input_buffer(self, incremental_state):
|
||||
return get_incremental_state(
|
||||
self,
|
||||
incremental_state,
|
||||
'f',
|
||||
) or {}
|
||||
|
||||
def _set_input_buffer(self, incremental_state, buffer):
|
||||
set_incremental_state(
|
||||
self,
|
||||
incremental_state,
|
||||
'f',
|
||||
buffer,
|
||||
)
|
||||
|
||||
def clear_buffer(self, incremental_state):
|
||||
if incremental_state is not None:
|
||||
saved_state = self._get_input_buffer(incremental_state)
|
||||
if 'prev_input' in saved_state:
|
||||
del saved_state['prev_input']
|
||||
self._set_input_buffer(incremental_state, saved_state)
|
||||
|
||||
|
||||
class MultiheadAttention(nn.Module):
|
||||
def __init__(self, embed_dim, num_heads, kdim=None, vdim=None, dropout=0., bias=True,
|
||||
add_bias_kv=False, add_zero_attn=False, self_attention=False,
|
||||
encoder_decoder_attention=False):
|
||||
super().__init__()
|
||||
self.embed_dim = embed_dim
|
||||
self.kdim = kdim if kdim is not None else embed_dim
|
||||
self.vdim = vdim if vdim is not None else embed_dim
|
||||
self.qkv_same_dim = self.kdim == embed_dim and self.vdim == embed_dim
|
||||
|
||||
self.num_heads = num_heads
|
||||
self.dropout = dropout
|
||||
self.head_dim = embed_dim // num_heads
|
||||
assert self.head_dim * num_heads == self.embed_dim, "embed_dim must be divisible by num_heads"
|
||||
self.scaling = self.head_dim ** -0.5
|
||||
|
||||
self.self_attention = self_attention
|
||||
self.encoder_decoder_attention = encoder_decoder_attention
|
||||
|
||||
assert not self.self_attention or self.qkv_same_dim, 'Self-attention requires query, key and ' \
|
||||
'value to be of the same size'
|
||||
|
||||
if self.qkv_same_dim:
|
||||
self.in_proj_weight = Parameter(torch.Tensor(3 * embed_dim, embed_dim))
|
||||
else:
|
||||
self.k_proj_weight = Parameter(torch.Tensor(embed_dim, self.kdim))
|
||||
self.v_proj_weight = Parameter(torch.Tensor(embed_dim, self.vdim))
|
||||
self.q_proj_weight = Parameter(torch.Tensor(embed_dim, embed_dim))
|
||||
|
||||
if bias:
|
||||
self.in_proj_bias = Parameter(torch.Tensor(3 * embed_dim))
|
||||
else:
|
||||
self.register_parameter('in_proj_bias', None)
|
||||
|
||||
self.out_proj = nn.Linear(embed_dim, embed_dim, bias=bias)
|
||||
|
||||
if add_bias_kv:
|
||||
self.bias_k = Parameter(torch.Tensor(1, 1, embed_dim))
|
||||
self.bias_v = Parameter(torch.Tensor(1, 1, embed_dim))
|
||||
else:
|
||||
self.bias_k = self.bias_v = None
|
||||
|
||||
self.add_zero_attn = add_zero_attn
|
||||
|
||||
self.reset_parameters()
|
||||
|
||||
self.enable_torch_version = False
|
||||
if hasattr(F, "multi_head_attention_forward"):
|
||||
self.enable_torch_version = True
|
||||
else:
|
||||
self.enable_torch_version = False
|
||||
self.last_attn_probs = None
|
||||
|
||||
def reset_parameters(self):
|
||||
if self.qkv_same_dim:
|
||||
nn.init.xavier_uniform_(self.in_proj_weight)
|
||||
else:
|
||||
nn.init.xavier_uniform_(self.k_proj_weight)
|
||||
nn.init.xavier_uniform_(self.v_proj_weight)
|
||||
nn.init.xavier_uniform_(self.q_proj_weight)
|
||||
|
||||
nn.init.xavier_uniform_(self.out_proj.weight)
|
||||
if self.in_proj_bias is not None:
|
||||
nn.init.constant_(self.in_proj_bias, 0.)
|
||||
nn.init.constant_(self.out_proj.bias, 0.)
|
||||
if self.bias_k is not None:
|
||||
nn.init.xavier_normal_(self.bias_k)
|
||||
if self.bias_v is not None:
|
||||
nn.init.xavier_normal_(self.bias_v)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
query, key, value,
|
||||
key_padding_mask=None,
|
||||
incremental_state=None,
|
||||
need_weights=True,
|
||||
static_kv=False,
|
||||
attn_mask=None,
|
||||
before_softmax=False,
|
||||
need_head_weights=False,
|
||||
enc_dec_attn_constraint_mask=None,
|
||||
reset_attn_weight=None
|
||||
):
|
||||
"""Input shape: Time x Batch x Channel
|
||||
|
||||
Args:
|
||||
key_padding_mask (ByteTensor, optional): mask to exclude
|
||||
keys that are pads, of shape `(batch, src_len)`, where
|
||||
padding elements are indicated by 1s.
|
||||
need_weights (bool, optional): return the attention weights,
|
||||
averaged over heads (default: False).
|
||||
attn_mask (ByteTensor, optional): typically used to
|
||||
implement causal attention, where the mask prevents the
|
||||
attention from looking forward in time (default: None).
|
||||
before_softmax (bool, optional): return the raw attention
|
||||
weights and values before the attention softmax.
|
||||
need_head_weights (bool, optional): return the attention
|
||||
weights for each head. Implies *need_weights*. Default:
|
||||
return the average attention weights over all heads.
|
||||
"""
|
||||
if need_head_weights:
|
||||
need_weights = True
|
||||
|
||||
tgt_len, bsz, embed_dim = query.size()
|
||||
assert embed_dim == self.embed_dim
|
||||
assert list(query.size()) == [tgt_len, bsz, embed_dim]
|
||||
if self.enable_torch_version and incremental_state is None and not static_kv and reset_attn_weight is None:
|
||||
if self.qkv_same_dim:
|
||||
return F.multi_head_attention_forward(query, key, value,
|
||||
self.embed_dim, self.num_heads,
|
||||
self.in_proj_weight,
|
||||
self.in_proj_bias, self.bias_k, self.bias_v,
|
||||
self.add_zero_attn, self.dropout,
|
||||
self.out_proj.weight, self.out_proj.bias,
|
||||
self.training, key_padding_mask, need_weights,
|
||||
attn_mask)
|
||||
else:
|
||||
return F.multi_head_attention_forward(query, key, value,
|
||||
self.embed_dim, self.num_heads,
|
||||
torch.empty([0]),
|
||||
self.in_proj_bias, self.bias_k, self.bias_v,
|
||||
self.add_zero_attn, self.dropout,
|
||||
self.out_proj.weight, self.out_proj.bias,
|
||||
self.training, key_padding_mask, need_weights,
|
||||
attn_mask, use_separate_proj_weight=True,
|
||||
q_proj_weight=self.q_proj_weight,
|
||||
k_proj_weight=self.k_proj_weight,
|
||||
v_proj_weight=self.v_proj_weight)
|
||||
|
||||
if incremental_state is not None:
|
||||
saved_state = self._get_input_buffer(incremental_state)
|
||||
if 'prev_key' in saved_state:
|
||||
# previous time steps are cached - no need to recompute
|
||||
# key and value if they are static
|
||||
if static_kv:
|
||||
assert self.encoder_decoder_attention and not self.self_attention
|
||||
key = value = None
|
||||
else:
|
||||
saved_state = None
|
||||
|
||||
if self.self_attention:
|
||||
# self-attention
|
||||
q, k, v = self.in_proj_qkv(query)
|
||||
elif self.encoder_decoder_attention:
|
||||
# encoder-decoder attention
|
||||
q = self.in_proj_q(query)
|
||||
if key is None:
|
||||
assert value is None
|
||||
k = v = None
|
||||
else:
|
||||
k = self.in_proj_k(key)
|
||||
v = self.in_proj_v(key)
|
||||
|
||||
else:
|
||||
q = self.in_proj_q(query)
|
||||
k = self.in_proj_k(key)
|
||||
v = self.in_proj_v(value)
|
||||
q *= self.scaling
|
||||
|
||||
if self.bias_k is not None:
|
||||
assert self.bias_v is not None
|
||||
k = torch.cat([k, self.bias_k.repeat(1, bsz, 1)])
|
||||
v = torch.cat([v, self.bias_v.repeat(1, bsz, 1)])
|
||||
if attn_mask is not None:
|
||||
attn_mask = torch.cat([attn_mask, attn_mask.new_zeros(attn_mask.size(0), 1)], dim=1)
|
||||
if key_padding_mask is not None:
|
||||
key_padding_mask = torch.cat(
|
||||
[key_padding_mask, key_padding_mask.new_zeros(key_padding_mask.size(0), 1)], dim=1)
|
||||
|
||||
q = q.contiguous().view(tgt_len, bsz * self.num_heads, self.head_dim).transpose(0, 1)
|
||||
if k is not None:
|
||||
k = k.contiguous().view(-1, bsz * self.num_heads, self.head_dim).transpose(0, 1)
|
||||
if v is not None:
|
||||
v = v.contiguous().view(-1, bsz * self.num_heads, self.head_dim).transpose(0, 1)
|
||||
|
||||
if saved_state is not None:
|
||||
# saved states are stored with shape (bsz, num_heads, seq_len, head_dim)
|
||||
if 'prev_key' in saved_state:
|
||||
prev_key = saved_state['prev_key'].view(bsz * self.num_heads, -1, self.head_dim)
|
||||
if static_kv:
|
||||
k = prev_key
|
||||
else:
|
||||
k = torch.cat((prev_key, k), dim=1)
|
||||
if 'prev_value' in saved_state:
|
||||
prev_value = saved_state['prev_value'].view(bsz * self.num_heads, -1, self.head_dim)
|
||||
if static_kv:
|
||||
v = prev_value
|
||||
else:
|
||||
v = torch.cat((prev_value, v), dim=1)
|
||||
if 'prev_key_padding_mask' in saved_state and saved_state['prev_key_padding_mask'] is not None:
|
||||
prev_key_padding_mask = saved_state['prev_key_padding_mask']
|
||||
if static_kv:
|
||||
key_padding_mask = prev_key_padding_mask
|
||||
else:
|
||||
key_padding_mask = torch.cat((prev_key_padding_mask, key_padding_mask), dim=1)
|
||||
|
||||
saved_state['prev_key'] = k.view(bsz, self.num_heads, -1, self.head_dim)
|
||||
saved_state['prev_value'] = v.view(bsz, self.num_heads, -1, self.head_dim)
|
||||
saved_state['prev_key_padding_mask'] = key_padding_mask
|
||||
|
||||
self._set_input_buffer(incremental_state, saved_state)
|
||||
|
||||
src_len = k.size(1)
|
||||
|
||||
# This is part of a workaround to get around fork/join parallelism
|
||||
# not supporting Optional types.
|
||||
if key_padding_mask is not None and key_padding_mask.shape == torch.Size([]):
|
||||
key_padding_mask = None
|
||||
|
||||
if key_padding_mask is not None:
|
||||
assert key_padding_mask.size(0) == bsz
|
||||
assert key_padding_mask.size(1) == src_len
|
||||
|
||||
if self.add_zero_attn:
|
||||
src_len += 1
|
||||
k = torch.cat([k, k.new_zeros((k.size(0), 1) + k.size()[2:])], dim=1)
|
||||
v = torch.cat([v, v.new_zeros((v.size(0), 1) + v.size()[2:])], dim=1)
|
||||
if attn_mask is not None:
|
||||
attn_mask = torch.cat([attn_mask, attn_mask.new_zeros(attn_mask.size(0), 1)], dim=1)
|
||||
if key_padding_mask is not None:
|
||||
key_padding_mask = torch.cat(
|
||||
[key_padding_mask, torch.zeros(key_padding_mask.size(0), 1).type_as(key_padding_mask)], dim=1)
|
||||
|
||||
attn_weights = torch.bmm(q, k.transpose(1, 2))
|
||||
attn_weights = self.apply_sparse_mask(attn_weights, tgt_len, src_len, bsz)
|
||||
|
||||
assert list(attn_weights.size()) == [bsz * self.num_heads, tgt_len, src_len]
|
||||
|
||||
if attn_mask is not None:
|
||||
if len(attn_mask.shape) == 2:
|
||||
attn_mask = attn_mask.unsqueeze(0)
|
||||
elif len(attn_mask.shape) == 3:
|
||||
attn_mask = attn_mask[:, None].repeat([1, self.num_heads, 1, 1]).reshape(
|
||||
bsz * self.num_heads, tgt_len, src_len)
|
||||
attn_weights = attn_weights + attn_mask
|
||||
|
||||
if enc_dec_attn_constraint_mask is not None: # bs x head x L_kv
|
||||
attn_weights = attn_weights.view(bsz, self.num_heads, tgt_len, src_len)
|
||||
attn_weights = attn_weights.masked_fill(
|
||||
enc_dec_attn_constraint_mask.unsqueeze(2).bool(),
|
||||
-1e8,
|
||||
)
|
||||
attn_weights = attn_weights.view(bsz * self.num_heads, tgt_len, src_len)
|
||||
|
||||
if key_padding_mask is not None:
|
||||
# don't attend to padding symbols
|
||||
attn_weights = attn_weights.view(bsz, self.num_heads, tgt_len, src_len)
|
||||
attn_weights = attn_weights.masked_fill(
|
||||
key_padding_mask.unsqueeze(1).unsqueeze(2),
|
||||
-1e8,
|
||||
)
|
||||
attn_weights = attn_weights.view(bsz * self.num_heads, tgt_len, src_len)
|
||||
|
||||
attn_logits = attn_weights.view(bsz, self.num_heads, tgt_len, src_len)
|
||||
|
||||
if before_softmax:
|
||||
return attn_weights, v
|
||||
|
||||
attn_weights_float = softmax(attn_weights, dim=-1)
|
||||
attn_weights = attn_weights_float.type_as(attn_weights)
|
||||
attn_probs = F.dropout(attn_weights_float.type_as(attn_weights), p=self.dropout, training=self.training)
|
||||
|
||||
if reset_attn_weight is not None:
|
||||
if reset_attn_weight:
|
||||
self.last_attn_probs = attn_probs.detach()
|
||||
else:
|
||||
assert self.last_attn_probs is not None
|
||||
attn_probs = self.last_attn_probs
|
||||
attn = torch.bmm(attn_probs, v)
|
||||
assert list(attn.size()) == [bsz * self.num_heads, tgt_len, self.head_dim]
|
||||
attn = attn.transpose(0, 1).contiguous().view(tgt_len, bsz, embed_dim)
|
||||
attn = self.out_proj(attn)
|
||||
|
||||
if need_weights:
|
||||
attn_weights = attn_weights_float.view(bsz, self.num_heads, tgt_len, src_len).transpose(1, 0)
|
||||
if not need_head_weights:
|
||||
# average attention weights over heads
|
||||
attn_weights = attn_weights.mean(dim=0)
|
||||
else:
|
||||
attn_weights = None
|
||||
|
||||
return attn, (attn_weights, attn_logits)
|
||||
|
||||
def in_proj_qkv(self, query):
|
||||
return self._in_proj(query).chunk(3, dim=-1)
|
||||
|
||||
def in_proj_q(self, query):
|
||||
if self.qkv_same_dim:
|
||||
return self._in_proj(query, end=self.embed_dim)
|
||||
else:
|
||||
bias = self.in_proj_bias
|
||||
if bias is not None:
|
||||
bias = bias[:self.embed_dim]
|
||||
return F.linear(query, self.q_proj_weight, bias)
|
||||
|
||||
def in_proj_k(self, key):
|
||||
if self.qkv_same_dim:
|
||||
return self._in_proj(key, start=self.embed_dim, end=2 * self.embed_dim)
|
||||
else:
|
||||
weight = self.k_proj_weight
|
||||
bias = self.in_proj_bias
|
||||
if bias is not None:
|
||||
bias = bias[self.embed_dim:2 * self.embed_dim]
|
||||
return F.linear(key, weight, bias)
|
||||
|
||||
def in_proj_v(self, value):
|
||||
if self.qkv_same_dim:
|
||||
return self._in_proj(value, start=2 * self.embed_dim)
|
||||
else:
|
||||
weight = self.v_proj_weight
|
||||
bias = self.in_proj_bias
|
||||
if bias is not None:
|
||||
bias = bias[2 * self.embed_dim:]
|
||||
return F.linear(value, weight, bias)
|
||||
|
||||
def _in_proj(self, input, start=0, end=None):
|
||||
weight = self.in_proj_weight
|
||||
bias = self.in_proj_bias
|
||||
weight = weight[start:end, :]
|
||||
if bias is not None:
|
||||
bias = bias[start:end]
|
||||
return F.linear(input, weight, bias)
|
||||
|
||||
def _get_input_buffer(self, incremental_state):
|
||||
return get_incremental_state(
|
||||
self,
|
||||
incremental_state,
|
||||
'attn_state',
|
||||
) or {}
|
||||
|
||||
def _set_input_buffer(self, incremental_state, buffer):
|
||||
set_incremental_state(
|
||||
self,
|
||||
incremental_state,
|
||||
'attn_state',
|
||||
buffer,
|
||||
)
|
||||
|
||||
def apply_sparse_mask(self, attn_weights, tgt_len, src_len, bsz):
|
||||
return attn_weights
|
||||
|
||||
def clear_buffer(self, incremental_state=None):
|
||||
if incremental_state is not None:
|
||||
saved_state = self._get_input_buffer(incremental_state)
|
||||
if 'prev_key' in saved_state:
|
||||
del saved_state['prev_key']
|
||||
if 'prev_value' in saved_state:
|
||||
del saved_state['prev_value']
|
||||
self._set_input_buffer(incremental_state, saved_state)
|
||||
|
||||
|
||||
class EncSALayer(nn.Module):
|
||||
def __init__(self, c, num_heads, dropout, attention_dropout=0.1,
|
||||
relu_dropout=0.1, kernel_size=9, padding='SAME', act='gelu'):
|
||||
super().__init__()
|
||||
self.c = c
|
||||
self.dropout = dropout
|
||||
self.num_heads = num_heads
|
||||
if num_heads > 0:
|
||||
self.layer_norm1 = LayerNorm(c)
|
||||
self.self_attn = MultiheadAttention(
|
||||
self.c, num_heads, self_attention=True, dropout=attention_dropout, bias=False)
|
||||
self.layer_norm2 = LayerNorm(c)
|
||||
self.ffn = TransformerFFNLayer(
|
||||
c, 4 * c, kernel_size=kernel_size, dropout=relu_dropout, padding=padding, act=act)
|
||||
|
||||
def forward(self, x, encoder_padding_mask=None, **kwargs):
|
||||
layer_norm_training = kwargs.get('layer_norm_training', None)
|
||||
if layer_norm_training is not None:
|
||||
self.layer_norm1.training = layer_norm_training
|
||||
self.layer_norm2.training = layer_norm_training
|
||||
if self.num_heads > 0:
|
||||
residual = x
|
||||
x = self.layer_norm1(x)
|
||||
x, _, = self.self_attn(
|
||||
query=x,
|
||||
key=x,
|
||||
value=x,
|
||||
key_padding_mask=encoder_padding_mask
|
||||
)
|
||||
x = F.dropout(x, self.dropout, training=self.training)
|
||||
x = residual + x
|
||||
x = x * (1 - encoder_padding_mask.float()).transpose(0, 1)[..., None]
|
||||
|
||||
residual = x
|
||||
x = self.layer_norm2(x)
|
||||
x = self.ffn(x)
|
||||
x = F.dropout(x, self.dropout, training=self.training)
|
||||
x = residual + x
|
||||
x = x * (1 - encoder_padding_mask.float()).transpose(0, 1)[..., None]
|
||||
return x
|
||||
|
||||
|
||||
class DecSALayer(nn.Module):
|
||||
def __init__(self, c, num_heads, dropout, attention_dropout=0.1, relu_dropout=0.1,
|
||||
kernel_size=9, act='gelu'):
|
||||
super().__init__()
|
||||
self.c = c
|
||||
self.dropout = dropout
|
||||
self.layer_norm1 = LayerNorm(c)
|
||||
self.self_attn = MultiheadAttention(
|
||||
c, num_heads, self_attention=True, dropout=attention_dropout, bias=False
|
||||
)
|
||||
self.layer_norm2 = LayerNorm(c)
|
||||
self.encoder_attn = MultiheadAttention(
|
||||
c, num_heads, encoder_decoder_attention=True, dropout=attention_dropout, bias=False,
|
||||
)
|
||||
self.layer_norm3 = LayerNorm(c)
|
||||
self.ffn = TransformerFFNLayer(
|
||||
c, 4 * c, padding='LEFT', kernel_size=kernel_size, dropout=relu_dropout, act=act)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x,
|
||||
encoder_out=None,
|
||||
encoder_padding_mask=None,
|
||||
incremental_state=None,
|
||||
self_attn_mask=None,
|
||||
self_attn_padding_mask=None,
|
||||
attn_out=None,
|
||||
reset_attn_weight=None,
|
||||
**kwargs,
|
||||
):
|
||||
layer_norm_training = kwargs.get('layer_norm_training', None)
|
||||
if layer_norm_training is not None:
|
||||
self.layer_norm1.training = layer_norm_training
|
||||
self.layer_norm2.training = layer_norm_training
|
||||
self.layer_norm3.training = layer_norm_training
|
||||
residual = x
|
||||
x = self.layer_norm1(x)
|
||||
x, _ = self.self_attn(
|
||||
query=x,
|
||||
key=x,
|
||||
value=x,
|
||||
key_padding_mask=self_attn_padding_mask,
|
||||
incremental_state=incremental_state,
|
||||
attn_mask=self_attn_mask
|
||||
)
|
||||
x = F.dropout(x, self.dropout, training=self.training)
|
||||
x = residual + x
|
||||
|
||||
attn_logits = None
|
||||
if encoder_out is not None or attn_out is not None:
|
||||
residual = x
|
||||
x = self.layer_norm2(x)
|
||||
if encoder_out is not None:
|
||||
x, attn = self.encoder_attn(
|
||||
query=x,
|
||||
key=encoder_out,
|
||||
value=encoder_out,
|
||||
key_padding_mask=encoder_padding_mask,
|
||||
incremental_state=incremental_state,
|
||||
static_kv=True,
|
||||
enc_dec_attn_constraint_mask=get_incremental_state(self, incremental_state,
|
||||
'enc_dec_attn_constraint_mask'),
|
||||
reset_attn_weight=reset_attn_weight
|
||||
)
|
||||
attn_logits = attn[1]
|
||||
elif attn_out is not None:
|
||||
x = self.encoder_attn.in_proj_v(attn_out)
|
||||
if encoder_out is not None or attn_out is not None:
|
||||
x = F.dropout(x, self.dropout, training=self.training)
|
||||
x = residual + x
|
||||
|
||||
residual = x
|
||||
x = self.layer_norm3(x)
|
||||
x = self.ffn(x, incremental_state=incremental_state)
|
||||
x = F.dropout(x, self.dropout, training=self.training)
|
||||
x = residual + x
|
||||
return x, attn_logits
|
||||
|
||||
def clear_buffer(self, input, encoder_out=None, encoder_padding_mask=None, incremental_state=None):
|
||||
self.encoder_attn.clear_buffer(incremental_state)
|
||||
self.ffn.clear_buffer(incremental_state)
|
||||
|
||||
def set_buffer(self, name, tensor, incremental_state):
|
||||
return set_incremental_state(self, incremental_state, name, tensor)
|
||||
|
||||
|
||||
class TransformerEncoderLayer(nn.Module):
|
||||
def __init__(self, hidden_size, dropout, kernel_size=9, num_heads=2):
|
||||
super().__init__()
|
||||
self.hidden_size = hidden_size
|
||||
self.dropout = dropout
|
||||
self.num_heads = num_heads
|
||||
self.op = EncSALayer(
|
||||
hidden_size, num_heads, dropout=dropout,
|
||||
attention_dropout=0.0, relu_dropout=dropout,
|
||||
kernel_size=kernel_size)
|
||||
|
||||
def forward(self, x, **kwargs):
|
||||
return self.op(x, **kwargs)
|
||||
|
||||
|
||||
class TransformerDecoderLayer(nn.Module):
|
||||
def __init__(self, hidden_size, dropout, kernel_size=9, num_heads=2):
|
||||
super().__init__()
|
||||
self.hidden_size = hidden_size
|
||||
self.dropout = dropout
|
||||
self.num_heads = num_heads
|
||||
self.op = DecSALayer(
|
||||
hidden_size, num_heads, dropout=dropout,
|
||||
attention_dropout=0.0, relu_dropout=dropout,
|
||||
kernel_size=kernel_size)
|
||||
|
||||
def forward(self, x, **kwargs):
|
||||
return self.op(x, **kwargs)
|
||||
|
||||
def clear_buffer(self, *args):
|
||||
return self.op.clear_buffer(*args)
|
||||
|
||||
def set_buffer(self, *args):
|
||||
return self.op.set_buffer(*args)
|
||||
|
||||
|
||||
class FFTBlocks(nn.Module):
|
||||
def __init__(self, hidden_size, num_layers, ffn_kernel_size=9, dropout=0.0,
|
||||
num_heads=2, use_pos_embed=True, use_last_norm=True,
|
||||
use_pos_embed_alpha=True):
|
||||
super().__init__()
|
||||
self.num_layers = num_layers
|
||||
embed_dim = self.hidden_size = hidden_size
|
||||
self.dropout = dropout
|
||||
self.use_pos_embed = use_pos_embed
|
||||
self.use_last_norm = use_last_norm
|
||||
if use_pos_embed:
|
||||
self.max_source_positions = DEFAULT_MAX_TARGET_POSITIONS
|
||||
self.padding_idx = 0
|
||||
self.pos_embed_alpha = nn.Parameter(torch.Tensor([1])) if use_pos_embed_alpha else 1
|
||||
self.embed_positions = SinusoidalPositionalEmbedding(
|
||||
embed_dim, self.padding_idx, init_size=DEFAULT_MAX_TARGET_POSITIONS,
|
||||
)
|
||||
|
||||
self.layers = nn.ModuleList([])
|
||||
self.layers.extend([
|
||||
TransformerEncoderLayer(self.hidden_size, self.dropout,
|
||||
kernel_size=ffn_kernel_size, num_heads=num_heads)
|
||||
for _ in range(self.num_layers)
|
||||
])
|
||||
if self.use_last_norm:
|
||||
self.layer_norm = nn.LayerNorm(embed_dim)
|
||||
else:
|
||||
self.layer_norm = None
|
||||
|
||||
def forward(self, x, padding_mask=None, attn_mask=None, return_hiddens=False):
|
||||
"""
|
||||
:param x: [B, T, C]
|
||||
:param padding_mask: [B, T]
|
||||
:return: [B, T, C] or [L, B, T, C]
|
||||
"""
|
||||
padding_mask = x.abs().sum(-1).eq(0).data if padding_mask is None else padding_mask
|
||||
nonpadding_mask_TB = 1 - padding_mask.transpose(0, 1).float()[:, :, None] # [T, B, 1]
|
||||
if self.use_pos_embed:
|
||||
positions = self.pos_embed_alpha * self.embed_positions(x[..., 0])
|
||||
x = x + positions
|
||||
x = F.dropout(x, p=self.dropout, training=self.training)
|
||||
# B x T x C -> T x B x C
|
||||
x = x.transpose(0, 1) * nonpadding_mask_TB
|
||||
hiddens = []
|
||||
for layer in self.layers:
|
||||
x = layer(x, encoder_padding_mask=padding_mask, attn_mask=attn_mask) * nonpadding_mask_TB
|
||||
hiddens.append(x)
|
||||
if self.use_last_norm:
|
||||
x = self.layer_norm(x) * nonpadding_mask_TB
|
||||
if return_hiddens:
|
||||
x = torch.stack(hiddens, 0) # [L, T, B, C]
|
||||
x = x.transpose(1, 2) # [L, B, T, C]
|
||||
else:
|
||||
x = x.transpose(0, 1) # [B, T, C]
|
||||
return x
|
||||
|
||||
|
||||
class FastSpeechEncoder(FFTBlocks):
|
||||
def __init__(self, dict_size, hidden_size=256, num_layers=4, kernel_size=9, num_heads=2,
|
||||
dropout=0.0):
|
||||
super().__init__(hidden_size, num_layers, kernel_size, num_heads=num_heads,
|
||||
use_pos_embed=False, dropout=dropout) # use_pos_embed_alpha for compatibility
|
||||
self.embed_tokens = Embedding(dict_size, hidden_size, 0)
|
||||
self.embed_scale = math.sqrt(hidden_size)
|
||||
self.padding_idx = 0
|
||||
self.embed_positions = SinusoidalPositionalEmbedding(
|
||||
hidden_size, self.padding_idx, init_size=DEFAULT_MAX_TARGET_POSITIONS,
|
||||
)
|
||||
|
||||
def forward(self, txt_tokens, attn_mask=None):
|
||||
"""
|
||||
|
||||
:param txt_tokens: [B, T]
|
||||
:return: {
|
||||
'encoder_out': [B x T x C]
|
||||
}
|
||||
"""
|
||||
encoder_padding_mask = txt_tokens.eq(self.padding_idx).data
|
||||
x = self.forward_embedding(txt_tokens) # [B, T, H]
|
||||
if self.num_layers > 0:
|
||||
x = super(FastSpeechEncoder, self).forward(x, encoder_padding_mask, attn_mask=attn_mask)
|
||||
return x
|
||||
|
||||
def forward_embedding(self, txt_tokens):
|
||||
# embed tokens and positions
|
||||
x = self.embed_scale * self.embed_tokens(txt_tokens)
|
||||
positions = self.embed_positions(txt_tokens)
|
||||
x = x + positions
|
||||
x = F.dropout(x, p=self.dropout, training=self.training)
|
||||
return x
|
||||
|
||||
|
||||
class FastSpeechDecoder(FFTBlocks):
|
||||
def __init__(self, hidden_size=256, num_layers=4, kernel_size=9, num_heads=2):
|
||||
super().__init__(hidden_size, num_layers, kernel_size, num_heads=num_heads)
|
||||
@@ -0,0 +1,109 @@
|
||||
import torch
|
||||
from torch import nn
|
||||
from packaging import version
|
||||
|
||||
def fused_add_tanh_sigmoid_multiply(input_a, input_b, n_channels):
|
||||
n_channels_int = n_channels[0]
|
||||
in_act = input_a + input_b
|
||||
t_act = torch.tanh(in_act[:, :n_channels_int, :])
|
||||
s_act = torch.sigmoid(in_act[:, n_channels_int:, :])
|
||||
acts = t_act * s_act
|
||||
return acts
|
||||
|
||||
jit_fused_add_tanh_sigmoid_multiply = fused_add_tanh_sigmoid_multiply
|
||||
|
||||
def script_function():
|
||||
if version.parse(torch.__version__) >= version.parse('2.0'):
|
||||
global jit_fused_add_tanh_sigmoid_multiply
|
||||
jit_fused_add_tanh_sigmoid_multiply = torch.jit.script(fused_add_tanh_sigmoid_multiply)
|
||||
|
||||
|
||||
class WN(torch.nn.Module):
|
||||
def __init__(self, hidden_size, kernel_size, dilation_rate, n_layers, c_cond=0,
|
||||
p_dropout=0, share_cond_layers=False, is_BTC=False):
|
||||
super(WN, self).__init__()
|
||||
assert (kernel_size % 2 == 1)
|
||||
assert (hidden_size % 2 == 0)
|
||||
self.is_BTC = is_BTC
|
||||
self.hidden_size = hidden_size
|
||||
self.kernel_size = kernel_size
|
||||
self.dilation_rate = dilation_rate
|
||||
self.n_layers = n_layers
|
||||
self.gin_channels = c_cond
|
||||
self.p_dropout = p_dropout
|
||||
self.share_cond_layers = share_cond_layers
|
||||
|
||||
self.in_layers = torch.nn.ModuleList()
|
||||
self.res_skip_layers = torch.nn.ModuleList()
|
||||
self.drop = nn.Dropout(p_dropout)
|
||||
|
||||
if c_cond != 0 and not share_cond_layers:
|
||||
cond_layer = torch.nn.Conv1d(c_cond, 2 * hidden_size * n_layers, 1)
|
||||
self.cond_layer = torch.nn.utils.weight_norm(cond_layer, name='weight')
|
||||
|
||||
for i in range(n_layers):
|
||||
dilation = dilation_rate ** i
|
||||
padding = int((kernel_size * dilation - dilation) / 2)
|
||||
in_layer = torch.nn.Conv1d(hidden_size, 2 * hidden_size, kernel_size,
|
||||
dilation=dilation, padding=padding)
|
||||
in_layer = torch.nn.utils.weight_norm(in_layer, name='weight')
|
||||
self.in_layers.append(in_layer)
|
||||
|
||||
# last one is not necessary
|
||||
if i < n_layers - 1:
|
||||
res_skip_channels = 2 * hidden_size
|
||||
else:
|
||||
res_skip_channels = hidden_size
|
||||
|
||||
res_skip_layer = torch.nn.Conv1d(hidden_size, res_skip_channels, 1)
|
||||
res_skip_layer = torch.nn.utils.weight_norm(res_skip_layer, name='weight')
|
||||
self.res_skip_layers.append(res_skip_layer)
|
||||
|
||||
script_function()
|
||||
|
||||
def forward(self, x, nonpadding=None, cond=None):
|
||||
if self.is_BTC:
|
||||
x = x.transpose(1, 2)
|
||||
cond = cond.transpose(1, 2) if cond is not None else None
|
||||
nonpadding = nonpadding.transpose(1, 2) if nonpadding is not None else None
|
||||
if nonpadding is None:
|
||||
nonpadding = 1
|
||||
output = torch.zeros_like(x)
|
||||
n_channels_tensor = torch.IntTensor([self.hidden_size])
|
||||
|
||||
if cond is not None and not self.share_cond_layers:
|
||||
cond = self.cond_layer(cond)
|
||||
|
||||
for i in range(self.n_layers):
|
||||
x_in = self.in_layers[i](x)
|
||||
x_in = self.drop(x_in)
|
||||
if cond is not None:
|
||||
cond_offset = i * 2 * self.hidden_size
|
||||
cond_l = cond[:, cond_offset:cond_offset + 2 * self.hidden_size, :]
|
||||
else:
|
||||
cond_l = torch.zeros_like(x_in)
|
||||
|
||||
if version.parse(torch.__version__) >= version.parse('2.0'):
|
||||
acts = jit_fused_add_tanh_sigmoid_multiply(x_in, cond_l, n_channels_tensor)
|
||||
else:
|
||||
acts = fused_add_tanh_sigmoid_multiply(x_in, cond_l, n_channels_tensor)
|
||||
|
||||
res_skip_acts = self.res_skip_layers[i](acts)
|
||||
if i < self.n_layers - 1:
|
||||
x = (x + res_skip_acts[:, :self.hidden_size, :]) * nonpadding
|
||||
output = output + res_skip_acts[:, self.hidden_size:, :]
|
||||
else:
|
||||
output = output + res_skip_acts
|
||||
output = output * nonpadding
|
||||
if self.is_BTC:
|
||||
output = output.transpose(1, 2)
|
||||
return output
|
||||
|
||||
def remove_weight_norm(self):
|
||||
def remove_weight_norm(m):
|
||||
try:
|
||||
nn.utils.remove_weight_norm(m)
|
||||
except ValueError: # this module didn't have weight norm
|
||||
return
|
||||
|
||||
self.apply(remove_weight_norm)
|
||||
@@ -0,0 +1 @@
|
||||
"""Pitch extractor modules for ROSVOT."""
|
||||
@@ -0,0 +1,6 @@
|
||||
from .constants import *
|
||||
from .model import E2E0
|
||||
from .utils import to_local_average_f0, to_viterbi_f0
|
||||
from .inference import RMVPE
|
||||
from .spec import MelSpectrogram
|
||||
from .extractor import extract
|
||||
@@ -0,0 +1,9 @@
|
||||
SAMPLE_RATE = 16000
|
||||
|
||||
N_CLASS = 360
|
||||
|
||||
N_MELS = 128
|
||||
MEL_FMIN = 30
|
||||
MEL_FMAX = 8000
|
||||
WINDOW_LENGTH = 1024
|
||||
CONST = 1997.3794084376191
|
||||
@@ -0,0 +1,173 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from .constants import N_MELS
|
||||
|
||||
|
||||
class ConvBlockRes(nn.Module):
|
||||
def __init__(self, in_channels, out_channels, momentum=0.01):
|
||||
super(ConvBlockRes, self).__init__()
|
||||
self.conv = nn.Sequential(
|
||||
nn.Conv2d(in_channels=in_channels,
|
||||
out_channels=out_channels,
|
||||
kernel_size=(3, 3),
|
||||
stride=(1, 1),
|
||||
padding=(1, 1),
|
||||
bias=False),
|
||||
nn.BatchNorm2d(out_channels, momentum=momentum),
|
||||
nn.ReLU(),
|
||||
|
||||
nn.Conv2d(in_channels=out_channels,
|
||||
out_channels=out_channels,
|
||||
kernel_size=(3, 3),
|
||||
stride=(1, 1),
|
||||
padding=(1, 1),
|
||||
bias=False),
|
||||
nn.BatchNorm2d(out_channels, momentum=momentum),
|
||||
nn.ReLU(),
|
||||
)
|
||||
if in_channels != out_channels:
|
||||
self.shortcut = nn.Conv2d(in_channels, out_channels, (1, 1))
|
||||
self.is_shortcut = True
|
||||
else:
|
||||
self.is_shortcut = False
|
||||
|
||||
def forward(self, x):
|
||||
if self.is_shortcut:
|
||||
return self.conv(x) + self.shortcut(x)
|
||||
else:
|
||||
return self.conv(x) + x
|
||||
|
||||
|
||||
class ResEncoderBlock(nn.Module):
|
||||
def __init__(self, in_channels, out_channels, kernel_size, n_blocks=1, momentum=0.01):
|
||||
super(ResEncoderBlock, self).__init__()
|
||||
self.n_blocks = n_blocks
|
||||
self.conv = nn.ModuleList()
|
||||
self.conv.append(ConvBlockRes(in_channels, out_channels, momentum))
|
||||
for i in range(n_blocks - 1):
|
||||
self.conv.append(ConvBlockRes(out_channels, out_channels, momentum))
|
||||
self.kernel_size = kernel_size
|
||||
if self.kernel_size is not None:
|
||||
self.pool = nn.AvgPool2d(kernel_size=kernel_size)
|
||||
|
||||
def forward(self, x):
|
||||
for i in range(self.n_blocks):
|
||||
x = self.conv[i](x)
|
||||
if self.kernel_size is not None:
|
||||
return x, self.pool(x)
|
||||
else:
|
||||
return x
|
||||
|
||||
|
||||
class ResDecoderBlock(nn.Module):
|
||||
def __init__(self, in_channels, out_channels, stride, n_blocks=1, momentum=0.01):
|
||||
super(ResDecoderBlock, self).__init__()
|
||||
out_padding = (0, 1) if stride == (1, 2) else (1, 1)
|
||||
self.n_blocks = n_blocks
|
||||
self.conv1 = nn.Sequential(
|
||||
nn.ConvTranspose2d(in_channels=in_channels,
|
||||
out_channels=out_channels,
|
||||
kernel_size=(3, 3),
|
||||
stride=stride,
|
||||
padding=(1, 1),
|
||||
output_padding=out_padding,
|
||||
bias=False),
|
||||
nn.BatchNorm2d(out_channels, momentum=momentum),
|
||||
nn.ReLU(),
|
||||
)
|
||||
self.conv2 = nn.ModuleList()
|
||||
self.conv2.append(ConvBlockRes(out_channels * 2, out_channels, momentum))
|
||||
for i in range(n_blocks-1):
|
||||
self.conv2.append(ConvBlockRes(out_channels, out_channels, momentum))
|
||||
|
||||
def forward(self, x, concat_tensor):
|
||||
x = self.conv1(x)
|
||||
x = torch.cat((x, concat_tensor), dim=1)
|
||||
for i in range(self.n_blocks):
|
||||
x = self.conv2[i](x)
|
||||
return x
|
||||
|
||||
|
||||
class Encoder(nn.Module):
|
||||
def __init__(self, in_channels, in_size, n_encoders, kernel_size, n_blocks, out_channels=16, momentum=0.01):
|
||||
super(Encoder, self).__init__()
|
||||
self.n_encoders = n_encoders
|
||||
self.bn = nn.BatchNorm2d(in_channels, momentum=momentum)
|
||||
self.layers = nn.ModuleList()
|
||||
self.latent_channels = []
|
||||
for i in range(self.n_encoders):
|
||||
self.layers.append(ResEncoderBlock(in_channels, out_channels, kernel_size, n_blocks, momentum=momentum))
|
||||
self.latent_channels.append([out_channels, in_size])
|
||||
in_channels = out_channels
|
||||
out_channels *= 2
|
||||
in_size //= 2
|
||||
self.out_size = in_size
|
||||
self.out_channel = out_channels
|
||||
|
||||
def forward(self, x):
|
||||
concat_tensors = []
|
||||
x = self.bn(x)
|
||||
for i in range(self.n_encoders):
|
||||
_, x = self.layers[i](x)
|
||||
concat_tensors.append(_)
|
||||
return x, concat_tensors
|
||||
|
||||
|
||||
class Intermediate(nn.Module):
|
||||
def __init__(self, in_channels, out_channels, n_inters, n_blocks, momentum=0.01):
|
||||
super(Intermediate, self).__init__()
|
||||
self.n_inters = n_inters
|
||||
self.layers = nn.ModuleList()
|
||||
self.layers.append(ResEncoderBlock(in_channels, out_channels, None, n_blocks, momentum))
|
||||
for i in range(self.n_inters-1):
|
||||
self.layers.append(ResEncoderBlock(out_channels, out_channels, None, n_blocks, momentum))
|
||||
|
||||
def forward(self, x):
|
||||
for i in range(self.n_inters):
|
||||
x = self.layers[i](x)
|
||||
return x
|
||||
|
||||
|
||||
class Decoder(nn.Module):
|
||||
def __init__(self, in_channels, n_decoders, stride, n_blocks, momentum=0.01):
|
||||
super(Decoder, self).__init__()
|
||||
self.layers = nn.ModuleList()
|
||||
self.n_decoders = n_decoders
|
||||
for i in range(self.n_decoders):
|
||||
out_channels = in_channels // 2
|
||||
self.layers.append(ResDecoderBlock(in_channels, out_channels, stride, n_blocks, momentum))
|
||||
in_channels = out_channels
|
||||
|
||||
def forward(self, x, concat_tensors):
|
||||
for i in range(self.n_decoders):
|
||||
x = self.layers[i](x, concat_tensors[-1-i])
|
||||
return x
|
||||
|
||||
|
||||
class TimbreFilter(nn.Module):
|
||||
def __init__(self, latent_rep_channels):
|
||||
super(TimbreFilter, self).__init__()
|
||||
self.layers = nn.ModuleList()
|
||||
for latent_rep in latent_rep_channels:
|
||||
self.layers.append(ConvBlockRes(latent_rep[0], latent_rep[0]))
|
||||
|
||||
def forward(self, x_tensors):
|
||||
out_tensors = []
|
||||
for i, layer in enumerate(self.layers):
|
||||
out_tensors.append(layer(x_tensors[i]))
|
||||
return out_tensors
|
||||
|
||||
|
||||
class DeepUnet0(nn.Module):
|
||||
def __init__(self, kernel_size, n_blocks, en_de_layers=5, inter_layers=4, in_channels=1, en_out_channels=16):
|
||||
super(DeepUnet0, self).__init__()
|
||||
self.encoder = Encoder(in_channels, N_MELS, en_de_layers, kernel_size, n_blocks, en_out_channels)
|
||||
self.intermediate = Intermediate(self.encoder.out_channel // 2, self.encoder.out_channel, inter_layers, n_blocks)
|
||||
self.tf = TimbreFilter(self.encoder.latent_channels)
|
||||
self.decoder = Decoder(self.encoder.out_channel, en_de_layers, kernel_size, n_blocks)
|
||||
|
||||
def forward(self, x):
|
||||
x, concat_tensors = self.encoder(x)
|
||||
x = self.intermediate(x)
|
||||
x = self.decoder(x, concat_tensors)
|
||||
return x
|
||||
@@ -0,0 +1,183 @@
|
||||
import math
|
||||
import os
|
||||
|
||||
from tqdm import tqdm
|
||||
import librosa
|
||||
import numpy as np
|
||||
import torch
|
||||
from torch.utils.data import Dataset, DataLoader, DistributedSampler
|
||||
import torch.multiprocessing as mp
|
||||
from torch.distributed import init_process_group
|
||||
import torch.distributed as dist
|
||||
|
||||
from .inference import RMVPE
|
||||
from ....utils.commons.dataset_utils import batch_by_size, build_dataloader
|
||||
# import utils
|
||||
from ....utils.audio import get_wav_num_frames
|
||||
|
||||
"""
|
||||
A convenient API for batch inference
|
||||
update: add ddp
|
||||
"""
|
||||
|
||||
class RMVPEInferDataset(Dataset):
|
||||
def __init__(self, wav_fns: list, id_and_sizes=None, sr=24000, hop_size=128, num_workers=0):
|
||||
if id_and_sizes is None:
|
||||
id_and_sizes = []
|
||||
if type(wav_fns[0]) == str: # wav_paths
|
||||
for idx, wav_path in enumerate(wav_fns):
|
||||
total_frames = get_wav_num_frames(wav_path, sr)
|
||||
id_and_sizes.append((idx, round(total_frames / hop_size)))
|
||||
else: # numpy arrays, mono wavs
|
||||
for idx, wav in enumerate(wav_fns):
|
||||
id_and_sizes.append((idx, round(wav.shape[-1] / hop_size)))
|
||||
self.wav_fns = wav_fns
|
||||
self.id_and_sizes = id_and_sizes
|
||||
self.sr = sr
|
||||
self.num_workers = num_workers
|
||||
|
||||
def __getitem__(self, idx):
|
||||
if type(self.wav_fns[idx]) == str:
|
||||
wav_fn = self.wav_fns[idx]
|
||||
wav, _ = librosa.core.load(wav_fn, sr=self.sr)
|
||||
else:
|
||||
wav = self.wav_fns[idx]
|
||||
return idx, wav
|
||||
|
||||
def collater(self, samples: list):
|
||||
return samples
|
||||
|
||||
def __len__(self):
|
||||
return len(self.wav_fns)
|
||||
|
||||
def ordered_indices(self):
|
||||
"""Return an ordered list of indices. Batches will be constructed based
|
||||
on this order."""
|
||||
return np.arange(len(self))
|
||||
|
||||
def num_tokens(self, index):
|
||||
return self.id_and_sizes[index][1]
|
||||
|
||||
@torch.no_grad()
|
||||
def extract(wav_fns: list, id_and_sizes=None, ckpt=None, sr=24000, hop_size=128, bsz=128, max_tokens=100000,
|
||||
fmax=900, fmin=50, ds_workers=0):
|
||||
all_gpu_ids = [int(x) for x in os.environ.get("CUDA_VISIBLE_DEVICES", "").split(",") if x != '']
|
||||
num_gpus = len(all_gpu_ids)
|
||||
dist_config = {
|
||||
"dist_backend": "nccl",
|
||||
"dist_url": "tcp://localhost:54189",
|
||||
"world_size": 1
|
||||
}
|
||||
# https://discuss.pytorch.org/t/how-to-fix-a-sigsegv-in-pytorch-when-using-distributed-training-e-g-ddp/113518/10#:~:text=Using%20start%20and%20join%20avoids
|
||||
# https://github.com/pytorch/pytorch/issues/40403#issuecomment-648515174
|
||||
# mp.set_start_method('spawn')
|
||||
if num_gpus > 1:
|
||||
result_queue = mp.Queue()
|
||||
for rank in range(num_gpus):
|
||||
mp.Process(target=extract_worker, args=(rank, wav_fns, id_and_sizes, ckpt, sr, hop_size, bsz, max_tokens, fmax,
|
||||
fmin, dist_config, num_gpus, ds_workers, result_queue,)).start()
|
||||
f0_res = [None] * len(wav_fns)
|
||||
for _ in range(num_gpus):
|
||||
f0_res_dict = result_queue.get()
|
||||
for idx in f0_res_dict:
|
||||
f0_res[idx] = f0_res_dict[idx]
|
||||
del f0_res_dict
|
||||
else:
|
||||
# f0_res = extract_one_process(wav_fns, id_and_sizes, ckpt, sr, hop_size, bsz, max_tokens, fmax, fmin)
|
||||
f0_res_dict = extract_worker(0, wav_fns, id_and_sizes, ckpt, sr, hop_size, bsz, max_tokens, fmax,
|
||||
fmin, dist_config, num_gpus, ds_workers, None)
|
||||
f0_res = [None] * len(wav_fns)
|
||||
for idx in f0_res_dict:
|
||||
f0_res[idx] = f0_res_dict[idx]
|
||||
return f0_res
|
||||
|
||||
@torch.no_grad()
|
||||
def extract_worker(rank, wav_fns: list, id_and_sizes=None, ckpt=None, sr=24000, hop_size=128, bsz=128, max_tokens=100000,
|
||||
fmax=900, fmin=50, dist_config=None, num_gpus=1, ds_workers=0, q=None):
|
||||
# print(f"rank: {rank}")
|
||||
if num_gpus > 1:
|
||||
init_process_group(backend=dist_config['dist_backend'], init_method=dist_config['dist_url'],
|
||||
world_size=dist_config['world_size'] * num_gpus, rank=rank)
|
||||
dataset = RMVPEInferDataset(wav_fns, id_and_sizes, sr, hop_size, num_workers=ds_workers)
|
||||
# ds_sampler = DistributedSampler(dataset, shuffle=False) if num_gpus > 1 else None
|
||||
# loader = DataLoader(dataset, sampler=ds_sampler, collate_fn=dataset.collator, batch_size=1, num_workers=40, drop_last=False)
|
||||
loader = build_dataloader(dataset, shuffle=False, max_tokens=max_tokens, max_sentences=bsz, use_ddp=num_gpus > 1)
|
||||
loader = tqdm(loader, desc=f'| Processing f0 in [n_ranks={num_gpus}; max_tokens={max_tokens}; max_sentences={bsz}]') if rank == 0 else loader
|
||||
|
||||
device = torch.device(f"cuda:{int(rank)}")
|
||||
model = RMVPE(ckpt, device=device)
|
||||
f0_res_dict = {}
|
||||
for batch in loader:
|
||||
if batch is None or len(batch) == 0:
|
||||
continue
|
||||
idxs = [item[0] for item in batch]
|
||||
wavs = [item[1] for item in batch]
|
||||
lengths = [(wav.shape[0] + hop_size - 1) // hop_size for wav in wavs]
|
||||
with torch.no_grad():
|
||||
f0s, uvs = model.get_pitch_batch(
|
||||
wavs, sample_rate=sr,
|
||||
hop_size=hop_size,
|
||||
lengths=lengths,
|
||||
fmax=fmax,
|
||||
fmin=fmin
|
||||
)
|
||||
for i, idx in enumerate(idxs):
|
||||
f0_res_dict[idx] = f0s[i]
|
||||
if q is not None:
|
||||
q.put(f0_res_dict)
|
||||
else:
|
||||
return f0_res_dict
|
||||
|
||||
# old version
|
||||
def extract_one_process(wav_fns: list, id_and_sizes=None, ckpt=None, sr=24000, hop_size=128, bsz=128, max_tokens=100000,
|
||||
fmax=900, fmin=50, device='cuda'):
|
||||
assert ckpt is not None
|
||||
rmvpe = RMVPE(ckpt, device=device)
|
||||
if id_and_sizes is None:
|
||||
id_and_sizes = []
|
||||
if type(wav_fns[0]) == str: # wav_paths
|
||||
for idx, wav_path in enumerate(wav_fns):
|
||||
total_frames = get_wav_num_frames(wav_path, sr)
|
||||
id_and_sizes.append((idx, round(total_frames / hop_size)))
|
||||
else: # numpy arrays, mono wavs
|
||||
for idx, wav in enumerate(wav_fns):
|
||||
id_and_sizes.append((idx, round(wav.shape[-1] / hop_size)))
|
||||
get_size = lambda x: x[1]
|
||||
bs = batch_by_size(id_and_sizes, get_size, max_tokens=max_tokens, max_sentences=bsz)
|
||||
for i in range(len(bs)):
|
||||
bs[i] = [bs[i][j][0] for j in range(len(bs[i]))]
|
||||
|
||||
f0_res = [None] * len(wav_fns)
|
||||
for batch in tqdm(bs, total=len(bs), desc=f'| Processing f0 in [max_tokens={max_tokens}; max_sentences={bsz}]'):
|
||||
wavs, mel_lengths, lengths = [], [], []
|
||||
for idx in batch:
|
||||
if type(wav_fns[idx]) == str:
|
||||
wav_fn = wav_fns[idx]
|
||||
wav, _ = librosa.core.load(wav_fn, sr=sr)
|
||||
else:
|
||||
wav = wav_fns[idx]
|
||||
wavs.append(wav)
|
||||
mel_lengths.append(math.ceil((wav.shape[0] + 1) / hop_size))
|
||||
lengths.append((wav.shape[0] + hop_size - 1) // hop_size)
|
||||
|
||||
with torch.no_grad():
|
||||
f0s, uvs = rmvpe.get_pitch_batch(
|
||||
wavs, sample_rate=sr,
|
||||
hop_size=hop_size,
|
||||
lengths=lengths,
|
||||
fmax=fmax,
|
||||
fmin=fmin
|
||||
)
|
||||
|
||||
for i, idx in enumerate(batch):
|
||||
f0_res[idx] = f0s[i]
|
||||
|
||||
if rmvpe is not None:
|
||||
rmvpe.release_cuda()
|
||||
torch.cuda.empty_cache()
|
||||
rmvpe = None
|
||||
|
||||
return f0_res
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,134 @@
|
||||
import math
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torchaudio.transforms import Resample
|
||||
import pyworld as pw
|
||||
|
||||
from ....utils.audio.pitch_utils import interp_f0, resample_align_curve
|
||||
from .constants import *
|
||||
from .model import E2E0
|
||||
from .spec import MelSpectrogram
|
||||
from .utils import to_local_average_f0, to_viterbi_f0
|
||||
|
||||
|
||||
class RMVPE:
|
||||
def __init__(self, model_path, hop_length=160, device=None):
|
||||
self.resample_kernel = {}
|
||||
if device is None:
|
||||
self.device = 'cuda' if torch.cuda.is_available() else 'cpu'
|
||||
else:
|
||||
self.device = device
|
||||
self.model = E2E0(4, 1, (2, 2)).eval().to(self.device)
|
||||
ckpt = torch.load(model_path, map_location=self.device)
|
||||
self.model.load_state_dict(ckpt['model'], strict=False)
|
||||
self.mel_extractor = MelSpectrogram(
|
||||
N_MELS, SAMPLE_RATE, WINDOW_LENGTH, hop_length, None, MEL_FMIN, MEL_FMAX
|
||||
).to(self.device)
|
||||
self.hop_length = hop_length
|
||||
|
||||
@torch.no_grad()
|
||||
def mel2hidden(self, mel):
|
||||
n_frames = mel.shape[-1]
|
||||
mel = F.pad(mel, (0, 32 * ((n_frames - 1) // 32 + 1) - n_frames), mode='constant')
|
||||
hidden = self.model(mel)
|
||||
return hidden[:, :n_frames]
|
||||
|
||||
def decode(self, hidden, thred=0.03, use_viterbi=False):
|
||||
if use_viterbi:
|
||||
f0 = to_viterbi_f0(hidden, thred=thred)
|
||||
else:
|
||||
f0 = to_local_average_f0(hidden, thred=thred)
|
||||
return f0
|
||||
|
||||
def postprocess(self, f0, fmin=50, fmax=1000, audio=None, min_gap=2):
|
||||
if audio is not None:
|
||||
# this doesn't work. deprecated
|
||||
t = np.arange(0, f0.shape[0] * self.hop_length / 16000, self.hop_length / 16000)
|
||||
f0 = pw.stonemask(audio.astype(np.float64), f0.astype(np.float64), t, 16000).astype(float)
|
||||
f0[f0 < fmin] = 0
|
||||
f0[f0 > fmax] = 0
|
||||
# eliminate glitch
|
||||
# min_gap: if successive positive f0 positions < min_gap, zero these positions
|
||||
# eg: if min_gap=2, [0, 500, 500, 0] => [0, 0, 0, 0]
|
||||
for idx in range(f0.shape[0] - min_gap - 1):
|
||||
if f0[idx] == 0 and f0[idx + min_gap + 1] == 0 and np.sum(f0[idx: idx + min_gap + 2]) > 0:
|
||||
f0[idx: idx + min_gap + 2] = 0
|
||||
return f0
|
||||
|
||||
def infer_from_audio(self, audio, sample_rate=16000, thred=0.03, use_viterbi=False):
|
||||
audio = torch.from_numpy(audio).float().unsqueeze(0).to(self.device)
|
||||
if sample_rate == 16000:
|
||||
audio_res = audio
|
||||
else:
|
||||
key_str = str(sample_rate)
|
||||
if key_str not in self.resample_kernel:
|
||||
self.resample_kernel[key_str] = Resample(sample_rate, 16000, lowpass_filter_width=128)
|
||||
self.resample_kernel[key_str] = self.resample_kernel[key_str].to(self.device)
|
||||
audio_res = self.resample_kernel[key_str](audio)
|
||||
mel = self.mel_extractor(audio_res, center=True)
|
||||
hidden = self.mel2hidden(mel)
|
||||
f0 = self.decode(hidden, thred=thred, use_viterbi=use_viterbi).squeeze(0)
|
||||
return f0
|
||||
|
||||
def get_pitch(self, waveform, sample_rate, hop_size, length, interp_uv=False, fmin=50, fmax=1000):
|
||||
f0 = self.infer_from_audio(waveform, sample_rate=sample_rate)
|
||||
f0 = self.postprocess(f0, fmin, fmax)
|
||||
uv = f0 == 0
|
||||
time_step = hop_size / sample_rate
|
||||
f0_res = resample_align_curve(f0, 0.01, time_step, length)
|
||||
uv_res = resample_align_curve(uv.astype(np.float32), 0.01, time_step, length) > 0.5
|
||||
if not interp_uv:
|
||||
f0_res[uv_res] = 0
|
||||
return f0_res, uv_res
|
||||
|
||||
def infer_from_audio_batch(self, audios, sample_rate=16000, thred=0.03, use_viterbi=False):
|
||||
from ....utils.commons.dataset_utils import collate_1d_or_2d
|
||||
if isinstance(audios, list):
|
||||
audios = [torch.from_numpy(audio).float() for audio in audios]
|
||||
sizes = [math.ceil((audio.shape[0] + 1) / self.hop_length) for audio in audios]
|
||||
audios = collate_1d_or_2d(audios, 0.0).to(self.device)
|
||||
elif isinstance(audios, torch.Tensor):
|
||||
sizes = None
|
||||
if audios.device != self.device:
|
||||
audios = audios.to(self.device)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
if sample_rate == 16000:
|
||||
audios_res = audios
|
||||
else:
|
||||
key_str = str(sample_rate)
|
||||
if key_str not in self.resample_kernel:
|
||||
self.resample_kernel[key_str] = Resample(sample_rate, 16000, lowpass_filter_width=128)
|
||||
self.resample_kernel[key_str] = self.resample_kernel[key_str].to(self.device)
|
||||
audios_res = self.resample_kernel[key_str](audios)
|
||||
mels = self.mel_extractor(audios_res, center=True)
|
||||
hiddens = self.mel2hidden(mels)
|
||||
f0 = self.decode(hiddens, thred=thred, use_viterbi=use_viterbi)
|
||||
f0s = []
|
||||
for i in range(f0.shape[0]):
|
||||
f = f0[i, :sizes[i]] if sizes is not None else f0[i, :]
|
||||
f0s.append(f)
|
||||
return f0s
|
||||
|
||||
def get_pitch_batch(self, waveforms, sample_rate, hop_size, lengths, interp_uv=False, fmin=50, fmax=1000):
|
||||
# hop_size, sample_rate: tgt params
|
||||
f0s = self.infer_from_audio_batch(waveforms, sample_rate=sample_rate)
|
||||
f0s_res, uvs_res = [], []
|
||||
for idx, f0 in enumerate(f0s):
|
||||
f0 = self.postprocess(f0, fmin, fmax, min_gap=6)
|
||||
uv = f0 == 0
|
||||
length = lengths[idx]
|
||||
time_step = hop_size / sample_rate
|
||||
f0_res = resample_align_curve(f0, 0.01, time_step, length)
|
||||
uv_res = resample_align_curve(uv.astype(np.float32), 0.01, time_step, length) > 0.5
|
||||
if not interp_uv:
|
||||
f0_res[uv_res] = 0
|
||||
f0s_res.append(f0_res)
|
||||
uvs_res.append(uv_res)
|
||||
return f0s_res, uvs_res
|
||||
|
||||
def release_cuda(self):
|
||||
self.model = self.model.cpu()
|
||||
self.mel_extractor = self.mel_extractor.cpu()
|
||||
@@ -0,0 +1,32 @@
|
||||
from torch import nn
|
||||
|
||||
from .constants import *
|
||||
from .deepunet import DeepUnet0
|
||||
from .seq import BiGRU
|
||||
|
||||
|
||||
class E2E0(nn.Module):
|
||||
def __init__(self, n_blocks, n_gru, kernel_size, en_de_layers=5, inter_layers=4, in_channels=1,
|
||||
en_out_channels=16):
|
||||
super(E2E0, self).__init__()
|
||||
self.unet = DeepUnet0(kernel_size, n_blocks, en_de_layers, inter_layers, in_channels, en_out_channels)
|
||||
self.cnn = nn.Conv2d(en_out_channels, 3, (3, 3), padding=(1, 1))
|
||||
if n_gru:
|
||||
self.fc = nn.Sequential(
|
||||
BiGRU(3 * N_MELS, 256, n_gru),
|
||||
nn.Linear(512, N_CLASS),
|
||||
nn.Dropout(0.25),
|
||||
nn.Sigmoid()
|
||||
)
|
||||
else:
|
||||
self.fc = nn.Sequential(
|
||||
nn.Linear(3 * N_MELS, N_CLASS),
|
||||
nn.Dropout(0.25),
|
||||
nn.Sigmoid()
|
||||
)
|
||||
|
||||
def forward(self, mel):
|
||||
mel = mel.transpose(-1, -2).unsqueeze(1)
|
||||
x = self.cnn(self.unet(mel)).transpose(1, 2).flatten(-2)
|
||||
x = self.fc(x)
|
||||
return x
|
||||
@@ -0,0 +1,10 @@
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
class BiGRU(nn.Module):
|
||||
def __init__(self, input_features, hidden_features, num_layers):
|
||||
super(BiGRU, self).__init__()
|
||||
self.gru = nn.GRU(input_features, hidden_features, num_layers=num_layers, batch_first=True, bidirectional=True)
|
||||
|
||||
def forward(self, x):
|
||||
return self.gru(x)[0]
|
||||
@@ -0,0 +1,72 @@
|
||||
import torch
|
||||
import numpy as np
|
||||
import torch.nn.functional as F
|
||||
from librosa.filters import mel
|
||||
|
||||
|
||||
class MelSpectrogram(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
n_mel_channels,
|
||||
sampling_rate,
|
||||
win_length,
|
||||
hop_length,
|
||||
n_fft=None,
|
||||
mel_fmin=0,
|
||||
mel_fmax=None,
|
||||
clamp=1e-5
|
||||
):
|
||||
super().__init__()
|
||||
n_fft = win_length if n_fft is None else n_fft
|
||||
self.hann_window = {}
|
||||
mel_basis = mel(
|
||||
sr=sampling_rate,
|
||||
n_fft=n_fft,
|
||||
n_mels=n_mel_channels,
|
||||
fmin=mel_fmin,
|
||||
fmax=mel_fmax,
|
||||
htk=True)
|
||||
mel_basis = torch.from_numpy(mel_basis).float()
|
||||
self.register_buffer("mel_basis", mel_basis)
|
||||
self.n_fft = win_length if n_fft is None else n_fft
|
||||
self.hop_length = hop_length
|
||||
self.win_length = win_length
|
||||
self.sampling_rate = sampling_rate
|
||||
self.n_mel_channels = n_mel_channels
|
||||
self.clamp = clamp
|
||||
|
||||
def forward(self, audio, keyshift=0, speed=1, center=True):
|
||||
factor = 2 ** (keyshift / 12)
|
||||
n_fft_new = int(np.round(self.n_fft * factor))
|
||||
win_length_new = int(np.round(self.win_length * factor))
|
||||
hop_length_new = int(np.round(self.hop_length * speed))
|
||||
|
||||
keyshift_key = str(keyshift) + '_' + str(audio.device)
|
||||
if keyshift_key not in self.hann_window:
|
||||
self.hann_window[keyshift_key] = torch.hann_window(win_length_new).to(audio.device)
|
||||
if center:
|
||||
pad_left = win_length_new // 2
|
||||
pad_right = (win_length_new + 1) // 2
|
||||
audio = F.pad(audio, (pad_left, pad_right))
|
||||
|
||||
fft = torch.stft(
|
||||
audio,
|
||||
n_fft=n_fft_new,
|
||||
hop_length=hop_length_new,
|
||||
win_length=win_length_new,
|
||||
window=self.hann_window[keyshift_key],
|
||||
center=False,
|
||||
return_complex=True
|
||||
)
|
||||
magnitude = fft.abs()
|
||||
|
||||
if keyshift != 0:
|
||||
size = self.n_fft // 2 + 1
|
||||
resize = magnitude.size(1)
|
||||
if resize < size:
|
||||
magnitude = F.pad(magnitude, (0, 0, 0, size - resize))
|
||||
magnitude = magnitude[:, :size, :] * self.win_length / win_length_new
|
||||
|
||||
mel_output = torch.matmul(self.mel_basis, magnitude)
|
||||
log_mel_spec = torch.log(torch.clamp(mel_output, min=self.clamp))
|
||||
return log_mel_spec
|
||||
@@ -0,0 +1,43 @@
|
||||
import librosa
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from .constants import *
|
||||
|
||||
|
||||
def to_local_average_f0(hidden, center=None, thred=0.03):
|
||||
idx = torch.arange(N_CLASS, device=hidden.device)[None, None, :] # [B=1, T=1, N]
|
||||
idx_cents = idx * 20 + CONST # [B=1, N]
|
||||
if center is None:
|
||||
center = torch.argmax(hidden, dim=2, keepdim=True) # [B, T, 1]
|
||||
start = torch.clip(center - 4, min=0) # [B, T, 1]
|
||||
end = torch.clip(center + 5, max=N_CLASS) # [B, T, 1]
|
||||
idx_mask = (idx >= start) & (idx < end) # [B, T, N]
|
||||
weights = hidden * idx_mask # [B, T, N]
|
||||
product_sum = torch.sum(weights * idx_cents, dim=2) # [B, T]
|
||||
weight_sum = torch.sum(weights, dim=2) # [B, T]
|
||||
cents = product_sum / (weight_sum + (weight_sum == 0)) # avoid dividing by zero, [B, T]
|
||||
f0 = 10 * 2 ** (cents / 1200)
|
||||
uv = hidden.max(dim=2)[0] < thred # [B, T]
|
||||
f0 = f0 * ~uv
|
||||
return f0.cpu().numpy()
|
||||
|
||||
|
||||
def to_viterbi_f0(hidden, thred=0.03):
|
||||
# Create viterbi transition matrix
|
||||
if not hasattr(to_viterbi_f0, 'transition'):
|
||||
xx, yy = np.meshgrid(range(N_CLASS), range(N_CLASS))
|
||||
transition = np.maximum(30 - abs(xx - yy), 0)
|
||||
transition = transition / transition.sum(axis=1, keepdims=True)
|
||||
to_viterbi_f0.transition = transition
|
||||
|
||||
# Convert to probability
|
||||
prob = hidden.squeeze(0).cpu().numpy()
|
||||
prob = prob.T
|
||||
prob = prob / prob.sum(axis=0)
|
||||
|
||||
# Perform viterbi decoding
|
||||
path = librosa.sequence.viterbi(prob, to_viterbi_f0.transition).astype(np.int64)
|
||||
center = torch.from_numpy(path).unsqueeze(0).unsqueeze(-1).to(hidden.device)
|
||||
|
||||
return to_local_average_f0(hidden, center=center, thred=thred)
|
||||
@@ -0,0 +1 @@
|
||||
"""Core ROSVOT model components."""
|
||||
@@ -0,0 +1,295 @@
|
||||
from copy import deepcopy
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import numpy as np
|
||||
|
||||
from ...utils.commons.hparams import hparams
|
||||
from ...utils.commons.gpu_mem_track import MemTracker
|
||||
from ..commons.layers import Embedding
|
||||
from ..commons.conv import ResidualBlock, ConvBlocks
|
||||
from ..commons.conformer.conformer import ConformerLayers
|
||||
from .unet import Unet
|
||||
|
||||
def regulate_boundary(bd_logits, threshold, min_gap=18, ref_bd=None, ref_bd_min_gap=8, non_padding=None):
|
||||
# this doesn't preserve gradient
|
||||
device = bd_logits.device
|
||||
bd_logits = torch.sigmoid(bd_logits).data.cpu()
|
||||
# bd_logits[0] = bd_logits[-1] = 1e-5 # avoid itv invalid problem
|
||||
bd = (bd_logits > threshold).long()
|
||||
bd_res = torch.zeros_like(bd).long()
|
||||
for i in range(bd.shape[0]):
|
||||
bd_i = bd[i]
|
||||
last_bd_idx = -1
|
||||
start = -1
|
||||
for j in range(bd_i.shape[0]):
|
||||
if bd_i[j] == 1:
|
||||
if 0 <= start < j:
|
||||
continue
|
||||
elif start < 0:
|
||||
start = j
|
||||
else:
|
||||
if 0 <= start < j:
|
||||
if j - 1 > start:
|
||||
bd_idx = start + int(torch.argmax(bd_logits[i, start: j]).item())
|
||||
else:
|
||||
bd_idx = start
|
||||
if bd_idx - last_bd_idx < min_gap and last_bd_idx > 0:
|
||||
bd_idx = round((bd_idx + last_bd_idx) / 2)
|
||||
bd_res[i, last_bd_idx] = 0
|
||||
bd_res[i, bd_idx] = 1
|
||||
last_bd_idx = bd_idx
|
||||
start = -1
|
||||
|
||||
# assert ref_bd_min_gap <= min_gap // 2
|
||||
if ref_bd is not None and ref_bd_min_gap > 0:
|
||||
ref = ref_bd.data.cpu()
|
||||
for i in range(bd_res.shape[0]):
|
||||
ref_bd_i = ref[i]
|
||||
ref_bd_i_js = []
|
||||
for j in range(ref_bd_i.shape[0]):
|
||||
if ref_bd_i[j] == 1:
|
||||
ref_bd_i_js.append(j)
|
||||
seg_sum = torch.sum(bd_res[i, max(0, j - ref_bd_min_gap): j + ref_bd_min_gap])
|
||||
if seg_sum == 0:
|
||||
bd_res[i, j] = 1
|
||||
elif seg_sum == 1 and bd_res[i, j] != 1:
|
||||
bd_res[i, max(0, j - ref_bd_min_gap): j + ref_bd_min_gap] = \
|
||||
ref_bd_i[max(0, j - ref_bd_min_gap): j + ref_bd_min_gap]
|
||||
elif seg_sum > 1:
|
||||
for k in range(1, ref_bd_min_gap+1):
|
||||
if bd_res[i, max(0, j - k)] == 1 and ref_bd_i[max(0, j - k)] != 1:
|
||||
bd_res[i, max(0, j - k)] = 0
|
||||
break
|
||||
if bd_res[i, min(bd_res.shape[1] - 1, j + k)] == 1 and ref_bd_i[min(bd_res.shape[1] - 1, j + k)] != 1:
|
||||
bd_res[i, min(bd_res.shape[1] - 1, j + k)] = 0
|
||||
break
|
||||
bd_res[i, j] = 1
|
||||
# final check
|
||||
assert torch.sum(bd_res[i, ref_bd_i_js]) == len(ref_bd_i_js), \
|
||||
f"{torch.sum(bd_res[i, ref_bd_i_js])} {len(ref_bd_i_js)}"
|
||||
|
||||
bd_res = bd_res.to(device)
|
||||
|
||||
# force valid begin and end
|
||||
bd_res[:, 0] = 0
|
||||
if non_padding is not None:
|
||||
for i in range(bd_res.shape[0]):
|
||||
bd_res[i, sum(non_padding[i]) - 1:] = 0
|
||||
else:
|
||||
bd_res[:, -1] = 0
|
||||
|
||||
return bd_res
|
||||
|
||||
class BackboneNet(nn.Module):
|
||||
def __init__(self, hparams):
|
||||
super().__init__()
|
||||
self.hidden_size = hidden_size = hparams['hidden_size']
|
||||
self.dropout = hparams.get('dropout', 0.0)
|
||||
updown_rates = [2, 2, 2]
|
||||
channel_multiples = [1, 1, 1]
|
||||
if hparams.get('updown_rates', None) is not None:
|
||||
updown_rates = [int(i) for i in hparams.get('updown_rates', None).split('-')]
|
||||
if hparams.get('channel_multiples', None) is not None:
|
||||
channel_multiples = [float(i) for i in hparams.get('channel_multiples', None).split('-')]
|
||||
assert len(updown_rates) == len(channel_multiples)
|
||||
# convs
|
||||
if hparams.get('bkb_net', 'conv') == 'conv':
|
||||
self.net = Unet(hidden_size, down_layers=len(updown_rates), mid_layers=hparams.get('bkb_layers', 12),
|
||||
up_layers=len(updown_rates), kernel_size=3, updown_rates=updown_rates,
|
||||
channel_multiples=channel_multiples, dropout=0, is_BTC=True,
|
||||
constant_channels=False, mid_net=None, use_skip_layer=hparams.get('unet_skip_layer', False))
|
||||
# conformer
|
||||
elif hparams.get('bkb_net', 'conv') == 'conformer':
|
||||
mid_net = ConformerLayers(
|
||||
hidden_size, num_layers=hparams.get('bkb_layers', 12), kernel_size=hparams.get('conformer_kernel', 9),
|
||||
dropout=self.dropout, num_heads=4)
|
||||
self.net = Unet(hidden_size, down_layers=len(updown_rates), up_layers=len(updown_rates), kernel_size=3,
|
||||
updown_rates=updown_rates, channel_multiples=channel_multiples, dropout=0,
|
||||
is_BTC=True, constant_channels=False, mid_net=mid_net,
|
||||
use_skip_layer=hparams.get('unet_skip_layer', False))
|
||||
|
||||
def forward(self, x):
|
||||
return self.net(x)
|
||||
|
||||
class PitchDecoder(nn.Module):
|
||||
def __init__(self, hparams):
|
||||
super().__init__()
|
||||
self.hidden_size = hidden_size = hparams['hidden_size']
|
||||
self.dropout = hparams.get('dropout', 0.0)
|
||||
self.note_bd_out = nn.Linear(hidden_size, 1)
|
||||
self.note_bd_temperature = max(1e-7, hparams.get('note_bd_temperature', 1.0))
|
||||
|
||||
# note prediction
|
||||
self.pitch_attn_num_head = hparams.get('pitch_attn_num_head', 1)
|
||||
self.multihead_dot_attn = nn.Linear(hidden_size, self.pitch_attn_num_head)
|
||||
self.post = ConvBlocks(hidden_size, out_dims=hidden_size, dilations=None, kernel_size=3,
|
||||
layers_in_block=1, c_multiple=1, dropout=self.dropout, num_layers=1,
|
||||
post_net_kernel=3, act_type='leakyrelu')
|
||||
self.pitch_out = nn.Linear(hidden_size, hparams.get('note_num', 100) + 4)
|
||||
self.note_num = hparams.get('note_num', 100)
|
||||
self.note_start = hparams.get('note_start', 30)
|
||||
self.pitch_temperature = max(1e-7, hparams.get('note_pitch_temperature', 1.0))
|
||||
|
||||
def forward(self, feat, note_bd, train=True):
|
||||
bsz, T, _ = feat.shape
|
||||
|
||||
attn = torch.sigmoid(self.multihead_dot_attn(feat)) # [B, T, C] -> [B, T, num_head]
|
||||
attn = F.dropout(attn, self.dropout, train)
|
||||
attn_feat = feat.unsqueeze(3) * attn.unsqueeze(2) # [B, T, C, 1] x [B, T, 1, num_head] -> [B, T, C, num_head]
|
||||
attn_feat = torch.mean(attn_feat, dim=-1) # [B, T, C, num_head] -> [B, T, C]
|
||||
mel2note = torch.cumsum(note_bd, 1)
|
||||
note_length = torch.max(torch.sum(note_bd, dim=1)).item() + 1 # max length
|
||||
note_lengths = torch.sum(note_bd, dim=1) + 1 # [B]
|
||||
# print('note_length', note_length)
|
||||
|
||||
attn = torch.mean(attn, dim=-1, keepdim=True) # [B, T, num_head] -> [B, T, 1]
|
||||
denom = mel2note.new_zeros(bsz, note_length, dtype=attn.dtype).scatter_add_(
|
||||
dim=1, index=mel2note, src=attn.squeeze(-1)
|
||||
) # [B, T] -> [B, note_length] count the note frames of each note (with padding excluded)
|
||||
frame2note = mel2note.unsqueeze(-1).repeat(1, 1, self.hidden_size) # [B, T] -> [B, T, C], with padding included
|
||||
note_aggregate = frame2note.new_zeros(bsz, note_length, self.hidden_size, dtype=attn_feat.dtype).scatter_add_(
|
||||
dim=1, index=frame2note, src=attn_feat
|
||||
) # [B, T, C] -> [B, note_length, C]
|
||||
note_aggregate = note_aggregate / (denom.unsqueeze(-1) + 1e-5)
|
||||
note_aggregate = F.dropout(note_aggregate, self.dropout, train)
|
||||
note_logits = self.post(note_aggregate)
|
||||
note_logits = self.pitch_out(note_logits) / self.pitch_temperature
|
||||
# note_logits = torch.clamp(note_logits, min=-16., max=16.) # don't know need it or not
|
||||
|
||||
note_pred = torch.softmax(note_logits, dim=-1) # [B, note_length, note_num]
|
||||
note_pred = torch.argmax(note_pred, dim=-1) # [B, note_length]
|
||||
# for some reason, note idx maybe 130 (why?)
|
||||
note_pred[note_pred > self.note_num] = 0
|
||||
note_pred[note_pred < self.note_start] = 0
|
||||
|
||||
return note_lengths, note_logits, note_pred
|
||||
|
||||
class MidiExtractor(nn.Module):
|
||||
def __init__(self, hparams):
|
||||
super(MidiExtractor, self).__init__()
|
||||
self.hparams = deepcopy(hparams)
|
||||
self.hidden_size = hidden_size = hparams['hidden_size']
|
||||
self.dropout = hparams.get('dropout', 0.0)
|
||||
self.note_bd_threshold = hparams.get('note_bd_threshold', 0.5)
|
||||
self.note_bd_min_gap = round(hparams.get('note_bd_min_gap', 100) * hparams['audio_sample_rate'] / 1000 / hparams['hop_size'])
|
||||
self.note_bd_ref_min_gap = round(hparams.get('note_bd_ref_min_gap', 50) * hparams['audio_sample_rate'] / 1000 / hparams['hop_size'])
|
||||
|
||||
self.mel_proj = nn.Conv1d(hparams['use_mel_bins'], hidden_size, kernel_size=3, padding=1)
|
||||
self.mel_encoder = ConvBlocks(hidden_size, out_dims=hidden_size, dilations=None, kernel_size=3,
|
||||
layers_in_block=2, c_multiple=1, dropout=self.dropout, num_layers=1,
|
||||
post_net_kernel=3, act_type='leakyrelu')
|
||||
self.use_pitch = hparams.get('use_pitch_embed', True)
|
||||
if self.use_pitch:
|
||||
self.pitch_embed = Embedding(300, hidden_size, 0, 'kaiming')
|
||||
self.uv_embed = Embedding(3, hidden_size, 0, 'kaiming')
|
||||
self.use_wbd = hparams.get('use_wbd', True)
|
||||
if self.use_wbd:
|
||||
self.word_bd_embed = Embedding(3, hidden_size, 0, 'kaiming')
|
||||
self.cond_encoder = ConvBlocks(hidden_size, out_dims=hidden_size, dilations=None, kernel_size=3,
|
||||
layers_in_block=1, c_multiple=1, dropout=self.dropout, num_layers=1,
|
||||
post_net_kernel=3, act_type='leakyrelu')
|
||||
|
||||
# backbone
|
||||
self.net = BackboneNet(hparams)
|
||||
|
||||
# note bd prediction
|
||||
self.note_bd_out = nn.Linear(hidden_size, 1)
|
||||
self.note_bd_temperature = max(1e-7, hparams.get('note_bd_temperature', 1.0))
|
||||
|
||||
# note prediction
|
||||
self.pitch_decoder = PitchDecoder(hparams)
|
||||
|
||||
self.reset_parameters()
|
||||
|
||||
def run_encoder(self, mel=None, word_bd=None, pitch=None, uv=None, non_padding=None):
|
||||
mel_embed = self.mel_proj(mel.transpose(1, 2)).transpose(1, 2)
|
||||
mel_embed = self.mel_encoder(mel_embed)
|
||||
pitch_embed = word_bd_embed = 0
|
||||
if self.use_pitch and pitch is not None and uv is not None:
|
||||
pitch_embed = self.pitch_embed(pitch) + self.uv_embed(uv) # [B, T, C]
|
||||
if self.use_wbd and word_bd is not None:
|
||||
word_bd_embed = self.word_bd_embed(word_bd)
|
||||
feat = self.cond_encoder(mel_embed + pitch_embed + word_bd_embed)
|
||||
|
||||
return feat
|
||||
|
||||
def forward(self, mel=None, word_bd=None, note_bd=None, pitch=None, uv=None, non_padding=None, train=True):
|
||||
ret = {}
|
||||
bsz, T, _ = mel.shape
|
||||
|
||||
feat = self.run_encoder(mel, word_bd, pitch, uv, non_padding)
|
||||
feat = self.net(feat) # [B, T, C]
|
||||
|
||||
# note bd prediction
|
||||
note_bd_logits = self.note_bd_out(F.dropout(feat, self.dropout, train)).squeeze(-1) / self.note_bd_temperature
|
||||
note_bd_logits = torch.clamp(note_bd_logits, min=-16., max=16.)
|
||||
ret['note_bd_logits'] = note_bd_logits # [B, T]
|
||||
if note_bd is None or not train:
|
||||
note_bd = regulate_boundary(note_bd_logits, self.note_bd_threshold, self.note_bd_min_gap,
|
||||
word_bd, self.note_bd_ref_min_gap, non_padding)
|
||||
ret['note_bd_pred'] = note_bd # [B, T]
|
||||
|
||||
# note pitch prediction
|
||||
note_lengths, note_logits, note_pred = self.pitch_decoder(feat, note_bd, train)
|
||||
ret['note_lengths'], ret['note_logits'], ret['note_pred'] = note_lengths, note_logits, note_pred
|
||||
|
||||
return ret
|
||||
|
||||
def reset_parameters(self):
|
||||
nn.init.kaiming_normal_(self.pitch_decoder.multihead_dot_attn.weight, mode='fan_in')
|
||||
nn.init.kaiming_normal_(self.note_bd_out.weight, mode='fan_in')
|
||||
nn.init.kaiming_normal_(self.pitch_decoder.pitch_out.weight, mode='fan_in')
|
||||
nn.init.kaiming_normal_(self.mel_proj.weight, mode='fan_in')
|
||||
nn.init.constant_(self.pitch_decoder.multihead_dot_attn.bias, 0.0)
|
||||
nn.init.constant_(self.note_bd_out.bias, 0.0)
|
||||
nn.init.constant_(self.pitch_decoder.pitch_out.bias, 0.0)
|
||||
|
||||
|
||||
class WordbdExtractor(MidiExtractor):
|
||||
def __init__(self, hparams):
|
||||
super().__init__(hparams)
|
||||
self.use_wbd = False
|
||||
self.word_bd_embed = None
|
||||
self.note_bd_out = self.note_bd_temperature = self.pitch_decoder = None
|
||||
|
||||
self.word_bd_threshold = hparams.get('word_bd_threshold', 0.5)
|
||||
self.word_bd_min_gap = round(
|
||||
hparams.get('word_bd_min_gap', 100) * hparams['audio_sample_rate'] / 1000 / hparams['hop_size'])
|
||||
|
||||
self.word_bd_out = nn.Linear(self.hidden_size, 1)
|
||||
self.word_bd_temperature = max(1e-7, hparams.get('word_bd_temperature', 1.0))
|
||||
nn.init.kaiming_normal_(self.word_bd_out.weight, mode='fan_in')
|
||||
nn.init.constant_(self.word_bd_out.bias, 0.0)
|
||||
|
||||
def forward(self, mel=None, pitch=None, uv=None, non_padding=None, train=True):
|
||||
# gpu_tracker.track()
|
||||
ret = {}
|
||||
bsz, T, _ = mel.shape
|
||||
|
||||
feat = self.run_encoder(mel=mel, pitch=pitch, uv=uv, non_padding=non_padding)
|
||||
feat = self.net(feat) # [B, T, C]
|
||||
|
||||
word_bd_logits = self.word_bd_out(F.dropout(feat, self.dropout, train)).squeeze(-1) / self.word_bd_temperature
|
||||
word_bd_logits = torch.clamp(word_bd_logits, min=-16., max=16.)
|
||||
ret['word_bd_logits'] = word_bd_logits # [B, T]
|
||||
|
||||
if not train:
|
||||
word_bd = regulate_boundary(word_bd_logits, self.word_bd_threshold, self.word_bd_min_gap,
|
||||
non_padding=non_padding)
|
||||
ret['word_bd_pred'] = word_bd # [B, T]
|
||||
|
||||
return ret
|
||||
|
||||
def reset_parameters(self):
|
||||
if self.use_pitch:
|
||||
nn.init.kaiming_normal_(self.pitch_embed.weight, mode='fan_in')
|
||||
nn.init.kaiming_normal_(self.uv_embed.weight, mode='fan_in')
|
||||
nn.init.kaiming_normal_(self.mel_proj.weight, mode='fan_in')
|
||||
if self.use_pitch:
|
||||
nn.init.constant_(self.pitch_embed.weight[self.pitch_embed.padding_idx], 0.0)
|
||||
nn.init.constant_(self.uv_embed.weight[self.uv_embed.padding_idx], 0.0)
|
||||
|
||||
|
||||
@@ -0,0 +1,172 @@
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from ..commons.layers import LayerNorm, Embedding
|
||||
from ..commons.conv import ConvBlocks, ResidualBlock, get_norm_builder, get_act_builder
|
||||
|
||||
class UnetDown(nn.Module):
|
||||
def __init__(self, hidden_size, n_layers, kernel_size, down_rates, channel_multiples=None, dropout=0.0,
|
||||
is_BTC=True, constant_channels=False):
|
||||
super(UnetDown, self).__init__()
|
||||
assert n_layers == len(down_rates) # downs, down sample rate
|
||||
down_rates = [int(i) for i in down_rates]
|
||||
self.n_layers = n_layers
|
||||
self.hidden_size = hidden_size
|
||||
self.is_BTC = is_BTC
|
||||
channel_multiples = channel_multiples if channel_multiples is not None else down_rates
|
||||
self.layers = nn.ModuleList()
|
||||
self.downs = nn.ModuleList()
|
||||
in_channels = hidden_size
|
||||
for i in range(self.n_layers):
|
||||
out_channels = int(in_channels * channel_multiples[i]) if not constant_channels else in_channels
|
||||
self.layers.append(nn.Sequential(
|
||||
ResidualBlock(in_channels, kernel_size, dilation=1, n=1, norm_type='ln', dropout=dropout,
|
||||
c_multiple=1, ln_eps=1e-5, act_type='leakyrelu'),
|
||||
nn.Conv1d(in_channels, out_channels, kernel_size, padding=(kernel_size - 1) // 2),
|
||||
ResidualBlock(out_channels, kernel_size, dilation=1, n=1, norm_type='ln',
|
||||
dropout=dropout, c_multiple=1, ln_eps=1e-5, act_type='leakyrelu')
|
||||
))
|
||||
self.downs.append(nn.Sequential(
|
||||
nn.AvgPool1d(down_rates[i])
|
||||
))
|
||||
in_channels = out_channels
|
||||
self.last_norm = get_norm_builder('ln', out_channels)()
|
||||
self.post_net = nn.Conv1d(out_channels, out_channels, kernel_size=kernel_size,
|
||||
padding=kernel_size // 2)
|
||||
|
||||
def forward(self, x, **kwargs):
|
||||
# x [B, T, C]
|
||||
if self.is_BTC:
|
||||
x = x.transpose(1, 2) # [B, C, T]
|
||||
skip_xs = []
|
||||
for i in range(self.n_layers):
|
||||
skip_x = self.layers[i](x)
|
||||
x = self.downs[i](skip_x)
|
||||
if self.is_BTC:
|
||||
skip_xs.append(skip_x.transpose(1, 2)) # [B, T, C]
|
||||
else:
|
||||
skip_xs.append(skip_x)
|
||||
x = self.post_net(self.last_norm(x))
|
||||
if self.is_BTC:
|
||||
x = x.transpose(1, 2)
|
||||
return x, skip_xs
|
||||
|
||||
class UnetMid(nn.Module):
|
||||
def __init__(self, hidden_size, kernel_size, n_layers=None, in_dims=None, out_dims=None,
|
||||
dropout=0.0, is_BTC=True, net=None):
|
||||
super(UnetMid, self).__init__()
|
||||
in_dims = in_dims if in_dims is not None else hidden_size
|
||||
out_dims = out_dims if out_dims is not None else hidden_size
|
||||
self.pre = nn.Conv1d(in_dims, hidden_size, kernel_size, padding=kernel_size // 2)
|
||||
self.post = nn.Conv1d(hidden_size, out_dims, kernel_size, padding=kernel_size // 2)
|
||||
self.is_BTC = is_BTC
|
||||
if net is not None:
|
||||
self.net = net
|
||||
else:
|
||||
self.net = ConvBlocks(hidden_size, out_dims=hidden_size, dilations=None, kernel_size=kernel_size,
|
||||
layers_in_block=2, c_multiple=2, dropout=dropout, num_layers=n_layers,
|
||||
post_net_kernel=3, act_type='leakyrelu', is_BTC=is_BTC)
|
||||
|
||||
def forward(self, x, cond=None, **kwargs):
|
||||
# x [B, T, C]
|
||||
if self.is_BTC:
|
||||
x = self.pre(x.transpose(1, 2)).transpose(1, 2)
|
||||
else:
|
||||
x = self.pre(x)
|
||||
if cond is None:
|
||||
cond = 0
|
||||
x = self.net(x + cond)
|
||||
if self.is_BTC:
|
||||
x = self.post(x.transpose(1, 2)).transpose(1, 2)
|
||||
else:
|
||||
x = self.post(x)
|
||||
return x
|
||||
|
||||
class UnetUp(nn.Module):
|
||||
def __init__(self, hidden_size, n_layers, kernel_size, up_rates, channel_multiples=None, dropout=0.0,
|
||||
is_BTC=True, constant_channels=False, use_skip_layer=False, skip_scale=1.0):
|
||||
super(UnetUp, self).__init__()
|
||||
assert n_layers == len(up_rates) # this is reversed in up module, from the output to the interface with middle
|
||||
up_rates = [int(i) for i in up_rates]
|
||||
self.n_layers = n_layers
|
||||
self.hidden_size = hidden_size
|
||||
self.is_BTC = is_BTC
|
||||
self.skip_scale = skip_scale
|
||||
channel_multiples = channel_multiples if channel_multiples is not None else up_rates
|
||||
# in_channels = int(np.cumprod(channel_multiples)[-1] * hidden_size) if not constant_channels else hidden_size
|
||||
self.in_channels_lst = (np.cumprod([1] + channel_multiples) * hidden_size).astype(int) if not constant_channels \
|
||||
else [hidden_size for _ in range(self.n_layers + 1)]
|
||||
in_channels = self.in_channels_lst[-1]
|
||||
self.ups = nn.ModuleList()
|
||||
self.skip_layers = nn.ModuleList()
|
||||
self.layers = nn.ModuleList()
|
||||
for i in range(self.n_layers-1, -1, -1):
|
||||
out_channels = self.in_channels_lst[i] if not constant_channels else in_channels
|
||||
self.ups.append(nn.Sequential(
|
||||
nn.ConvTranspose1d(in_channels, in_channels, kernel_size=kernel_size, stride=up_rates[i],
|
||||
padding=kernel_size//2, output_padding=up_rates[i]-1),
|
||||
get_norm_builder('ln', in_channels)(),
|
||||
get_act_builder('leakyrelu')()
|
||||
))
|
||||
self.layers.append(nn.Sequential(
|
||||
# ResidualBlock(in_channels*2, kernel_size, dilation=1, n=1, norm_type='ln', dropout=dropout,
|
||||
# c_multiple=1, ln_eps=1e-5, act_type='leakyrelu'),
|
||||
nn.Conv1d(in_channels*2, out_channels, kernel_size, padding=(kernel_size - 1) // 2),
|
||||
ResidualBlock(out_channels, kernel_size, dilation=1, n=1, norm_type='ln',
|
||||
dropout=dropout, c_multiple=1, ln_eps=1e-5, act_type='leakyrelu')
|
||||
))
|
||||
if use_skip_layer:
|
||||
self.skip_layers.append(
|
||||
ResidualBlock(in_channels, kernel_size, dilation=1, n=1, norm_type='ln', dropout=dropout,
|
||||
c_multiple=1, ln_eps=1e-5, act_type='leakyrelu')
|
||||
)
|
||||
else:
|
||||
self.skip_layers.append(nn.Identity())
|
||||
|
||||
in_channels = out_channels
|
||||
self.out_channels = out_channels
|
||||
self.last_norm = get_norm_builder('ln', out_channels)()
|
||||
self.post_net = nn.Conv1d(out_channels, out_channels, kernel_size=kernel_size,
|
||||
padding=kernel_size // 2)
|
||||
|
||||
def forward(self, x, skips, **kwargs):
|
||||
# x [B, T, C]
|
||||
if self.is_BTC:
|
||||
x = x.transpose(1, 2) # [B, C, T]
|
||||
for i in range(self.n_layers):
|
||||
x = self.ups[i](x)
|
||||
skip_x = skips[self.n_layers - i - 1] if not self.is_BTC \
|
||||
else skips[self.n_layers - i - 1].transpose(1, 2) # [B, T, C] -> [B, C, T]
|
||||
skip_x = self.skip_layers[i](skip_x) * self.skip_scale
|
||||
x = torch.cat((x, skip_x), dim=1) # [B, C, T]
|
||||
x = self.layers[i](x)
|
||||
x = self.post_net(self.last_norm(x))
|
||||
if self.is_BTC:
|
||||
x = x.transpose(1, 2)
|
||||
return x
|
||||
|
||||
class Unet(nn.Module):
|
||||
def __init__(self, hidden_size, down_layers, up_layers, kernel_size,
|
||||
updown_rates, mid_layers=None, channel_multiples=None, dropout=0.0,
|
||||
is_BTC=True, constant_channels=False, mid_net=None, use_skip_layer=False, skip_scale=1.0):
|
||||
super(Unet, self).__init__()
|
||||
assert len(updown_rates) == down_layers == up_layers, f"{len(updown_rates)}, {down_layers}, {up_layers}"
|
||||
if channel_multiples is not None:
|
||||
assert len(channel_multiples) == len(updown_rates)
|
||||
else:
|
||||
channel_multiples = updown_rates
|
||||
self.down = UnetDown(hidden_size, down_layers, kernel_size, updown_rates,
|
||||
channel_multiples, dropout, is_BTC, constant_channels)
|
||||
down_out_dims = int(np.cumprod(channel_multiples)[-1] * hidden_size) if not constant_channels else hidden_size
|
||||
self.mid = UnetMid(hidden_size, kernel_size, mid_layers,
|
||||
in_dims=down_out_dims, out_dims=down_out_dims, dropout=dropout, is_BTC=is_BTC, net=mid_net)
|
||||
self.up = UnetUp(hidden_size, up_layers, kernel_size, updown_rates,
|
||||
channel_multiples, dropout, is_BTC, constant_channels, use_skip_layer, skip_scale)
|
||||
|
||||
def forward(self, x, mid_cond=None, **kwargs):
|
||||
x, skips = self.down(x)
|
||||
x = self.mid(x, mid_cond)
|
||||
x = self.up(x, skips)
|
||||
return x
|
||||
@@ -0,0 +1,15 @@
|
||||
|
||||
def seed_everything(seed: int, seed_cudnn=False):
|
||||
import random, os
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
random.seed(seed)
|
||||
os.environ['PYTHONHASHSEED'] = str(seed)
|
||||
np.random.seed(seed)
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed(seed)
|
||||
if seed_cudnn:
|
||||
torch.backends.cudnn.deterministic = True
|
||||
torch.backends.cudnn.benchmark = True
|
||||
|
||||
@@ -0,0 +1,100 @@
|
||||
import librosa
|
||||
import numpy as np
|
||||
import wave
|
||||
import soundfile as sf
|
||||
|
||||
|
||||
def librosa_pad_lr(x, fsize, fshift, pad_sides=1):
|
||||
'''compute right padding (final frame) or both sides padding (first and final frames)
|
||||
'''
|
||||
assert pad_sides in (1, 2)
|
||||
# return int(fsize // 2)
|
||||
pad = (x.shape[0] // fshift + 1) * fshift - x.shape[0]
|
||||
if pad_sides == 1:
|
||||
return 0, pad
|
||||
else:
|
||||
return pad // 2, pad // 2 + pad % 2
|
||||
|
||||
|
||||
def amp_to_db(x):
|
||||
return 20 * np.log10(np.maximum(1e-5, x))
|
||||
|
||||
|
||||
def db_to_amp(x):
|
||||
return 10.0 ** (x * 0.05)
|
||||
|
||||
|
||||
def normalize(S, min_level_db):
|
||||
return (S - min_level_db) / -min_level_db
|
||||
|
||||
|
||||
def denormalize(D, min_level_db):
|
||||
return (D * -min_level_db) + min_level_db
|
||||
|
||||
|
||||
def librosa_wav2spec(wav_path,
|
||||
fft_size=1024,
|
||||
hop_size=256,
|
||||
win_length=1024,
|
||||
window="hann",
|
||||
num_mels=80,
|
||||
fmin=80,
|
||||
fmax=-1,
|
||||
eps=1e-6,
|
||||
sample_rate=22050,
|
||||
loud_norm=False,
|
||||
trim_long_sil=False):
|
||||
import pyloudnorm as pyln
|
||||
if isinstance(wav_path, str):
|
||||
if trim_long_sil:
|
||||
from .vad import trim_long_silences
|
||||
wav, _, _ = trim_long_silences(wav_path, sample_rate)
|
||||
else:
|
||||
wav, _ = librosa.core.load(wav_path, sr=sample_rate)
|
||||
else:
|
||||
wav = wav_path
|
||||
wav_orig = np.copy(wav)
|
||||
|
||||
if loud_norm:
|
||||
meter = pyln.Meter(sample_rate) # create BS.1770 meter
|
||||
loudness = meter.integrated_loudness(wav)
|
||||
wav = pyln.normalize.loudness(wav, loudness, -22.0)
|
||||
if np.abs(wav).max() > 1:
|
||||
wav = wav / np.abs(wav).max()
|
||||
|
||||
# get amplitude spectrogram
|
||||
x_stft = librosa.stft(wav, n_fft=fft_size, hop_length=hop_size,
|
||||
win_length=win_length, window=window, pad_mode="constant")
|
||||
linear_spc = np.abs(x_stft) # (n_bins, T)
|
||||
|
||||
# get mel basis
|
||||
fmin = 0 if fmin == -1 else fmin
|
||||
fmax = sample_rate / 2 if fmax == -1 else fmax
|
||||
mel_basis = librosa.filters.mel(sr=sample_rate, n_fft=fft_size, n_mels=num_mels, fmin=fmin, fmax=fmax)
|
||||
|
||||
# calculate mel spec
|
||||
mel = mel_basis @ linear_spc
|
||||
mel = np.log10(np.maximum(eps, mel)) # (n_mel_bins, T)
|
||||
l_pad, r_pad = librosa_pad_lr(wav, fft_size, hop_size, 1)
|
||||
wav = np.pad(wav, (l_pad, r_pad), mode='constant', constant_values=0.0)
|
||||
wav = wav[:mel.shape[1] * hop_size]
|
||||
|
||||
# log linear spec
|
||||
linear_spc = np.log10(np.maximum(eps, linear_spc))
|
||||
return {'wav': wav, 'mel': mel.T, 'linear': linear_spc.T, 'mel_basis': mel_basis, 'wav_orig': wav_orig}
|
||||
|
||||
def get_wav_num_frames(path, sr=None):
|
||||
try:
|
||||
with wave.open(path, 'rb') as f:
|
||||
sr_ = f.getframerate()
|
||||
if sr is None:
|
||||
sr = sr_
|
||||
return int(f.getnframes() / (sr_ / sr))
|
||||
except wave.Error:
|
||||
wav_file, sr_ = sf.read(path, dtype='float32')
|
||||
if sr is None:
|
||||
sr = sr_
|
||||
return int(len(wav_file) / (sr_ / sr))
|
||||
except:
|
||||
wav_file, sr_ = librosa.core.load(path, sr=sr)
|
||||
return len(wav_file)
|
||||
@@ -0,0 +1,90 @@
|
||||
import re
|
||||
|
||||
import torch
|
||||
import numpy as np
|
||||
|
||||
from ..text.text_encoder import is_sil_phoneme
|
||||
|
||||
|
||||
def get_mel2ph(tg_fn, ph, mel, hop_size, audio_sample_rate, min_sil_duration=0):
|
||||
from textgrid import TextGrid
|
||||
ph_list = ph.split(" ")
|
||||
itvs = TextGrid.fromFile(tg_fn)[1]
|
||||
itvs_ = []
|
||||
for i in range(len(itvs)):
|
||||
if itvs[i].maxTime - itvs[i].minTime < min_sil_duration and i > 0 and is_sil_phoneme(itvs[i].mark):
|
||||
itvs_[-1].maxTime = itvs[i].maxTime
|
||||
else:
|
||||
itvs_.append(itvs[i])
|
||||
itvs.intervals = itvs_
|
||||
itv_marks = [itv.mark for itv in itvs]
|
||||
tg_len = len([x for x in itvs if not is_sil_phoneme(x.mark)])
|
||||
ph_len = len([x for x in ph_list if not is_sil_phoneme(x)])
|
||||
assert tg_len == ph_len, (tg_len, ph_len, itv_marks, ph_list, tg_fn)
|
||||
mel2ph = np.zeros([mel.shape[0]], int)
|
||||
i_itv = 0
|
||||
i_ph = 0
|
||||
while i_itv < len(itvs):
|
||||
itv = itvs[i_itv]
|
||||
ph = ph_list[i_ph]
|
||||
itv_ph = itv.mark
|
||||
start_frame = int(itv.minTime * audio_sample_rate / hop_size + 0.5)
|
||||
end_frame = int(itv.maxTime * audio_sample_rate / hop_size + 0.5)
|
||||
if is_sil_phoneme(itv_ph) and not is_sil_phoneme(ph):
|
||||
mel2ph[start_frame:end_frame] = i_ph
|
||||
i_itv += 1
|
||||
elif not is_sil_phoneme(itv_ph) and is_sil_phoneme(ph):
|
||||
i_ph += 1
|
||||
else:
|
||||
if not ((is_sil_phoneme(itv_ph) and is_sil_phoneme(ph)) \
|
||||
or re.sub(r'\d+', '', itv_ph.lower()) == re.sub(r'\d+', '', ph.lower())):
|
||||
print(f"| WARN: {tg_fn} phs are not same: ", itv_ph, ph, itv_marks, ph_list)
|
||||
mel2ph[start_frame:end_frame] = i_ph + 1
|
||||
i_ph += 1
|
||||
i_itv += 1
|
||||
mel2ph[-1] = mel2ph[-2]
|
||||
assert not np.any(mel2ph == 0)
|
||||
T_t = len(ph_list)
|
||||
dur = mel2token_to_dur(mel2ph, T_t)
|
||||
return mel2ph.tolist(), dur.tolist()
|
||||
|
||||
|
||||
def split_audio_by_mel2ph(audio, mel2ph, hop_size, audio_num_mel_bins):
|
||||
if isinstance(audio, torch.Tensor):
|
||||
audio = audio.numpy()
|
||||
if isinstance(mel2ph, torch.Tensor):
|
||||
mel2ph = mel2ph.numpy()
|
||||
assert len(audio.shape) == 1, len(mel2ph.shape) == 1
|
||||
split_locs = []
|
||||
for i in range(1, len(mel2ph)):
|
||||
if mel2ph[i] != mel2ph[i - 1]:
|
||||
split_loc = i * hop_size
|
||||
split_locs.append(split_loc)
|
||||
|
||||
new_audio = []
|
||||
for i in range(len(split_locs) - 1):
|
||||
new_audio.append(audio[split_locs[i]:split_locs[i + 1]])
|
||||
new_audio.append(np.zeros([0.5 * audio_num_mel_bins]))
|
||||
return np.concatenate(new_audio)
|
||||
|
||||
|
||||
def mel2token_to_dur(mel2token, T_txt=None, max_dur=None):
|
||||
is_torch = isinstance(mel2token, torch.Tensor)
|
||||
has_batch_dim = True
|
||||
if not is_torch:
|
||||
mel2token = torch.LongTensor(mel2token)
|
||||
if T_txt is None:
|
||||
T_txt = mel2token.max()
|
||||
if len(mel2token.shape) == 1:
|
||||
mel2token = mel2token[None, ...]
|
||||
has_batch_dim = False
|
||||
B, _ = mel2token.shape
|
||||
dur = mel2token.new_zeros(B, T_txt + 1).scatter_add(1, mel2token, torch.ones_like(mel2token))
|
||||
dur = dur[:, 1:]
|
||||
if max_dur is not None:
|
||||
dur = dur.clamp(max=max_dur)
|
||||
if not is_torch:
|
||||
dur = dur.numpy()
|
||||
if not has_batch_dim:
|
||||
dur = dur[0]
|
||||
return dur
|
||||
@@ -0,0 +1,22 @@
|
||||
import subprocess
|
||||
|
||||
import numpy as np
|
||||
from scipy.io import wavfile
|
||||
|
||||
|
||||
def save_wav(wav, path, sr, norm=False):
|
||||
if norm:
|
||||
wav = wav / np.abs(wav).max()
|
||||
wav = wav * 32767
|
||||
wavfile.write(path[:-4] + '.wav', sr, wav.astype(np.int16))
|
||||
if path[-4:] == '.mp3':
|
||||
to_mp3(path[:-4])
|
||||
|
||||
|
||||
def to_mp3(out_path):
|
||||
if out_path[-4:] == '.wav':
|
||||
out_path = out_path[:-4]
|
||||
subprocess.check_call(
|
||||
f'ffmpeg -threads 1 -loglevel error -i "{out_path}.wav" -vn -b:a 192k -y -hide_banner -async 1 "{out_path}.mp3"',
|
||||
shell=True, stdin=subprocess.PIPE)
|
||||
subprocess.check_call(f'rm -f "{out_path}.wav"', shell=True)
|
||||
@@ -0,0 +1,139 @@
|
||||
import math
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.utils.data
|
||||
from librosa.filters import mel as librosa_mel_fn
|
||||
from scipy.io.wavfile import read
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
MAX_WAV_VALUE = 32768.0
|
||||
|
||||
|
||||
def load_wav(full_path):
|
||||
sampling_rate, data = read(full_path)
|
||||
return data, sampling_rate
|
||||
|
||||
|
||||
def dynamic_range_compression(x, C=1, clip_val=1e-5):
|
||||
return np.log10(np.clip(x, a_min=clip_val, a_max=None) * C)
|
||||
|
||||
|
||||
def dynamic_range_decompression(x, C=1):
|
||||
return np.exp(x) / C
|
||||
|
||||
|
||||
def dynamic_range_compression_torch(x, C=1, clip_val=1e-5):
|
||||
return torch.log10(torch.clamp(x, min=clip_val) * C)
|
||||
|
||||
|
||||
def dynamic_range_decompression_torch(x, C=1):
|
||||
return torch.exp(x) / C
|
||||
|
||||
|
||||
def spectral_normalize_torch(magnitudes):
|
||||
output = dynamic_range_compression_torch(magnitudes)
|
||||
return output
|
||||
|
||||
|
||||
def spectral_de_normalize_torch(magnitudes):
|
||||
output = dynamic_range_decompression_torch(magnitudes)
|
||||
return output
|
||||
|
||||
|
||||
class MelNet(nn.Module):
|
||||
def __init__(self, hparams, device='cpu') -> None:
|
||||
super().__init__()
|
||||
self.n_fft = hparams['fft_size']
|
||||
self.num_mels = hparams['audio_num_mel_bins']
|
||||
self.sampling_rate = hparams['audio_sample_rate']
|
||||
self.hop_size = hparams['hop_size']
|
||||
self.win_size = hparams['win_size']
|
||||
self.fmin = hparams['fmin']
|
||||
self.fmax = hparams['fmax']
|
||||
self.device = device
|
||||
|
||||
mel = librosa_mel_fn(sr=self.sampling_rate, n_fft=self.n_fft, n_mels=self.num_mels, fmin=self.fmin,
|
||||
fmax=self.fmax)
|
||||
self.mel_basis = torch.from_numpy(mel).float().to(self.device)
|
||||
self.hann_window = torch.hann_window(self.win_size).to(self.device)
|
||||
|
||||
def to(self, device, **kwagrs):
|
||||
super().to(device=device, **kwagrs)
|
||||
self.mel_basis = self.mel_basis.to(device)
|
||||
self.hann_window = self.hann_window.to(device)
|
||||
self.device = device
|
||||
|
||||
def forward(self, y, center=False, complex=False):
|
||||
if isinstance(y, np.ndarray):
|
||||
y = torch.FloatTensor(y)
|
||||
if len(y.shape) == 1:
|
||||
y = y.unsqueeze(0)
|
||||
y = y.clamp(min=-1., max=1.).to(self.device)
|
||||
|
||||
pad_length = math.ceil(y.shape[1] / self.hop_size) * self.hop_size - y.shape[1]
|
||||
y = torch.nn.functional.pad(y.unsqueeze(1),
|
||||
[int((self.n_fft - self.hop_size) / 2),
|
||||
int((self.n_fft - self.hop_size) / 2 + pad_length)],
|
||||
mode='reflect')
|
||||
y = y.squeeze(1)
|
||||
|
||||
spec = torch.stft(y, self.n_fft, hop_length=self.hop_size, win_length=self.win_size, window=self.hann_window,
|
||||
center=center, pad_mode='reflect', normalized=False, onesided=True, return_complex=True)
|
||||
if not complex:
|
||||
spec = torch.view_as_real(spec)
|
||||
spec = torch.sqrt(spec.pow(2).sum(-1) + (1e-9)) # [B, n_fft, T]
|
||||
spec = torch.matmul(self.mel_basis, spec)
|
||||
spec = spectral_normalize_torch(spec)
|
||||
spec = spec.transpose(1, 2) # [B, T, n_fft]
|
||||
else:
|
||||
B, C, T, _ = spec.shape
|
||||
spec = spec.transpose(1, 2) # [B, T, n_fft, 2]
|
||||
return spec
|
||||
|
||||
|
||||
## below can be used in one gpu, but not ddp
|
||||
mel_basis = {}
|
||||
hann_window = {}
|
||||
|
||||
|
||||
def mel_spectrogram(y, hparams, center=False, complex=False): # y should be a tensor with shape (b,wav_len)
|
||||
# hop_size: 512 # For 22050Hz, 275 ~= 12.5 ms (0.0125 * sample_rate)
|
||||
# win_size: 2048 # For 22050Hz, 1100 ~= 50 ms (If None, win_size: fft_size) (0.05 * sample_rate)
|
||||
# fmin: 55 # Set this to 55 if your speaker is male! if female, 95 should help taking off noise. (To test depending on dataset. Pitch info: male~[65, 260], female~[100, 525])
|
||||
# fmax: 10000 # To be increased/reduced depending on data.
|
||||
# fft_size: 2048 # Extra window size is filled with 0 paddings to match this parameter
|
||||
# n_fft, num_mels, sampling_rate, hop_size, win_size, fmin, fmax,
|
||||
n_fft = hparams['fft_size']
|
||||
num_mels = hparams['audio_num_mel_bins']
|
||||
sampling_rate = hparams['audio_sample_rate']
|
||||
hop_size = hparams['hop_size']
|
||||
win_size = hparams['win_size']
|
||||
fmin = hparams['fmin']
|
||||
fmax = hparams['fmax']
|
||||
if isinstance(y, np.ndarray):
|
||||
y = torch.FloatTensor(y)
|
||||
if len(y.shape) == 1:
|
||||
y = y.unsqueeze(0)
|
||||
y = y.clamp(min=-1., max=1.)
|
||||
global mel_basis, hann_window
|
||||
if fmax not in mel_basis:
|
||||
mel = librosa_mel_fn(sampling_rate, n_fft, num_mels, fmin, fmax)
|
||||
mel_basis[str(fmax) + '_' + str(y.device)] = torch.from_numpy(mel).float().to(y.device)
|
||||
hann_window[str(y.device)] = torch.hann_window(win_size).to(y.device)
|
||||
|
||||
y = torch.nn.functional.pad(y.unsqueeze(1), [int((n_fft - hop_size) / 2), int((n_fft - hop_size) / 2)],
|
||||
mode='reflect')
|
||||
y = y.squeeze(1)
|
||||
|
||||
spec = torch.stft(y, n_fft, hop_length=hop_size, win_length=win_size, window=hann_window[str(y.device)],
|
||||
center=center, pad_mode='reflect', normalized=False, onesided=True, return_complex=complex)
|
||||
|
||||
if not complex:
|
||||
spec = torch.sqrt(spec.pow(2).sum(-1) + (1e-9))
|
||||
spec = torch.matmul(mel_basis[str(fmax) + '_' + str(y.device)], spec)
|
||||
spec = spectral_normalize_torch(spec)
|
||||
else:
|
||||
B, C, T, _ = spec.shape
|
||||
spec = spec.transpose(1, 2) # [B, T, n_fft, 2]
|
||||
return spec
|
||||
@@ -0,0 +1,60 @@
|
||||
import math
|
||||
import numpy as np
|
||||
|
||||
PITCH_EXTRACTOR = {}
|
||||
|
||||
|
||||
def register_pitch_extractor(name):
|
||||
def register_pitch_extractor_(cls):
|
||||
PITCH_EXTRACTOR[name] = cls
|
||||
return cls
|
||||
|
||||
return register_pitch_extractor_
|
||||
|
||||
|
||||
def get_pitch_extractor(name):
|
||||
return PITCH_EXTRACTOR[name]
|
||||
|
||||
|
||||
def extract_pitch_simple(wav):
|
||||
from ..commons.hparams import hparams
|
||||
return extract_pitch(hparams['pitch_extractor'], wav,
|
||||
hparams['hop_size'], hparams['audio_sample_rate'],
|
||||
f0_min=hparams['f0_min'], f0_max=hparams['f0_max'])
|
||||
|
||||
|
||||
def extract_pitch(extractor_name, wav_data, hop_size, audio_sample_rate, f0_min=75, f0_max=800, **kwargs):
|
||||
return get_pitch_extractor(extractor_name)(wav_data, hop_size, audio_sample_rate, f0_min, f0_max, **kwargs)
|
||||
|
||||
|
||||
@register_pitch_extractor('parselmouth')
|
||||
def parselmouth_pitch(wav_data, hop_size, audio_sample_rate, f0_min, f0_max,
|
||||
voicing_threshold=0.6, *args, **kwargs):
|
||||
import parselmouth
|
||||
time_step = hop_size / audio_sample_rate * 1000
|
||||
n_mel_frames = int(len(wav_data) // hop_size)
|
||||
f0_pm = parselmouth.Sound(wav_data, audio_sample_rate).to_pitch_ac(
|
||||
time_step=time_step / 1000, voicing_threshold=voicing_threshold,
|
||||
pitch_floor=f0_min, pitch_ceiling=f0_max).selected_array['frequency']
|
||||
pad_size = (n_mel_frames - len(f0_pm) + 1) // 2
|
||||
f0 = np.pad(f0_pm, [[pad_size, n_mel_frames - len(f0_pm) - pad_size]], mode='constant')
|
||||
return f0
|
||||
|
||||
@register_pitch_extractor('pyworld')
|
||||
def pyworld_pitch(wav_data, hop_size, audio_sample_rate, f0_min, f0_max,
|
||||
voicing_threshold=0.6, *args, **kwargs):
|
||||
import pyworld as pw
|
||||
# f0, _ = pw.harvest(wav_data.astype(np.double), audio_sample_rate, f0_floor=f0_min, f0_ceil=f0_max,
|
||||
# frame_period=hop_size * 1000 / audio_sample_rate)
|
||||
f0, _ = pw.dio(wav_data.astype(np.double), audio_sample_rate, f0_floor=f0_min, f0_ceil=f0_max, frame_period=hop_size * 1000 / audio_sample_rate)
|
||||
f0[f0 < f0_min] = 0.0
|
||||
f0[f0 > f0_max] = 0.0
|
||||
n_mel_frames = math.ceil(len(wav_data) / hop_size)
|
||||
if n_mel_frames > len(f0):
|
||||
pad_size = (n_mel_frames - len(f0) + 1) // 2
|
||||
f0 = np.pad(f0, [[pad_size, n_mel_frames - len(f0) - pad_size]], mode='constant')
|
||||
elif n_mel_frames < len(f0):
|
||||
left_del = (len(f0) - n_mel_frames + 1) // 2
|
||||
right_del = len(f0) - n_mel_frames - left_del
|
||||
f0 = f0[left_del: (-right_del if right_del > 0 else len(f0))]
|
||||
return f0
|
||||
@@ -0,0 +1,303 @@
|
||||
import numpy as np
|
||||
import torch
|
||||
import pretty_midi
|
||||
|
||||
def to_lf0(f0):
|
||||
f0[f0 < 1.0e-5] = 1.0e-6
|
||||
lf0 = f0.log() if isinstance(f0, torch.Tensor) else np.log(f0)
|
||||
lf0[f0 < 1.0e-5] = - 1.0E+10
|
||||
return lf0
|
||||
|
||||
|
||||
def to_f0(lf0):
|
||||
f0 = np.where(lf0 <= 0, 0.0, np.exp(lf0))
|
||||
return f0.flatten()
|
||||
|
||||
|
||||
def f0_to_coarse(f0, f0_bin=256, f0_max=900.0, f0_min=50.0):
|
||||
f0_mel_min = 1127 * np.log(1 + f0_min / 700)
|
||||
f0_mel_max = 1127 * np.log(1 + f0_max / 700)
|
||||
is_torch = isinstance(f0, torch.Tensor)
|
||||
f0_mel = 1127 * (1 + f0 / 700).log() if is_torch else 1127 * np.log(1 + f0 / 700)
|
||||
f0_mel[f0_mel > 0] = (f0_mel[f0_mel > 0] - f0_mel_min) * (f0_bin - 2) / (f0_mel_max - f0_mel_min) + 1
|
||||
|
||||
f0_mel[f0_mel <= 1] = 1
|
||||
f0_mel[f0_mel > f0_bin - 1] = f0_bin - 1
|
||||
f0_coarse = (f0_mel + 0.5).long() if is_torch else np.rint(f0_mel).astype(int)
|
||||
assert f0_coarse.max() <= f0_bin-1 and f0_coarse.min() >= 1, (f0_coarse.max(), f0_coarse.min(), f0.min(), f0.max())
|
||||
return f0_coarse
|
||||
|
||||
|
||||
def coarse_to_f0(f0_coarse, f0_bin=256, f0_max=900.0, f0_min=50.0):
|
||||
f0_mel_min = 1127 * np.log(1 + f0_min / 700)
|
||||
f0_mel_max = 1127 * np.log(1 + f0_max / 700)
|
||||
uv = f0_coarse == 1
|
||||
f0 = f0_mel_min + (f0_coarse - 1) * (f0_mel_max - f0_mel_min) / (f0_bin - 2)
|
||||
f0 = ((f0 / 1127).exp() - 1) * 700
|
||||
f0[uv] = 0
|
||||
return f0
|
||||
|
||||
|
||||
def norm_f0(f0, uv, pitch_norm='log', f0_mean=400, f0_std=100):
|
||||
is_torch = isinstance(f0, torch.Tensor)
|
||||
if pitch_norm == 'standard':
|
||||
f0 = (f0 - f0_mean) / f0_std
|
||||
if pitch_norm == 'log':
|
||||
f0 = torch.log2(f0 + 1e-8) if is_torch else np.log2(f0 + 1e-8)
|
||||
if uv is not None:
|
||||
f0[uv > 0] = 0
|
||||
return f0
|
||||
|
||||
|
||||
def norm_interp_f0(f0, pitch_norm='log', f0_mean=None, f0_std=None):
|
||||
is_torch = isinstance(f0, torch.Tensor)
|
||||
if is_torch:
|
||||
device = f0.device
|
||||
f0 = f0.data.cpu().numpy()
|
||||
uv = f0 == 0
|
||||
f0 = norm_f0(f0, uv, pitch_norm, f0_mean, f0_std)
|
||||
if sum(uv) == len(f0):
|
||||
f0[uv] = 0
|
||||
elif sum(uv) > 0:
|
||||
f0[uv] = np.interp(np.where(uv)[0], np.where(~uv)[0], f0[~uv])
|
||||
if is_torch:
|
||||
uv = torch.FloatTensor(uv)
|
||||
f0 = torch.FloatTensor(f0)
|
||||
f0 = f0.to(device)
|
||||
uv = uv.to(device)
|
||||
return f0, uv
|
||||
|
||||
|
||||
def denorm_f0(f0, uv, pitch_norm='log', f0_mean=400, f0_std=100, pitch_padding=None, min=50, max=900):
|
||||
is_torch = isinstance(f0, torch.Tensor)
|
||||
if pitch_norm == 'standard':
|
||||
f0 = f0 * f0_std + f0_mean
|
||||
if pitch_norm == 'log':
|
||||
f0 = 2 ** f0
|
||||
f0 = f0.clamp(min=min, max=max) if is_torch else np.clip(f0, a_min=min, a_max=max)
|
||||
if uv is not None:
|
||||
f0[uv > 0] = 0
|
||||
if pitch_padding is not None:
|
||||
f0[pitch_padding] = 0
|
||||
return f0
|
||||
|
||||
def interp_f0(f0, uv=None):
|
||||
if uv is None:
|
||||
uv = f0 == 0
|
||||
f0 = norm_f0(f0, uv)
|
||||
if uv.any() and not uv.all():
|
||||
f0[uv] = np.interp(np.where(uv)[0], np.where(~uv)[0], f0[~uv])
|
||||
return denorm_f0(f0, uv=None), uv
|
||||
|
||||
def resample_align_curve(points: np.ndarray, original_timestep: float, target_timestep: float, align_length=-1):
|
||||
t_max = (len(points) - 1) * original_timestep
|
||||
curve_interp = np.interp(
|
||||
np.arange(0, t_max, target_timestep),
|
||||
original_timestep * np.arange(len(points)),
|
||||
points
|
||||
).astype(points.dtype)
|
||||
if align_length > 0:
|
||||
delta_l = align_length - len(curve_interp)
|
||||
if delta_l < 0:
|
||||
curve_interp = curve_interp[:align_length]
|
||||
elif delta_l > 0:
|
||||
curve_interp = np.concatenate((curve_interp, np.full(delta_l, fill_value=curve_interp[-1])), axis=0)
|
||||
return curve_interp
|
||||
|
||||
def midi_to_hz(midi):
|
||||
if type(midi) == np.ndarray:
|
||||
non_mask = midi == 0
|
||||
freq_hz = 440.0 * 2.0 ** ((midi - 69.0) / 12.0)
|
||||
freq_hz[non_mask] = 0
|
||||
else:
|
||||
freq_hz = 440.0 * 2.0 ** ((midi - 69.0) / 12.0)
|
||||
return freq_hz
|
||||
|
||||
def hz_to_midi(hz):
|
||||
if type(hz) == torch.Tensor:
|
||||
non_mask = hz == 0
|
||||
midi = 69.0 + 12.0 * (torch.log2(hz) - torch.log2(torch.Tensor(440.0)))
|
||||
midi[non_mask] = 0
|
||||
elif type(hz) == np.ndarray:
|
||||
non_mask = hz == 0
|
||||
midi = 69.0 + 12.0 * (np.log2(hz) - np.log2(440.0))
|
||||
midi[non_mask] = 0
|
||||
else:
|
||||
midi = 69.0 + 12.0 * (np.log2(hz) - np.log2(440.0))
|
||||
if hz == 0:
|
||||
midi = 0
|
||||
return midi
|
||||
|
||||
def boundary2Interval(bd):
|
||||
# bd has a shape of [T] with T frames
|
||||
is_torch = isinstance(bd, torch.Tensor)
|
||||
if is_torch:
|
||||
device = bd.device
|
||||
bd = bd.data.cpu().numpy()
|
||||
assert len(bd.shape) == 1
|
||||
# force valid begin and end
|
||||
# bd[0] = 0 # took care of in regulate_boundary()
|
||||
# bd[-1] = 0
|
||||
ret = np.zeros(shape=(bd.sum() + 1, 2), dtype=int)
|
||||
ret_idx = 0
|
||||
ret[0, 0] = 0
|
||||
for i, u in enumerate(bd):
|
||||
if i == 0:
|
||||
continue
|
||||
if u == 1:
|
||||
ret[ret_idx, 1] = i
|
||||
ret[ret_idx+1, 0] = i
|
||||
ret_idx += 1
|
||||
ret[-1, 1] = bd.shape[0] - 1
|
||||
if is_torch:
|
||||
ret = torch.LongTensor(ret).to(device)
|
||||
return ret
|
||||
|
||||
def validate_pitch_and_itv(notes, note_itv):
|
||||
# notes [T]
|
||||
# note_itv [T, 2]
|
||||
assert notes.shape[0] == note_itv.shape[0]
|
||||
res_notes = []
|
||||
res_note_itv = []
|
||||
for idx in range(notes.shape[0]):
|
||||
pitch, itv = notes[idx], note_itv[idx]
|
||||
if itv[0] >= itv[1]:
|
||||
raise RuntimeError("The note duration should be positive")
|
||||
if pitch == 0:
|
||||
continue
|
||||
res_notes.append(pitch)
|
||||
res_note_itv.append([itv[0], itv[1]])
|
||||
res_notes = np.array(res_notes)
|
||||
res_note_itv = np.array(res_note_itv)
|
||||
return res_notes, res_note_itv
|
||||
|
||||
def save_midi(notes, note_itv, midi_path):
|
||||
# notes [T]
|
||||
# note_itv [T, 2]
|
||||
notes, note_itv = validate_pitch_and_itv(notes, note_itv)
|
||||
if notes.shape == (0,):
|
||||
return None
|
||||
assert notes.shape[0] == note_itv.shape[0]
|
||||
piano_chord = pretty_midi.PrettyMIDI()
|
||||
piano_program = pretty_midi.instrument_name_to_program('Acoustic Grand Piano')
|
||||
piano = pretty_midi.Instrument(program=piano_program)
|
||||
for idx in range(notes.shape[0]):
|
||||
pitch, itv = notes[idx], note_itv[idx]
|
||||
note = pretty_midi.Note(velocity=120, pitch=pitch, start=itv[0], end=itv[1])
|
||||
piano.notes.append(note)
|
||||
piano_chord.remove_invalid_notes()
|
||||
piano_chord.instruments.append(piano)
|
||||
piano_chord.write(midi_path)
|
||||
return piano_chord
|
||||
|
||||
def midi2NoteInterval(mid):
|
||||
assert type(mid) == pretty_midi.PrettyMIDI
|
||||
if len(mid.instruments) == 0 or len(mid.instruments[0].notes) == 0:
|
||||
return None
|
||||
ret = np.zeros(shape=(len(mid.instruments[0].notes), 2))
|
||||
for i, note in enumerate(mid.instruments[0].notes):
|
||||
ret[i, 0] = note.start
|
||||
ret[i, 1] = note.end
|
||||
return ret
|
||||
|
||||
def midi2NotePitch(mid):
|
||||
assert type(mid) == pretty_midi.PrettyMIDI
|
||||
if len(mid.instruments) == 0 or len(mid.instruments[0].notes) == 0:
|
||||
return None
|
||||
ret = np.zeros(shape=len(mid.instruments[0].notes))
|
||||
for i, note in enumerate(mid.instruments[0].notes):
|
||||
ret[i] = note.pitch
|
||||
return ret
|
||||
|
||||
def midi_onset_eval(mid_gt, mid_pred):
|
||||
import mir_eval
|
||||
interval_true = midi2NoteInterval(mid_gt)
|
||||
if interval_true is None:
|
||||
raise RuntimeError('Midi ground truth is None')
|
||||
interval_pred = midi2NoteInterval(mid_pred)
|
||||
if interval_pred is None:
|
||||
return 0, 0, 0
|
||||
onset_p, onset_r, onset_f = mir_eval.transcription.onset_precision_recall_f1(
|
||||
interval_true, interval_pred, onset_tolerance=0.05, strict=False, beta=1.0)
|
||||
return onset_p, onset_r, onset_f
|
||||
|
||||
def midi_offset_eval(mid_gt, mid_pred):
|
||||
import mir_eval
|
||||
interval_true = midi2NoteInterval(mid_gt)
|
||||
if interval_true is None:
|
||||
raise RuntimeError('Midi ground truth is None')
|
||||
interval_pred = midi2NoteInterval(mid_pred)
|
||||
if interval_pred is None:
|
||||
return 0, 0, 0
|
||||
offset_p, offset_r, offset_f = mir_eval.transcription.offset_precision_recall_f1(
|
||||
interval_true, interval_pred, offset_ratio=0.2, offset_min_tolerance=0.05, strict=False, beta=1.0)
|
||||
return offset_p, offset_r, offset_f
|
||||
|
||||
def midi_pitch_eval(mid_gt, mid_pred, offset_ratio=0.2):
|
||||
import mir_eval
|
||||
interval_true = midi2NoteInterval(mid_gt)
|
||||
pitch_true = midi_to_hz(midi2NotePitch(mid_gt))
|
||||
if interval_true is None or pitch_true is None:
|
||||
raise RuntimeError('Midi ground truth is None')
|
||||
interval_pred = midi2NoteInterval(mid_pred)
|
||||
pitch_pred = midi2NotePitch(mid_pred)
|
||||
if interval_pred is None:
|
||||
return 0, 0, 0, 0
|
||||
if pitch_pred is None:
|
||||
pitch_pred = np.zeros(interval_pred.shape[0])
|
||||
pitch_pred = midi_to_hz(pitch_pred)
|
||||
overlap_p, overlap_r, overlap_f, avg_overlap_ratio = mir_eval.transcription.precision_recall_f1_overlap(
|
||||
interval_true, pitch_true, interval_pred, pitch_pred, onset_tolerance=0.05, pitch_tolerance=50.0,
|
||||
offset_ratio=offset_ratio, offset_min_tolerance=0.05, strict=False, beta=1.0)
|
||||
return overlap_p, overlap_r, overlap_f, avg_overlap_ratio
|
||||
|
||||
def midi_COn_eval(mid_gt, mid_pred):
|
||||
return midi_onset_eval(mid_gt, mid_pred)
|
||||
|
||||
def midi_COnP_eval(mid_gt, mid_pred):
|
||||
return midi_pitch_eval(mid_gt, mid_pred, offset_ratio=None)
|
||||
|
||||
def midi_COnPOff_eval(mid_gt, mid_pred):
|
||||
return midi_pitch_eval(mid_gt, mid_pred)
|
||||
|
||||
def midi_melody_eval(mid_gt, mid_pred, hop_size=256, sample_rate=48000):
|
||||
interval_true = midi2NoteInterval(mid_gt)
|
||||
pitch_true = midi_to_hz(midi2NotePitch(mid_gt))
|
||||
if interval_true is None or pitch_true is None:
|
||||
raise RuntimeError('Midi ground truth is None')
|
||||
interval_pred = midi2NoteInterval(mid_pred)
|
||||
pitch_pred = midi2NotePitch(mid_pred)
|
||||
if interval_pred is None:
|
||||
return 0, 0, 0, 0
|
||||
if pitch_pred is None:
|
||||
pitch_pred = np.zeros(interval_pred.shape[0])
|
||||
pitch_pred = midi_to_hz(pitch_pred)
|
||||
|
||||
vr, vfa, rpa, rca, oa = melody_eval_pitch_and_itv(
|
||||
pitch_true, interval_true, pitch_pred, interval_pred, hop_size, sample_rate)
|
||||
|
||||
return vr, vfa, rpa, rca, oa
|
||||
|
||||
def melody_eval_pitch_and_itv(pitch_true, interval_true, pitch_pred, interval_pred, hop_size=256, sample_rate=48000):
|
||||
import mir_eval
|
||||
t_gt = np.arange(0, interval_true[-1][1], hop_size / sample_rate)
|
||||
freq_gt = np.zeros_like(t_gt)
|
||||
for idx in range(len(pitch_true)):
|
||||
freq_gt[min(len(freq_gt) - 1, round(interval_true[idx][0] * sample_rate / hop_size)): round(
|
||||
interval_true[idx][1] * sample_rate / hop_size)] = pitch_true[idx]
|
||||
|
||||
t_pred = np.arange(0, interval_pred[-1][1], hop_size / sample_rate)
|
||||
freq_pred = np.zeros_like(t_pred)
|
||||
for idx in range(len(pitch_pred)):
|
||||
freq_pred[min(len(freq_pred) - 1, round(interval_pred[idx][0] * sample_rate / hop_size)): round(
|
||||
interval_pred[idx][1] * sample_rate / hop_size)] = pitch_pred[idx]
|
||||
|
||||
ref_voicing, ref_cent, est_voicing, est_cent = mir_eval.melody.to_cent_voicing(t_gt, freq_gt,
|
||||
t_pred, freq_pred)
|
||||
vr, vfa = mir_eval.melody.voicing_measures(ref_voicing,
|
||||
est_voicing) # voicing recall, voicing false alarm
|
||||
rpa = mir_eval.melody.raw_pitch_accuracy(ref_voicing, ref_cent, est_voicing, est_cent)
|
||||
rca = mir_eval.melody.raw_chroma_accuracy(ref_voicing, ref_cent, est_voicing, est_cent)
|
||||
oa = mir_eval.melody.overall_accuracy(ref_voicing, ref_cent, est_voicing, est_cent)
|
||||
|
||||
return vr, vfa, rpa, rca, oa
|
||||
@@ -0,0 +1,78 @@
|
||||
from skimage.transform import resize
|
||||
import struct
|
||||
import webrtcvad
|
||||
from scipy.ndimage.morphology import binary_dilation
|
||||
import librosa
|
||||
import numpy as np
|
||||
import pyloudnorm as pyln
|
||||
import warnings
|
||||
|
||||
warnings.filterwarnings("ignore", message="Possible clipped samples in output")
|
||||
|
||||
int16_max = (2 ** 15) - 1
|
||||
|
||||
|
||||
def trim_long_silences(path, sr=None, return_raw_wav=False, norm=True, vad_max_silence_length=12):
|
||||
"""
|
||||
Ensures that segments without voice in the waveform remain no longer than a
|
||||
threshold determined by the VAD parameters in params.py.
|
||||
:param wav: the raw waveform as a numpy array of floats
|
||||
:param vad_max_silence_length: Maximum number of consecutive silent frames a segment can have.
|
||||
:return: the same waveform with silences trimmed away (length <= original wav length)
|
||||
"""
|
||||
|
||||
## Voice Activation Detection
|
||||
# Window size of the VAD. Must be either 10, 20 or 30 milliseconds.
|
||||
# This sets the granularity of the VAD. Should not need to be changed.
|
||||
sampling_rate = 16000
|
||||
wav_raw, sr = librosa.core.load(path, sr=sr)
|
||||
|
||||
if norm:
|
||||
meter = pyln.Meter(sr) # create BS.1770 meter
|
||||
loudness = meter.integrated_loudness(wav_raw)
|
||||
wav_raw = pyln.normalize.loudness(wav_raw, loudness, -20.0)
|
||||
if np.abs(wav_raw).max() > 1.0:
|
||||
wav_raw = wav_raw / np.abs(wav_raw).max()
|
||||
|
||||
wav = librosa.resample(wav_raw, sr, sampling_rate, res_type='kaiser_best')
|
||||
|
||||
vad_window_length = 30 # In milliseconds
|
||||
# Number of frames to average together when performing the moving average smoothing.
|
||||
# The larger this value, the larger the VAD variations must be to not get smoothed out.
|
||||
vad_moving_average_width = 8
|
||||
|
||||
# Compute the voice detection window size
|
||||
samples_per_window = (vad_window_length * sampling_rate) // 1000
|
||||
|
||||
# Trim the end of the audio to have a multiple of the window size
|
||||
wav = wav[:len(wav) - (len(wav) % samples_per_window)]
|
||||
|
||||
# Convert the float waveform to 16-bit mono PCM
|
||||
pcm_wave = struct.pack("%dh" % len(wav), *(np.round(wav * int16_max)).astype(np.int16))
|
||||
|
||||
# Perform voice activation detection
|
||||
voice_flags = []
|
||||
vad = webrtcvad.Vad(mode=3)
|
||||
for window_start in range(0, len(wav), samples_per_window):
|
||||
window_end = window_start + samples_per_window
|
||||
voice_flags.append(vad.is_speech(pcm_wave[window_start * 2:window_end * 2],
|
||||
sample_rate=sampling_rate))
|
||||
voice_flags = np.array(voice_flags)
|
||||
|
||||
# Smooth the voice detection with a moving average
|
||||
def moving_average(array, width):
|
||||
array_padded = np.concatenate((np.zeros((width - 1) // 2), array, np.zeros(width // 2)))
|
||||
ret = np.cumsum(array_padded, dtype=float)
|
||||
ret[width:] = ret[width:] - ret[:-width]
|
||||
return ret[width - 1:] / width
|
||||
|
||||
audio_mask = moving_average(voice_flags, vad_moving_average_width)
|
||||
audio_mask = np.round(audio_mask).astype(np.bool)
|
||||
|
||||
# Dilate the voiced regions
|
||||
audio_mask = binary_dilation(audio_mask, np.ones(vad_max_silence_length + 1))
|
||||
audio_mask = np.repeat(audio_mask, samples_per_window)
|
||||
audio_mask = resize(audio_mask, (len(wav_raw),)) > 0
|
||||
if return_raw_wav:
|
||||
return wav_raw, audio_mask, sr
|
||||
return wav_raw[audio_mask], audio_mask, sr
|
||||
@@ -0,0 +1,235 @@
|
||||
import logging
|
||||
import os
|
||||
import random
|
||||
import subprocess
|
||||
import sys
|
||||
from datetime import datetime
|
||||
import numpy as np
|
||||
import torch.utils.data
|
||||
from torch import nn
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
from .dataset_utils import data_loader
|
||||
from .hparams import hparams
|
||||
from .meters import AvgrageMeter
|
||||
from .tensor_utils import tensors_to_scalars
|
||||
from .trainer import Trainer
|
||||
|
||||
torch.multiprocessing.set_sharing_strategy(os.getenv('TORCH_SHARE_STRATEGY', 'file_system'))
|
||||
|
||||
log_format = '%(asctime)s %(message)s'
|
||||
logging.basicConfig(stream=sys.stdout, level=logging.INFO,
|
||||
format=log_format, datefmt='%m/%d %I:%M:%S %p')
|
||||
|
||||
|
||||
class BaseTask(nn.Module):
|
||||
def __init__(self, *args, **kwargs):
|
||||
super(BaseTask, self).__init__()
|
||||
self.current_epoch = 0
|
||||
self.global_step = 0
|
||||
self.trainer = None
|
||||
self.use_ddp = False
|
||||
self.gradient_clip_norm = hparams['clip_grad_norm']
|
||||
self.gradient_clip_val = hparams.get('clip_grad_value', 0)
|
||||
self.model = None
|
||||
self.training_losses_meter = None
|
||||
self.logger: SummaryWriter = None
|
||||
|
||||
######################
|
||||
# build model, dataloaders, optimizer, scheduler and tensorboard
|
||||
######################
|
||||
def build_model(self):
|
||||
raise NotImplementedError
|
||||
|
||||
@data_loader
|
||||
def train_dataloader(self):
|
||||
raise NotImplementedError
|
||||
|
||||
@data_loader
|
||||
def test_dataloader(self):
|
||||
raise NotImplementedError
|
||||
|
||||
@data_loader
|
||||
def val_dataloader(self):
|
||||
raise NotImplementedError
|
||||
|
||||
def build_scheduler(self, optimizer):
|
||||
return None
|
||||
|
||||
def build_optimizer(self, model):
|
||||
raise NotImplementedError
|
||||
|
||||
def configure_optimizers(self):
|
||||
optm = self.build_optimizer(self.model)
|
||||
self.scheduler = self.build_scheduler(optm)
|
||||
if isinstance(optm, (list, tuple)):
|
||||
return optm
|
||||
return [optm]
|
||||
|
||||
def build_tensorboard(self, save_dir, name, **kwargs):
|
||||
log_dir = os.path.join(save_dir, name)
|
||||
os.makedirs(log_dir, exist_ok=True)
|
||||
self.logger = SummaryWriter(log_dir=log_dir, **kwargs)
|
||||
|
||||
######################
|
||||
# training
|
||||
######################
|
||||
def on_train_start(self):
|
||||
pass
|
||||
|
||||
def on_train_end(self):
|
||||
pass
|
||||
|
||||
def on_epoch_start(self):
|
||||
self.training_losses_meter = {'total_loss': AvgrageMeter()}
|
||||
|
||||
def on_epoch_end(self):
|
||||
loss_outputs = {k: round(v.avg, 4) for k, v in self.training_losses_meter.items()}
|
||||
print(f"Epoch {self.current_epoch} ended. Steps: {self.global_step}. {loss_outputs}")
|
||||
|
||||
def _training_step(self, sample, batch_idx, optimizer_idx):
|
||||
"""
|
||||
|
||||
:param sample:
|
||||
:param batch_idx:
|
||||
:return: total loss: torch.Tensor, loss_log: dict
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def training_step(self, sample, batch_idx, optimizer_idx=-1):
|
||||
"""
|
||||
|
||||
:param sample:
|
||||
:param batch_idx:
|
||||
:param optimizer_idx:
|
||||
:return: {'loss': torch.Tensor, 'progress_bar': dict, 'tb_log': dict}
|
||||
"""
|
||||
loss_ret = self._training_step(sample, batch_idx, optimizer_idx)
|
||||
if loss_ret is None:
|
||||
return {'loss': None}
|
||||
total_loss, log_outputs = loss_ret
|
||||
log_outputs = tensors_to_scalars(log_outputs)
|
||||
for k, v in log_outputs.items():
|
||||
if k not in self.training_losses_meter:
|
||||
self.training_losses_meter[k] = AvgrageMeter()
|
||||
if not np.isnan(v):
|
||||
self.training_losses_meter[k].update(v)
|
||||
self.training_losses_meter['total_loss'].update(total_loss.item())
|
||||
|
||||
if optimizer_idx >= 0:
|
||||
log_outputs[f'lr_{optimizer_idx}'] = self.trainer.optimizers[optimizer_idx].param_groups[0]['lr']
|
||||
|
||||
progress_bar_log = log_outputs
|
||||
tb_log = {f'tr/{k}': v for k, v in log_outputs.items()}
|
||||
return {
|
||||
'loss': total_loss,
|
||||
'progress_bar': progress_bar_log,
|
||||
'tb_log': tb_log
|
||||
}
|
||||
|
||||
def on_before_optimization(self, opt_idx):
|
||||
if self.gradient_clip_norm > 0:
|
||||
torch.nn.utils.clip_grad_norm_(self.parameters(), self.gradient_clip_norm)
|
||||
if self.gradient_clip_val > 0:
|
||||
torch.nn.utils.clip_grad_value_(self.parameters(), self.gradient_clip_val)
|
||||
|
||||
def on_after_optimization(self, epoch, batch_idx, optimizer, optimizer_idx):
|
||||
if self.scheduler is not None:
|
||||
# self.scheduler.step(self.global_step // hparams['accumulate_grad_batches'])
|
||||
# the code above causes EPOCH_DEPRECATION_WARNING, changed it and changed the optimizer init with
|
||||
# step_size divided by accumulate_grad_batches
|
||||
self.scheduler.step()
|
||||
|
||||
######################
|
||||
# validation
|
||||
######################
|
||||
def validation_start(self):
|
||||
pass
|
||||
|
||||
def validation_step(self, sample, batch_idx):
|
||||
"""
|
||||
|
||||
:param sample:
|
||||
:param batch_idx:
|
||||
:return: output: {"losses": {...}, "total_loss": float, ...} or (total loss: torch.Tensor, loss_log: dict)
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def validation_end(self, outputs):
|
||||
"""
|
||||
|
||||
:param outputs:
|
||||
:return: loss_output: dict
|
||||
"""
|
||||
all_losses_meter = {'total_loss': AvgrageMeter()}
|
||||
for output in outputs:
|
||||
if len(output) == 0 or output is None:
|
||||
continue
|
||||
if isinstance(output, dict):
|
||||
assert 'losses' in output, 'Key "losses" should exist in validation output.'
|
||||
n = output.pop('nsamples', 1)
|
||||
losses = tensors_to_scalars(output['losses'])
|
||||
total_loss = output.get('total_loss', sum(losses.values()))
|
||||
else:
|
||||
assert len(output) == 2, 'Validation output should only consist of two elements: (total_loss, losses)'
|
||||
n = 1
|
||||
total_loss, losses = output
|
||||
losses = tensors_to_scalars(losses)
|
||||
if isinstance(total_loss, torch.Tensor):
|
||||
total_loss = total_loss.item()
|
||||
for k, v in losses.items():
|
||||
if k not in all_losses_meter:
|
||||
all_losses_meter[k] = AvgrageMeter()
|
||||
all_losses_meter[k].update(v, n)
|
||||
all_losses_meter['total_loss'].update(total_loss, n)
|
||||
loss_output = {k: round(v.avg, 4) for k, v in all_losses_meter.items()}
|
||||
print(f"| Validation results@{self.global_step}: {loss_output}")
|
||||
return {
|
||||
'tb_log': {f'val/{k}': v for k, v in loss_output.items()},
|
||||
'val_loss': loss_output['total_loss']
|
||||
}
|
||||
|
||||
######################
|
||||
# testing
|
||||
######################
|
||||
def test_start(self):
|
||||
pass
|
||||
|
||||
def test_step(self, sample, batch_idx):
|
||||
return self.validation_step(sample, batch_idx)
|
||||
|
||||
def test_end(self, outputs):
|
||||
return self.validation_end(outputs)
|
||||
|
||||
######################
|
||||
# start training/testing
|
||||
######################
|
||||
@classmethod
|
||||
def start(cls):
|
||||
os.environ['MASTER_PORT'] = str(random.randint(15000, 30000))
|
||||
random.seed(hparams['seed'])
|
||||
np.random.seed(hparams['seed'])
|
||||
work_dir = hparams['work_dir']
|
||||
trainer = Trainer(
|
||||
work_dir=work_dir,
|
||||
val_check_interval=hparams['val_check_interval'],
|
||||
tb_log_interval=hparams['tb_log_interval'],
|
||||
max_updates=hparams['max_updates'],
|
||||
num_sanity_val_steps=hparams['num_sanity_val_steps'] if not hparams['validate'] else 10000,
|
||||
accumulate_grad_batches=hparams['accumulate_grad_batches'],
|
||||
print_nan_grads=hparams['print_nan_grads'],
|
||||
resume_from_checkpoint=hparams.get('resume_from_checkpoint', 0),
|
||||
amp=hparams['amp'],
|
||||
monitor_key=hparams['valid_monitor_key'],
|
||||
monitor_mode=hparams['valid_monitor_mode'],
|
||||
num_ckpt_keep=hparams['num_ckpt_keep'],
|
||||
save_best=hparams['save_best'],
|
||||
seed=hparams['seed'],
|
||||
debug=hparams['debug']
|
||||
)
|
||||
if not hparams['infer']: # train
|
||||
trainer.fit(cls)
|
||||
else:
|
||||
trainer.test(cls)
|
||||
|
||||
def on_keyboard_interrupt(self):
|
||||
pass
|
||||
@@ -0,0 +1,68 @@
|
||||
import glob
|
||||
import os
|
||||
import re
|
||||
import torch
|
||||
|
||||
|
||||
def get_last_checkpoint(work_dir, steps=None):
|
||||
checkpoint = None
|
||||
last_ckpt_path = None
|
||||
ckpt_paths = get_all_ckpts(work_dir, steps)
|
||||
if len(ckpt_paths) > 0:
|
||||
last_ckpt_path = ckpt_paths[0]
|
||||
checkpoint = torch.load(last_ckpt_path, map_location='cpu')
|
||||
return checkpoint, last_ckpt_path
|
||||
|
||||
|
||||
def get_all_ckpts(work_dir, steps=None):
|
||||
if steps is None:
|
||||
ckpt_path_pattern = f'{work_dir}/model_ckpt_steps_*.ckpt'
|
||||
else:
|
||||
ckpt_path_pattern = f'{work_dir}/model_ckpt_steps_{steps}.ckpt'
|
||||
return sorted(glob.glob(ckpt_path_pattern),
|
||||
key=lambda x: -int(re.findall('.*steps\_(\d+)\.ckpt', x)[0]))
|
||||
|
||||
|
||||
def load_ckpt(cur_model, ckpt_base_dir, model_name='model', force=True, strict=True, verbose=True):
|
||||
if os.path.isfile(ckpt_base_dir):
|
||||
base_dir = os.path.dirname(ckpt_base_dir)
|
||||
ckpt_path = ckpt_base_dir
|
||||
checkpoint = torch.load(ckpt_base_dir, map_location='cpu')
|
||||
else:
|
||||
base_dir = ckpt_base_dir
|
||||
checkpoint, ckpt_path = get_last_checkpoint(ckpt_base_dir)
|
||||
if checkpoint is not None:
|
||||
state_dict = checkpoint["state_dict"]
|
||||
if len([k for k in state_dict.keys() if '.' in k]) > 0:
|
||||
state_dict = {k[len(model_name) + 1:]: v for k, v in state_dict.items()
|
||||
if k.startswith(f'{model_name}.')}
|
||||
else:
|
||||
if '.' not in model_name:
|
||||
state_dict = state_dict[model_name]
|
||||
else:
|
||||
base_model_name = model_name.split('.')[0]
|
||||
rest_model_name = model_name[len(base_model_name) + 1:]
|
||||
state_dict = {
|
||||
k[len(rest_model_name) + 1:]: v for k, v in state_dict[base_model_name].items()
|
||||
if k.startswith(f'{rest_model_name}.')}
|
||||
if not strict:
|
||||
cur_model_state_dict = cur_model.state_dict()
|
||||
unmatched_keys = []
|
||||
for key, param in state_dict.items():
|
||||
if key in cur_model_state_dict:
|
||||
new_param = cur_model_state_dict[key]
|
||||
if new_param.shape != param.shape:
|
||||
unmatched_keys.append(key)
|
||||
print("| Unmatched keys: ", key, new_param.shape, param.shape)
|
||||
for key in unmatched_keys:
|
||||
del state_dict[key]
|
||||
# print(state_dict)
|
||||
cur_model.load_state_dict(state_dict, strict=strict)
|
||||
if verbose:
|
||||
print(f"| load '{model_name}' from '{ckpt_path}'.")
|
||||
else:
|
||||
e_msg = f"| ckpt not found in {base_dir}."
|
||||
if force:
|
||||
assert False, e_msg
|
||||
else:
|
||||
print(e_msg)
|
||||
@@ -0,0 +1,372 @@
|
||||
import os
|
||||
import sys
|
||||
import traceback
|
||||
import types
|
||||
from functools import wraps
|
||||
from itertools import chain
|
||||
import numpy as np
|
||||
import torch.utils.data
|
||||
import torch.nn.functional as F
|
||||
from torch.utils.data import ConcatDataset
|
||||
from .hparams import hparams
|
||||
|
||||
|
||||
def collate_1d_or_2d(values, pad_idx=0, left_pad=False, shift_right=False, max_len=None, shift_id=1):
|
||||
if len(values[0].shape) == 1:
|
||||
return collate_1d(values, pad_idx, left_pad, shift_right, max_len, shift_id)
|
||||
else:
|
||||
return collate_2d(values, pad_idx, left_pad, shift_right, max_len)
|
||||
|
||||
|
||||
def collate_1d(values, pad_idx=0, left_pad=False, shift_right=False, max_len=None, shift_id=1):
|
||||
"""Convert a list of 1d tensors into a padded 2d tensor."""
|
||||
size = max(v.size(0) for v in values) if max_len is None else max_len
|
||||
res = values[0].new(len(values), size).fill_(pad_idx)
|
||||
|
||||
def copy_tensor(src, dst):
|
||||
assert dst.numel() == src.numel()
|
||||
if shift_right:
|
||||
dst[1:] = src[:-1]
|
||||
dst[0] = shift_id
|
||||
else:
|
||||
dst.copy_(src)
|
||||
|
||||
for i, v in enumerate(values):
|
||||
copy_tensor(v, res[i][size - len(v):] if left_pad else res[i][:len(v)])
|
||||
return res
|
||||
|
||||
|
||||
def collate_2d(values, pad_idx=0, left_pad=False, shift_right=False, max_len=None):
|
||||
"""Convert a list of 2d tensors into a padded 3d tensor."""
|
||||
size = max(v.size(0) for v in values) if max_len is None else max_len
|
||||
res = values[0].new(len(values), size, values[0].shape[1]).fill_(pad_idx)
|
||||
|
||||
def copy_tensor(src, dst):
|
||||
assert dst.numel() == src.numel()
|
||||
if shift_right:
|
||||
dst[1:] = src[:-1]
|
||||
else:
|
||||
dst.copy_(src)
|
||||
|
||||
for i, v in enumerate(values):
|
||||
copy_tensor(v, res[i][size - len(v):] if left_pad else res[i][:len(v)])
|
||||
return res
|
||||
|
||||
def collate_xd(values, pad_value=0, max_len=None):
|
||||
size = ((max(v.size(0) for v in values) if max_len is None else max_len), *values[0].shape[1:])
|
||||
res = torch.full((len(values), *size), fill_value=pad_value, dtype=values[0].dtype, device=values[0].device)
|
||||
|
||||
for i, v in enumerate(values):
|
||||
res[i, :len(v), ...] = v
|
||||
return res
|
||||
|
||||
def pad_or_cut_1d(values: torch.tensor, tgt_len, pad_value=0):
|
||||
src_len = values.shape[0]
|
||||
if src_len < tgt_len:
|
||||
res = F.pad(values, [0, tgt_len - src_len], value=pad_value)
|
||||
else:
|
||||
res = values[:tgt_len]
|
||||
return res
|
||||
|
||||
def pad_or_cut_2d(values: torch.tensor, tgt_len, dim=-1, pad_value=0):
|
||||
if dim == 0 or dim == -2:
|
||||
src_len = values.shape[0]
|
||||
if src_len < tgt_len:
|
||||
res = F.pad(values, [0, 0, 0, tgt_len - src_len], value=pad_value)
|
||||
else:
|
||||
res = values[:tgt_len]
|
||||
elif dim == 1 or dim == -1:
|
||||
src_len = values.shape[1]
|
||||
if src_len < tgt_len:
|
||||
res = F.pad(values, [0, tgt_len - src_len], value=pad_value)
|
||||
else:
|
||||
res = values[:, :tgt_len]
|
||||
else:
|
||||
raise RuntimeError(f"Wrong dim number {dim} while the tensor only has {len(values.shape)} dimensions.")
|
||||
return res
|
||||
|
||||
def pad_or_cut_3d(values: torch.tensor, tgt_len, dim=-1, pad_value=0):
|
||||
if dim == 0 or dim == -3:
|
||||
src_len = values.shape[0]
|
||||
if src_len < tgt_len:
|
||||
res = F.pad(values, [0, 0, 0, 0, 0, tgt_len - src_len], value=pad_value)
|
||||
else:
|
||||
res = values[:tgt_len]
|
||||
elif dim == 1 or dim == -2:
|
||||
src_len = values.shape[1]
|
||||
if src_len < tgt_len:
|
||||
res = F.pad(values, [0, 0, 0, tgt_len - src_len], value=pad_value)
|
||||
else:
|
||||
res = values[:, :tgt_len]
|
||||
elif dim == 2 or dim == -1:
|
||||
src_len = values.shape[2]
|
||||
if src_len < tgt_len:
|
||||
res = F.pad(values, [0, tgt_len - src_len], value=pad_value)
|
||||
else:
|
||||
res = values[:, :, :tgt_len]
|
||||
else:
|
||||
raise RuntimeError(f"Wrong dim number {dim} while the tensor only has {len(values.shape)} dimensions.")
|
||||
return res
|
||||
|
||||
def pad_or_cut_xd(values, tgt_len, dim=-1, pad_value=0):
|
||||
if len(values.shape) == 1:
|
||||
return pad_or_cut_1d(values, tgt_len, pad_value)
|
||||
elif len(values.shape) == 2:
|
||||
return pad_or_cut_2d(values, tgt_len, dim, pad_value)
|
||||
elif len(values.shape) == 3:
|
||||
return pad_or_cut_3d(values, tgt_len, dim, pad_value)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
def _is_batch_full(batch, num_tokens, max_tokens, max_sentences):
|
||||
if len(batch) == 0:
|
||||
return 0
|
||||
if len(batch) == max_sentences:
|
||||
return 1
|
||||
if num_tokens > max_tokens:
|
||||
return 1
|
||||
return 0
|
||||
|
||||
|
||||
def batch_by_size(
|
||||
indices, num_tokens_fn, max_tokens=None, max_sentences=None,
|
||||
required_batch_size_multiple=1, distributed=False
|
||||
):
|
||||
"""
|
||||
Yield mini-batches of indices bucketed by size. Batches may contain
|
||||
sequences of different lengths.
|
||||
|
||||
Args:
|
||||
indices (List[int]): ordered list of dataset indices
|
||||
num_tokens_fn (callable): function that returns the number of tokens at
|
||||
a given index
|
||||
max_tokens (int, optional): max number of tokens in each batch
|
||||
(default: None).
|
||||
max_sentences (int, optional): max number of sentences in each
|
||||
batch (default: None).
|
||||
required_batch_size_multiple (int, optional): require batch size to
|
||||
be a multiple of N (default: 1).
|
||||
"""
|
||||
max_tokens = max_tokens if max_tokens is not None else sys.maxsize
|
||||
max_sentences = max_sentences if max_sentences is not None else sys.maxsize
|
||||
bsz_mult = required_batch_size_multiple
|
||||
|
||||
if isinstance(indices, types.GeneratorType):
|
||||
indices = np.fromiter(indices, dtype=np.int64, count=-1)
|
||||
|
||||
sample_len = 0
|
||||
sample_lens = []
|
||||
batch = []
|
||||
batches = []
|
||||
for i in range(len(indices)):
|
||||
idx = indices[i]
|
||||
num_tokens = num_tokens_fn(idx)
|
||||
sample_lens.append(num_tokens)
|
||||
sample_len = max(sample_len, num_tokens)
|
||||
|
||||
assert sample_len <= max_tokens, (
|
||||
"sentence at index {} of size {} exceeds max_tokens "
|
||||
"limit of {}!".format(idx, sample_len, max_tokens)
|
||||
)
|
||||
num_tokens = (len(batch) + 1) * sample_len
|
||||
|
||||
if _is_batch_full(batch, num_tokens, max_tokens, max_sentences):
|
||||
mod_len = max(
|
||||
bsz_mult * (len(batch) // bsz_mult),
|
||||
len(batch) % bsz_mult,
|
||||
)
|
||||
batches.append(batch[:mod_len])
|
||||
batch = batch[mod_len:]
|
||||
sample_lens = sample_lens[mod_len:]
|
||||
sample_len = max(sample_lens) if len(sample_lens) > 0 else 0
|
||||
batch.append(idx)
|
||||
if len(batch) > 0:
|
||||
batches.append(batch)
|
||||
return batches
|
||||
|
||||
|
||||
def build_dataloader(dataset, shuffle, max_tokens=None, max_sentences=None,
|
||||
required_batch_size_multiple=-1, endless=False, apply_batch_by_size=True, pin_memory=False, use_ddp=False):
|
||||
import torch.distributed as dist
|
||||
devices_cnt = torch.cuda.device_count()
|
||||
if devices_cnt == 0:
|
||||
devices_cnt = 1
|
||||
if not use_ddp:
|
||||
devices_cnt = 1
|
||||
if required_batch_size_multiple == -1:
|
||||
required_batch_size_multiple = devices_cnt
|
||||
|
||||
def shuffle_batches(batches):
|
||||
np.random.shuffle(batches)
|
||||
return batches
|
||||
|
||||
if max_tokens is not None:
|
||||
max_tokens *= devices_cnt
|
||||
if max_sentences is not None:
|
||||
max_sentences *= devices_cnt
|
||||
indices = dataset.ordered_indices()
|
||||
if apply_batch_by_size:
|
||||
batch_sampler = batch_by_size(
|
||||
indices, dataset.num_tokens, max_tokens=max_tokens, max_sentences=max_sentences,
|
||||
required_batch_size_multiple=required_batch_size_multiple,
|
||||
)
|
||||
else:
|
||||
batch_sampler = []
|
||||
for i in range(0, len(indices), max_sentences):
|
||||
batch_sampler.append(indices[i:i + max_sentences])
|
||||
|
||||
if shuffle:
|
||||
batches = shuffle_batches(list(batch_sampler))
|
||||
if endless:
|
||||
batches = [b for _ in range(1000) for b in shuffle_batches(list(batch_sampler))]
|
||||
else:
|
||||
batches = batch_sampler
|
||||
if endless:
|
||||
batches = [b for _ in range(1000) for b in batches]
|
||||
num_workers = dataset.num_workers
|
||||
if use_ddp:
|
||||
num_replicas = dist.get_world_size()
|
||||
rank = dist.get_rank()
|
||||
# batches = [x[rank::num_replicas] for x in batches if len(x) % num_replicas == 0]
|
||||
# ensure that every sample in the dataset is covered
|
||||
batches_ = []
|
||||
for x in batches:
|
||||
if len(x) % num_replicas == 0:
|
||||
batches_.append(x[rank::num_replicas])
|
||||
else:
|
||||
x_ = x + [x[-1]] * (len(x) - len(x) // num_replicas * num_replicas)
|
||||
batches_.append(x_[rank::num_replicas])
|
||||
batches = batches_
|
||||
return torch.utils.data.DataLoader(dataset,
|
||||
collate_fn=dataset.collater,
|
||||
batch_sampler=batches,
|
||||
num_workers=num_workers,
|
||||
pin_memory=pin_memory)
|
||||
|
||||
|
||||
def unpack_dict_to_list(samples):
|
||||
samples_ = []
|
||||
bsz = samples.get('outputs').size(0)
|
||||
for i in range(bsz):
|
||||
res = {}
|
||||
for k, v in samples.items():
|
||||
try:
|
||||
res[k] = v[i]
|
||||
except:
|
||||
pass
|
||||
samples_.append(res)
|
||||
return samples_
|
||||
|
||||
|
||||
def remove_padding(x, padding_idx=0):
|
||||
if x is None:
|
||||
return None
|
||||
assert len(x.shape) in [1, 2]
|
||||
if len(x.shape) == 2: # [T, H]
|
||||
return x[np.abs(x).sum(-1) != padding_idx]
|
||||
elif len(x.shape) == 1: # [T]
|
||||
return x[x != padding_idx]
|
||||
|
||||
|
||||
def data_loader(fn):
|
||||
"""
|
||||
Decorator to make any fx with this use the lazy property
|
||||
:param fn:
|
||||
:return:
|
||||
"""
|
||||
|
||||
wraps(fn)
|
||||
attr_name = '_lazy_' + fn.__name__
|
||||
|
||||
def _get_data_loader(self):
|
||||
try:
|
||||
value = getattr(self, attr_name)
|
||||
except AttributeError:
|
||||
try:
|
||||
value = fn(self) # Lazy evaluation, done only once.
|
||||
except AttributeError as e:
|
||||
# Guard against AttributeError suppression. (Issue #142)
|
||||
traceback.print_exc()
|
||||
error = f'{fn.__name__}: An AttributeError was encountered: ' + str(e)
|
||||
raise RuntimeError(error) from e
|
||||
setattr(self, attr_name, value) # Memoize evaluation.
|
||||
return value
|
||||
|
||||
return _get_data_loader
|
||||
|
||||
|
||||
class BaseDataset(torch.utils.data.Dataset):
|
||||
def __init__(self, shuffle):
|
||||
super().__init__()
|
||||
self.hparams = hparams
|
||||
self.shuffle = shuffle
|
||||
self.sort_by_len = hparams['sort_by_len']
|
||||
self.sizes = None
|
||||
|
||||
@property
|
||||
def _sizes(self):
|
||||
return self.sizes
|
||||
|
||||
def __getitem__(self, index):
|
||||
raise NotImplementedError
|
||||
|
||||
def collater(self, samples):
|
||||
raise NotImplementedError
|
||||
|
||||
def __len__(self):
|
||||
return len(self._sizes)
|
||||
|
||||
def num_tokens(self, index):
|
||||
return self.size(index)
|
||||
|
||||
def size(self, index):
|
||||
"""Return an example's size as a float or tuple. This value is used when
|
||||
filtering a dataset with ``--max-positions``."""
|
||||
return min(self._sizes[index], hparams['max_frames'])
|
||||
|
||||
def ordered_indices(self):
|
||||
"""Return an ordered list of indices. Batches will be constructed based
|
||||
on this order."""
|
||||
if self.shuffle:
|
||||
indices = np.random.permutation(len(self))
|
||||
if self.sort_by_len:
|
||||
indices = indices[np.argsort(np.array(self._sizes)[indices], kind='mergesort')]
|
||||
else:
|
||||
indices = np.arange(len(self))
|
||||
return indices.tolist()
|
||||
|
||||
@property
|
||||
def num_workers(self):
|
||||
return int(os.getenv('NUM_WORKERS', hparams['ds_workers']))
|
||||
|
||||
|
||||
class BaseConcatDataset(ConcatDataset):
|
||||
def collater(self, samples):
|
||||
return self.datasets[0].collater(samples)
|
||||
|
||||
@property
|
||||
def _sizes(self):
|
||||
if not hasattr(self, 'sizes'):
|
||||
self.sizes = list(chain.from_iterable([d._sizes for d in self.datasets]))
|
||||
return self.sizes
|
||||
|
||||
def size(self, index):
|
||||
return min(self._sizes[index], hparams['max_frames'])
|
||||
|
||||
def num_tokens(self, index):
|
||||
return self.size(index)
|
||||
|
||||
def ordered_indices(self):
|
||||
"""Return an ordered list of indices. Batches will be constructed based
|
||||
on this order."""
|
||||
if self.datasets[0].shuffle:
|
||||
indices = np.random.permutation(len(self))
|
||||
if self.datasets[0].sort_by_len:
|
||||
indices = indices[np.argsort(np.array(self._sizes)[indices], kind='mergesort')]
|
||||
else:
|
||||
indices = np.arange(len(self))
|
||||
return indices
|
||||
|
||||
@property
|
||||
def num_workers(self):
|
||||
return self.datasets[0].num_workers
|
||||
@@ -0,0 +1,164 @@
|
||||
from torch.nn.parallel import DistributedDataParallel
|
||||
from torch.nn.parallel.distributed import _find_tensors
|
||||
import torch.optim
|
||||
import torch.utils.data
|
||||
import torch
|
||||
from packaging import version
|
||||
|
||||
class DDP(DistributedDataParallel):
|
||||
"""
|
||||
Override the forward call in lightning so it goes to training and validation step respectively
|
||||
"""
|
||||
|
||||
def forward(self, *inputs, **kwargs): # pragma: no cover
|
||||
# if version.parse(torch.__version__[:6]) < version.parse("1.11"):
|
||||
if version.parse(torch.__version__) < version.parse("1.11"): # fix the hard [:6] problem
|
||||
self._sync_params()
|
||||
inputs, kwargs = self.scatter(inputs, kwargs, self.device_ids)
|
||||
assert len(self.device_ids) == 1
|
||||
if self.module.training:
|
||||
output = self.module.training_step(*inputs[0], **kwargs[0])
|
||||
elif self.module.testing:
|
||||
output = self.module.test_step(*inputs[0], **kwargs[0])
|
||||
else:
|
||||
output = self.module.validation_step(*inputs[0], **kwargs[0])
|
||||
if torch.is_grad_enabled():
|
||||
# We'll return the output object verbatim since it is a freeform
|
||||
# object. We need to find any tensors in this object, though,
|
||||
# because we need to figure out which parameters were used during
|
||||
# this forward pass, to ensure we short circuit reduction for any
|
||||
# unused parameters. Only if `find_unused_parameters` is set.
|
||||
if self.find_unused_parameters:
|
||||
self.reducer.prepare_for_backward(list(_find_tensors(output)))
|
||||
else:
|
||||
self.reducer.prepare_for_backward([])
|
||||
elif version.parse("1.11") <= version.parse(torch.__version__) < version.parse("2.0"):
|
||||
from torch.nn.parallel.distributed import \
|
||||
logging, Join, _DDPSink, _tree_flatten_with_rref, _tree_unflatten_with_rref
|
||||
with torch.autograd.profiler.record_function("DistributedDataParallel.forward"):
|
||||
if torch.is_grad_enabled() and self.require_backward_grad_sync:
|
||||
self.logger.set_runtime_stats_and_log()
|
||||
self.num_iterations += 1
|
||||
self.reducer.prepare_for_forward()
|
||||
|
||||
# Notify the join context that this process has not joined, if
|
||||
# needed
|
||||
work = Join.notify_join_context(self)
|
||||
if work:
|
||||
self.reducer._set_forward_pass_work_handle(
|
||||
work, self._divide_by_initial_world_size
|
||||
)
|
||||
|
||||
# Calling _rebuild_buckets before forward compuation,
|
||||
# It may allocate new buckets before deallocating old buckets
|
||||
# inside _rebuild_buckets. To save peak memory usage,
|
||||
# call _rebuild_buckets before the peak memory usage increases
|
||||
# during forward computation.
|
||||
# This should be called only once during whole training period.
|
||||
if torch.is_grad_enabled() and self.reducer._rebuild_buckets():
|
||||
logging.info("Reducer buckets have been rebuilt in this iteration.")
|
||||
self._has_rebuilt_buckets = True
|
||||
|
||||
# sync params according to location (before/after forward) user
|
||||
# specified as part of hook, if hook was specified.
|
||||
buffer_hook_registered = hasattr(self, 'buffer_hook')
|
||||
if self._check_sync_bufs_pre_fwd():
|
||||
self._sync_buffers()
|
||||
|
||||
if self._join_config.enable:
|
||||
# Notify joined ranks whether they should sync in backwards pass or not.
|
||||
self._check_global_requires_backward_grad_sync(is_joined_rank=False)
|
||||
|
||||
# modified part
|
||||
inputs, kwargs = self.scatter(inputs, kwargs, self.device_ids)
|
||||
if self.module.training:
|
||||
output = self.module.training_step(*inputs[0], **kwargs[0])
|
||||
elif self.module.testing:
|
||||
output = self.module.test_step(*inputs[0], **kwargs[0])
|
||||
else:
|
||||
output = self.module.validation_step(*inputs[0], **kwargs[0])
|
||||
|
||||
# sync params according to location (before/after forward) user
|
||||
# specified as part of hook, if hook was specified.
|
||||
if self._check_sync_bufs_post_fwd():
|
||||
self._sync_buffers()
|
||||
|
||||
if torch.is_grad_enabled() and self.require_backward_grad_sync:
|
||||
self.require_forward_param_sync = True
|
||||
# We'll return the output object verbatim since it is a freeform
|
||||
# object. We need to find any tensors in this object, though,
|
||||
# because we need to figure out which parameters were used during
|
||||
# this forward pass, to ensure we short circuit reduction for any
|
||||
# unused parameters. Only if `find_unused_parameters` is set.
|
||||
if self.find_unused_parameters and not self.static_graph:
|
||||
# Do not need to populate this for static graph.
|
||||
self.reducer.prepare_for_backward(list(_find_tensors(output)))
|
||||
else:
|
||||
self.reducer.prepare_for_backward([])
|
||||
else:
|
||||
self.require_forward_param_sync = False
|
||||
|
||||
# TODO: DDPSink is currently enabled for unused parameter detection and
|
||||
# static graph training for first iteration.
|
||||
if (self.find_unused_parameters and not self.static_graph) or (
|
||||
self.static_graph and self.num_iterations == 1
|
||||
):
|
||||
state_dict = {
|
||||
'static_graph': self.static_graph,
|
||||
'num_iterations': self.num_iterations,
|
||||
}
|
||||
|
||||
output_tensor_list, treespec, output_is_rref = _tree_flatten_with_rref(
|
||||
output
|
||||
)
|
||||
output_placeholders = [None for _ in range(len(output_tensor_list))]
|
||||
# Do not touch tensors that have no grad_fn, which can cause issues
|
||||
# such as https://github.com/pytorch/pytorch/issues/60733
|
||||
for i, output in enumerate(output_tensor_list):
|
||||
if torch.is_tensor(output) and output.grad_fn is None:
|
||||
output_placeholders[i] = output
|
||||
|
||||
# When find_unused_parameters=True, makes tensors which require grad
|
||||
# run through the DDPSink backward pass. When not all outputs are
|
||||
# used in loss, this makes those corresponding tensors receive
|
||||
# undefined gradient which the reducer then handles to ensure
|
||||
# param.grad field is not touched and we don't error out.
|
||||
passthrough_tensor_list = _DDPSink.apply(
|
||||
self.reducer,
|
||||
state_dict,
|
||||
*output_tensor_list,
|
||||
)
|
||||
for i in range(len(output_placeholders)):
|
||||
if output_placeholders[i] is None:
|
||||
output_placeholders[i] = passthrough_tensor_list[i]
|
||||
|
||||
# Reconstruct output data structure.
|
||||
output = _tree_unflatten_with_rref(
|
||||
output_placeholders, treespec, output_is_rref
|
||||
)
|
||||
else:
|
||||
# now pytorch version >= 2.0
|
||||
with torch.autograd.profiler.record_function("DistributedDataParallel.forward"):
|
||||
inputs, kwargs = self._pre_forward(*inputs, **kwargs)
|
||||
output = (
|
||||
# self.module.forward(*inputs, **kwargs)
|
||||
# if self._delay_all_reduce_all_params
|
||||
# else self._run_ddp_forward(*inputs, **kwargs)
|
||||
# modified: delete 'delay_all_reduce_named_params' function
|
||||
self._run_ddp_forward(*inputs, **kwargs)
|
||||
)
|
||||
return self._post_forward(output)
|
||||
return output
|
||||
|
||||
def _run_ddp_forward(self, *inputs, **kwargs):
|
||||
if version.parse(torch.__version__) >= version.parse("2.0"):
|
||||
with self._inside_ddp_forward():
|
||||
if self.module.training:
|
||||
output = self.module.training_step(*inputs, **kwargs)
|
||||
elif self.module.testing:
|
||||
output = self.module.test_step(*inputs, **kwargs)
|
||||
else:
|
||||
output = self.module.validation_step(*inputs, **kwargs)
|
||||
return output # type: ignore[index]
|
||||
else:
|
||||
return super(DDP, self)._run_ddp_forward(*inputs, **kwargs)
|
||||
@@ -0,0 +1,113 @@
|
||||
import gc
|
||||
import datetime
|
||||
import inspect
|
||||
|
||||
import torch
|
||||
import numpy as np
|
||||
|
||||
dtype_memory_size_dict = {
|
||||
torch.float64: 64/8,
|
||||
torch.double: 64/8,
|
||||
torch.float32: 32/8,
|
||||
torch.float: 32/8,
|
||||
torch.float16: 16/8,
|
||||
torch.half: 16/8,
|
||||
torch.int64: 64/8,
|
||||
torch.long: 64/8,
|
||||
torch.int32: 32/8,
|
||||
torch.int: 32/8,
|
||||
torch.int16: 16/8,
|
||||
torch.short: 16/6,
|
||||
torch.uint8: 8/8,
|
||||
torch.int8: 8/8,
|
||||
}
|
||||
# compatibility of torch1.0
|
||||
if getattr(torch, "bfloat16", None) is not None:
|
||||
dtype_memory_size_dict[torch.bfloat16] = 16/8
|
||||
if getattr(torch, "bool", None) is not None:
|
||||
dtype_memory_size_dict[torch.bool] = 8/8 # pytorch use 1 byte for a bool, see https://github.com/pytorch/pytorch/issues/41571
|
||||
|
||||
def get_mem_space(x):
|
||||
try:
|
||||
ret = dtype_memory_size_dict[x]
|
||||
except KeyError:
|
||||
print(f"dtype {x} is not supported!")
|
||||
return ret
|
||||
|
||||
class MemTracker(object):
|
||||
"""
|
||||
Class used to track pytorch memory usage
|
||||
Arguments:
|
||||
detail(bool, default True): whether the function shows the detail gpu memory usage
|
||||
path(str): where to save log file
|
||||
verbose(bool, default False): whether show the trivial exception
|
||||
device(int): GPU number, default is 0
|
||||
"""
|
||||
def __init__(self, detail=True, path='', verbose=False, device=0):
|
||||
self.print_detail = detail
|
||||
self.last_tensor_sizes = set()
|
||||
self.gpu_profile_fn = path + f'{datetime.datetime.now():%d-%b-%y-%H:%M:%S}-gpu_mem_track.txt'
|
||||
self.verbose = verbose
|
||||
self.begin = True
|
||||
self.device = device
|
||||
|
||||
def get_tensors(self):
|
||||
for obj in gc.get_objects():
|
||||
try:
|
||||
if torch.is_tensor(obj) or (hasattr(obj, 'data') and torch.is_tensor(obj.data)):
|
||||
tensor = obj
|
||||
else:
|
||||
continue
|
||||
if tensor.is_cuda:
|
||||
yield tensor
|
||||
except Exception as e:
|
||||
if self.verbose:
|
||||
print('A trivial exception occured: {}'.format(e))
|
||||
|
||||
def get_tensor_usage(self):
|
||||
sizes = [np.prod(np.array(tensor.size())) * get_mem_space(tensor.dtype) for tensor in self.get_tensors()]
|
||||
return np.sum(sizes) / 1024**2
|
||||
|
||||
def get_allocate_usage(self):
|
||||
return torch.cuda.memory_allocated() / 1024**2
|
||||
|
||||
def clear_cache(self):
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
def print_all_gpu_tensor(self, file=None):
|
||||
for x in self.get_tensors():
|
||||
print(x.size(), x.dtype, np.prod(np.array(x.size()))*get_mem_space(x.dtype)/1024**2, file=file)
|
||||
|
||||
def track(self):
|
||||
"""
|
||||
Track the GPU memory usage
|
||||
"""
|
||||
frameinfo = inspect.stack()[1]
|
||||
where_str = frameinfo.filename + ' line ' + str(frameinfo.lineno) + ': ' + frameinfo.function
|
||||
|
||||
with open(self.gpu_profile_fn, 'a+') as f:
|
||||
|
||||
if self.begin:
|
||||
f.write(f"GPU Memory Track | {datetime.datetime.now():%d-%b-%y-%H:%M:%S} |"
|
||||
f" Total Tensor Used Memory:{self.get_tensor_usage():<7.1f}Mb"
|
||||
f" Total Allocated Memory:{self.get_allocate_usage():<7.1f}Mb\n\n")
|
||||
self.begin = False
|
||||
|
||||
if self.print_detail is True:
|
||||
ts_list = [(tensor.size(), tensor.dtype) for tensor in self.get_tensors()]
|
||||
new_tensor_sizes = {(type(x),
|
||||
tuple(x.size()),
|
||||
ts_list.count((x.size(), x.dtype)),
|
||||
np.prod(np.array(x.size()))*get_mem_space(x.dtype)/1024**2,
|
||||
x.dtype) for x in self.get_tensors()}
|
||||
for t, s, n, m, data_type in new_tensor_sizes - self.last_tensor_sizes:
|
||||
f.write(f'+ | {str(n)} * Size:{str(s):<20} | Memory: {str(m*n)[:6]} M | {str(t):<20} | {data_type}\n')
|
||||
for t, s, n, m, data_type in self.last_tensor_sizes - new_tensor_sizes:
|
||||
f.write(f'- | {str(n)} * Size:{str(s):<20} | Memory: {str(m*n)[:6]} M | {str(t):<20} | {data_type}\n')
|
||||
|
||||
self.last_tensor_sizes = new_tensor_sizes
|
||||
|
||||
f.write(f"\nAt {where_str:<50}"
|
||||
f" Total Tensor Used Memory:{self.get_tensor_usage():<7.1f}Mb"
|
||||
f" Total Allocated Memory:{self.get_allocate_usage():<7.1f}Mb\n\n")
|
||||
@@ -0,0 +1,131 @@
|
||||
import argparse
|
||||
import os
|
||||
import yaml
|
||||
|
||||
global_print_hparams = True
|
||||
hparams = {}
|
||||
|
||||
|
||||
class Args:
|
||||
def __init__(self, **kwargs):
|
||||
for k, v in kwargs.items():
|
||||
self.__setattr__(k, v)
|
||||
|
||||
|
||||
def override_config(old_config: dict, new_config: dict):
|
||||
for k, v in new_config.items():
|
||||
if isinstance(v, dict) and k in old_config:
|
||||
override_config(old_config[k], new_config[k])
|
||||
else:
|
||||
old_config[k] = v
|
||||
|
||||
|
||||
def set_hparams(config='', exp_name='', hparams_str='', print_hparams=True, global_hparams=True, root_dir=''):
|
||||
if config == '' and exp_name == '':
|
||||
parser = argparse.ArgumentParser(description='')
|
||||
parser.add_argument('--config', type=str, default='',
|
||||
help='location of the data corpus')
|
||||
parser.add_argument('--exp_name', type=str, default='', help='exp_name')
|
||||
parser.add_argument('-hp', '--hparams', type=str, default='',
|
||||
help='location of the data corpus')
|
||||
parser.add_argument('--infer', action='store_true', help='infer')
|
||||
parser.add_argument('--validate', action='store_true', help='validate')
|
||||
parser.add_argument('--reset', action='store_true', help='reset hparams')
|
||||
parser.add_argument('--remove', action='store_true', help='remove old ckpt')
|
||||
parser.add_argument('--debug', action='store_true', help='debug')
|
||||
parser.add_argument('--root_dir', type=str, default='', help='root directory of the project.')
|
||||
args, unknown = parser.parse_known_args()
|
||||
print("| Unknow hparams: ", unknown)
|
||||
else:
|
||||
args = Args(config=config, exp_name=exp_name, hparams=hparams_str,
|
||||
infer=False, validate=False, reset=False, debug=False, remove=False, root_dir=root_dir)
|
||||
global hparams
|
||||
assert args.config != '' or args.exp_name != ''
|
||||
root_dir = args.root_dir
|
||||
if args.config != '':
|
||||
assert os.path.exists(os.path.join(root_dir, args.config)), f'| Wrong config path! root_dir: {root_dir}, config_path: {args.config}'
|
||||
|
||||
config_chains = []
|
||||
loaded_config = set()
|
||||
|
||||
def load_config(config_fn):
|
||||
# deep first inheritance and avoid the second visit of one node
|
||||
if not os.path.exists(os.path.join(root_dir, config_fn)):
|
||||
return {}
|
||||
with open(os.path.join(root_dir, config_fn)) as f:
|
||||
hparams_ = yaml.safe_load(f)
|
||||
loaded_config.add(config_fn)
|
||||
if 'base_config' in hparams_:
|
||||
ret_hparams = {}
|
||||
if not isinstance(hparams_['base_config'], list):
|
||||
hparams_['base_config'] = [hparams_['base_config']]
|
||||
for c in hparams_['base_config']:
|
||||
if c.startswith('.'):
|
||||
c = f'{os.path.dirname(config_fn)}/{c}'
|
||||
c = os.path.normpath(c)
|
||||
if c not in loaded_config:
|
||||
override_config(ret_hparams, load_config(c))
|
||||
override_config(ret_hparams, hparams_)
|
||||
else:
|
||||
ret_hparams = hparams_
|
||||
config_chains.append(config_fn)
|
||||
return ret_hparams
|
||||
|
||||
saved_hparams = {}
|
||||
args_work_dir = ''
|
||||
if args.exp_name != '':
|
||||
args_work_dir = os.path.join(root_dir, f'checkpoints/{args.exp_name}')
|
||||
ckpt_config_path = f'{args_work_dir}/config.yaml'
|
||||
if os.path.exists(ckpt_config_path):
|
||||
with open(ckpt_config_path) as f:
|
||||
saved_hparams_ = yaml.safe_load(f)
|
||||
if saved_hparams_ is not None:
|
||||
saved_hparams.update(saved_hparams_)
|
||||
hparams_ = {}
|
||||
if args.config != '':
|
||||
hparams_.update(load_config(args.config))
|
||||
if not args.reset:
|
||||
hparams_.update(saved_hparams)
|
||||
hparams_['work_dir'] = args_work_dir
|
||||
|
||||
# Support config overriding in command line. Support list type config overriding.
|
||||
# Examples: --hparams="a=1,b.c=2,d=[1 1 1]"
|
||||
if args.hparams != "":
|
||||
for new_hparam in args.hparams.split(","):
|
||||
k, v = new_hparam.split("=")
|
||||
v = v.strip("\'\" ")
|
||||
config_node = hparams_
|
||||
for k_ in k.split(".")[:-1]:
|
||||
config_node = config_node[k_]
|
||||
k = k.split(".")[-1]
|
||||
if v in ['True', 'False'] or type(config_node[k]) in [bool, list, dict]:
|
||||
if type(config_node[k]) == list:
|
||||
v = v.replace(" ", ",")
|
||||
config_node[k] = eval(v)
|
||||
else:
|
||||
config_node[k] = type(config_node[k])(v)
|
||||
if args_work_dir != '' and args.remove:
|
||||
answer = input("REMOVE old checkpoint? Y/N [Default: N]: ")
|
||||
if answer.lower() == "y":
|
||||
pass
|
||||
if args_work_dir != '' and (not os.path.exists(ckpt_config_path) or args.reset) and not args.infer:
|
||||
os.makedirs(hparams_['work_dir'], exist_ok=True)
|
||||
with open(ckpt_config_path, 'w') as f:
|
||||
yaml.safe_dump(hparams_, f)
|
||||
|
||||
hparams_['infer'] = args.infer
|
||||
hparams_['debug'] = args.debug
|
||||
hparams_['validate'] = args.validate
|
||||
hparams_['exp_name'] = args.exp_name
|
||||
global global_print_hparams
|
||||
if global_hparams:
|
||||
hparams.clear()
|
||||
hparams.update(hparams_)
|
||||
if print_hparams and global_print_hparams and global_hparams:
|
||||
# print('| Hparams chains: ', config_chains)
|
||||
# print('| Hparams: ')
|
||||
# for i, (k, v) in enumerate(sorted(hparams_.items())):
|
||||
# print(f"\033[;33;m{k}\033[0m: {v}, ", end="\n" if i % 5 == 4 else "")
|
||||
# print("")
|
||||
global_print_hparams = False
|
||||
return hparams_
|
||||
@@ -0,0 +1,71 @@
|
||||
import pickle
|
||||
from copy import deepcopy
|
||||
|
||||
import numpy as np
|
||||
|
||||
|
||||
class IndexedDataset:
|
||||
def __init__(self, path, num_cache=1):
|
||||
super().__init__()
|
||||
self.path = path
|
||||
self.data_file = None
|
||||
self.data_offsets = np.load(f"{path}.idx", allow_pickle=True).item()['offsets']
|
||||
self.data_file = open(f"{path}.data", 'rb', buffering=-1)
|
||||
self.cache = []
|
||||
self.num_cache = num_cache
|
||||
|
||||
def check_index(self, i):
|
||||
if i < 0 or i >= len(self.data_offsets) - 1:
|
||||
raise IndexError('index out of range')
|
||||
|
||||
def __del__(self):
|
||||
if self.data_file:
|
||||
self.data_file.close()
|
||||
|
||||
def __getitem__(self, i):
|
||||
self.check_index(i)
|
||||
if self.num_cache > 0:
|
||||
for c in self.cache:
|
||||
if c[0] == i:
|
||||
return c[1]
|
||||
self.data_file.seek(self.data_offsets[i])
|
||||
b = self.data_file.read(self.data_offsets[i + 1] - self.data_offsets[i])
|
||||
item = pickle.loads(b)
|
||||
if self.num_cache > 0:
|
||||
self.cache = [(i, deepcopy(item))] + self.cache[:-1]
|
||||
return item
|
||||
|
||||
def __len__(self):
|
||||
return len(self.data_offsets) - 1
|
||||
|
||||
class IndexedDatasetBuilder:
|
||||
def __init__(self, path):
|
||||
self.path = path
|
||||
self.out_file = open(f"{path}.data", 'wb')
|
||||
self.byte_offsets = [0]
|
||||
|
||||
def add_item(self, item):
|
||||
s = pickle.dumps(item)
|
||||
bytes = self.out_file.write(s)
|
||||
self.byte_offsets.append(self.byte_offsets[-1] + bytes)
|
||||
|
||||
def finalize(self):
|
||||
self.out_file.close()
|
||||
np.save(open(f"{self.path}.idx", 'wb'), {'offsets': self.byte_offsets})
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import random
|
||||
from tqdm import tqdm
|
||||
ds_path = '/tmp/indexed_ds_example'
|
||||
size = 100
|
||||
items = [{"a": np.random.normal(size=[10000, 10]),
|
||||
"b": np.random.normal(size=[10000, 10])} for i in range(size)]
|
||||
builder = IndexedDatasetBuilder(ds_path)
|
||||
for i in tqdm(range(size)):
|
||||
builder.add_item(items[i])
|
||||
builder.finalize()
|
||||
ds = IndexedDataset(ds_path)
|
||||
for i in tqdm(range(10000)):
|
||||
idx = random.randint(0, size - 1)
|
||||
assert (ds[idx]['a'] == items[idx]['a']).all()
|
||||
@@ -0,0 +1,53 @@
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
def sigmoid_focal_loss(
|
||||
inputs: torch.Tensor,
|
||||
targets: torch.Tensor,
|
||||
alpha: float = 0.25,
|
||||
gamma: float = 2,
|
||||
reduction: str = "none",
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Loss used in RetinaNet for dense detection: https://arxiv.org/abs/1708.02002.
|
||||
|
||||
Args:
|
||||
inputs (Tensor): A float tensor of arbitrary shape.
|
||||
The predictions for each example.
|
||||
targets (Tensor): A float tensor with the same shape as inputs. Stores the binary
|
||||
classification label for each element in inputs
|
||||
(0 for the negative class and 1 for the positive class).
|
||||
alpha (float): Weighting factor in range (0,1) to balance
|
||||
positive vs negative examples or -1 for ignore. Default: ``0.25``.
|
||||
gamma (float): Exponent of the modulating factor (1 - p_t) to
|
||||
balance easy vs hard examples. Default: ``2``.
|
||||
reduction (string): ``'none'`` | ``'mean'`` | ``'sum'``
|
||||
``'none'``: No reduction will be applied to the output.
|
||||
``'mean'``: The output will be averaged.
|
||||
``'sum'``: The output will be summed. Default: ``'none'``.
|
||||
Returns:
|
||||
Loss tensor with the reduction option applied.
|
||||
"""
|
||||
# Original implementation from https://github.com/facebookresearch/fvcore/blob/master/fvcore/nn/focal_loss.py
|
||||
p = torch.sigmoid(inputs)
|
||||
ce_loss = F.binary_cross_entropy_with_logits(inputs, targets, reduction="none")
|
||||
p_t = p * targets + (1 - p) * (1 - targets)
|
||||
loss = ce_loss * ((1 - p_t) ** gamma)
|
||||
|
||||
if alpha >= 0: # decrease the importance of negative samples
|
||||
alpha_t = alpha * targets + (1 - alpha) * (1 - targets)
|
||||
loss = alpha_t * loss
|
||||
|
||||
# Check reduction option and return loss accordingly
|
||||
if reduction == "none":
|
||||
pass
|
||||
elif reduction == "mean":
|
||||
loss = loss.mean()
|
||||
elif reduction == "sum":
|
||||
loss = loss.sum()
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Invalid Value for arg 'reduction': '{reduction} \n Supported reduction modes: 'none', 'mean', 'sum'"
|
||||
)
|
||||
return loss
|
||||
|
||||
@@ -0,0 +1,42 @@
|
||||
import time
|
||||
import torch
|
||||
|
||||
|
||||
class AvgrageMeter(object):
|
||||
|
||||
def __init__(self):
|
||||
self.reset()
|
||||
|
||||
def reset(self):
|
||||
self.avg = 0
|
||||
self.sum = 0
|
||||
self.cnt = 0
|
||||
|
||||
def update(self, val, n=1):
|
||||
self.sum += val * n
|
||||
self.cnt += n
|
||||
self.avg = self.sum / self.cnt
|
||||
|
||||
|
||||
class Timer:
|
||||
timer_map = {}
|
||||
|
||||
def __init__(self, name, enable=False):
|
||||
if name not in Timer.timer_map:
|
||||
Timer.timer_map[name] = 0
|
||||
self.name = name
|
||||
self.enable = enable
|
||||
|
||||
def __enter__(self):
|
||||
if self.enable:
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.synchronize()
|
||||
self.t = time.time()
|
||||
|
||||
def __exit__(self, exc_type, exc_val, exc_tb):
|
||||
if self.enable:
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.synchronize()
|
||||
Timer.timer_map[self.name] += time.time() - self.t
|
||||
if self.enable:
|
||||
print(f'[Timer] {self.name}: {Timer.timer_map[self.name]}')
|
||||
@@ -0,0 +1,180 @@
|
||||
import os
|
||||
import traceback
|
||||
from functools import partial
|
||||
from tqdm import tqdm
|
||||
import torch
|
||||
|
||||
def chunked_worker(worker_id, args_queue=None, results_queue=None, init_ctx_func=None):
|
||||
ctx = init_ctx_func(worker_id) if init_ctx_func is not None else None
|
||||
while True:
|
||||
args = args_queue.get()
|
||||
if args == '<KILL>':
|
||||
return
|
||||
job_idx, map_func, arg = args
|
||||
try:
|
||||
map_func_ = partial(map_func, ctx=ctx) if ctx is not None else map_func
|
||||
if isinstance(arg, dict):
|
||||
res = map_func_(**arg)
|
||||
elif isinstance(arg, (list, tuple)):
|
||||
res = map_func_(*arg)
|
||||
else:
|
||||
res = map_func_(arg)
|
||||
results_queue.put((job_idx, res))
|
||||
except:
|
||||
traceback.print_exc()
|
||||
results_queue.put((job_idx, None))
|
||||
|
||||
|
||||
class MultiprocessManager:
|
||||
def __init__(self, num_workers=None, init_ctx_func=None, multithread=False, queue_max=-1):
|
||||
if multithread:
|
||||
from multiprocessing.dummy import Queue, Process
|
||||
else:
|
||||
from multiprocessing import Queue, Process
|
||||
if num_workers is None:
|
||||
num_workers = int(os.getenv('N_PROC', os.cpu_count()))
|
||||
self.num_workers = num_workers
|
||||
self.results_queue = Queue(maxsize=-1)
|
||||
self.jobs_pending = []
|
||||
self.args_queue = Queue(maxsize=queue_max)
|
||||
self.workers = []
|
||||
self.total_jobs = 0
|
||||
self.multithread = multithread
|
||||
for i in range(num_workers):
|
||||
if multithread:
|
||||
p = Process(target=chunked_worker,
|
||||
args=(i, self.args_queue, self.results_queue, init_ctx_func))
|
||||
else:
|
||||
p = Process(target=chunked_worker,
|
||||
args=(i, self.args_queue, self.results_queue, init_ctx_func),
|
||||
daemon=True)
|
||||
self.workers.append(p)
|
||||
p.start()
|
||||
|
||||
def add_job(self, func, args):
|
||||
if not self.args_queue.full():
|
||||
self.args_queue.put((self.total_jobs, func, args))
|
||||
else:
|
||||
self.jobs_pending.append((self.total_jobs, func, args))
|
||||
self.total_jobs += 1
|
||||
|
||||
def get_results(self):
|
||||
self.n_finished = 0
|
||||
while self.n_finished < self.total_jobs:
|
||||
while len(self.jobs_pending) > 0 and not self.args_queue.full():
|
||||
self.args_queue.put(self.jobs_pending[0])
|
||||
self.jobs_pending = self.jobs_pending[1:]
|
||||
job_id, res = self.results_queue.get()
|
||||
yield job_id, res
|
||||
self.n_finished += 1
|
||||
for w in range(self.num_workers):
|
||||
self.args_queue.put("<KILL>")
|
||||
for w in self.workers:
|
||||
w.join()
|
||||
|
||||
def close(self):
|
||||
if not self.multithread:
|
||||
for w in self.workers:
|
||||
w.terminate()
|
||||
|
||||
def __len__(self):
|
||||
return self.total_jobs
|
||||
|
||||
|
||||
def multiprocess_run_tqdm(map_func, args, num_workers=None, ordered=True, init_ctx_func=None,
|
||||
multithread=False, queue_max=-1, desc=None):
|
||||
for i, res in tqdm(
|
||||
multiprocess_run(map_func, args, num_workers, ordered, init_ctx_func, multithread,
|
||||
queue_max=queue_max),
|
||||
total=len(args), desc=desc):
|
||||
yield i, res
|
||||
|
||||
|
||||
def multiprocess_run(map_func, args, num_workers=None, ordered=True, init_ctx_func=None, multithread=False,
|
||||
queue_max=-1):
|
||||
"""
|
||||
Multiprocessing running chunked jobs.
|
||||
|
||||
Examples:
|
||||
>>> for res in tqdm(multiprocess_run(job_func, args):
|
||||
>>> print(res)
|
||||
|
||||
:param map_func:
|
||||
:param args:
|
||||
:param num_workers:
|
||||
:param ordered:
|
||||
:param init_ctx_func:
|
||||
:param q_max_size:
|
||||
:param multithread:
|
||||
:return:
|
||||
"""
|
||||
if num_workers is None:
|
||||
num_workers = int(os.getenv('N_PROC', os.cpu_count()))
|
||||
manager = MultiprocessManager(num_workers, init_ctx_func, multithread, queue_max=queue_max)
|
||||
for arg in args:
|
||||
manager.add_job(map_func, arg)
|
||||
if ordered:
|
||||
n_jobs = len(args)
|
||||
results = ['<WAIT>' for _ in range(n_jobs)]
|
||||
i_now = 0
|
||||
for job_i, res in manager.get_results():
|
||||
results[job_i] = res
|
||||
while i_now < n_jobs and (not isinstance(results[i_now], str) or results[i_now] != '<WAIT>'):
|
||||
yield i_now, results[i_now]
|
||||
results[i_now] = None
|
||||
i_now += 1
|
||||
else:
|
||||
for job_i, res in manager.get_results():
|
||||
yield job_i, res
|
||||
manager.close()
|
||||
|
||||
# #### this is the old version of chunked_multiprocess_run
|
||||
def chunked_worker_old(worker_id, map_func, args, results_queue=None, init_ctx_func=None):
|
||||
ctx = init_ctx_func(worker_id) if init_ctx_func is not None else None
|
||||
for job_idx, arg in args:
|
||||
try:
|
||||
if not isinstance(arg, tuple) and not isinstance(arg, list):
|
||||
arg = [arg]
|
||||
if ctx is not None:
|
||||
res = map_func(*arg, ctx=ctx)
|
||||
else:
|
||||
res = map_func(*arg)
|
||||
results_queue.put((job_idx, res))
|
||||
except:
|
||||
traceback.print_exc()
|
||||
results_queue.put((job_idx, None))
|
||||
|
||||
def chunked_multiprocess_run(
|
||||
map_func, args, num_workers=None, ordered=True,
|
||||
init_ctx_func=None, q_max_size=1000, multithread=False):
|
||||
if multithread:
|
||||
from multiprocessing.dummy import Queue, Process
|
||||
else:
|
||||
from multiprocessing import Queue, Process
|
||||
args = zip(range(len(args)), args)
|
||||
args = list(args)
|
||||
n_jobs = len(args)
|
||||
if num_workers is None:
|
||||
num_workers = int(os.getenv('N_PROC', os.cpu_count()))
|
||||
results_queues = []
|
||||
if ordered:
|
||||
for i in range(num_workers):
|
||||
results_queues.append(Queue(maxsize=q_max_size // num_workers))
|
||||
else:
|
||||
results_queue = Queue(maxsize=q_max_size)
|
||||
for i in range(num_workers):
|
||||
results_queues.append(results_queue)
|
||||
workers = []
|
||||
for i in range(num_workers):
|
||||
args_worker = args[i::num_workers]
|
||||
p = Process(target=chunked_worker_old, args=(
|
||||
i, map_func, args_worker, results_queues[i], init_ctx_func), daemon=True)
|
||||
workers.append(p)
|
||||
p.start()
|
||||
for n_finished in range(n_jobs):
|
||||
results_queue = results_queues[n_finished % num_workers]
|
||||
job_idx, res = results_queue.get()
|
||||
assert job_idx == n_finished or not ordered, (job_idx, n_finished)
|
||||
yield res
|
||||
for w in workers:
|
||||
w.join()
|
||||
@@ -0,0 +1,74 @@
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
def get_filter_2d(kernel, kernel_size, channels, no_grad=True):
|
||||
# Reshape to 2d depthwise convolutional weight
|
||||
kernel = kernel.view(1, 1, kernel_size, kernel_size)
|
||||
kernel = kernel.repeat(channels, 1, 1, 1)
|
||||
|
||||
filter = nn.Conv2d(in_channels=channels, out_channels=channels, kernel_size=kernel_size, groups=channels,
|
||||
bias=False, padding=kernel_size // 2)
|
||||
|
||||
filter.weight.data = kernel
|
||||
if no_grad:
|
||||
filter.weight.requires_grad = False
|
||||
|
||||
return filter
|
||||
|
||||
def get_filter_1d(kernel, kernel_size, channels, no_grad=True):
|
||||
kernel = kernel.view(1, 1, kernel_size)
|
||||
kernel = kernel.repeat(channels, 1, 1)
|
||||
|
||||
filter = nn.Conv1d(in_channels=channels, out_channels=channels, kernel_size=kernel_size, groups=channels,
|
||||
bias=False, padding=kernel_size // 2)
|
||||
|
||||
filter.weight.data = kernel
|
||||
if no_grad:
|
||||
filter.weight.requires_grad = False
|
||||
|
||||
return filter
|
||||
|
||||
def get_gaussian_kernel_2d(kernel_size, sigma):
|
||||
# Create a x, y coordinate grid of shape (kernel_size, kernel_size, 2)
|
||||
x_coord = torch.arange(kernel_size)
|
||||
x_grid = x_coord.repeat(kernel_size).view(kernel_size, kernel_size)
|
||||
y_grid = x_grid.t()
|
||||
xy_grid = torch.stack([x_grid, y_grid], dim=-1).float()
|
||||
|
||||
mean = (kernel_size - 1) / 2.
|
||||
variance = sigma ** 2.
|
||||
|
||||
# Calculate the 2-dimensional gaussian kernel which is
|
||||
# the product of two gaussian distributions for two different
|
||||
# variables (in this case called x and y)
|
||||
gaussian_kernel = (1. / (2. * np.pi * variance)) * torch.exp(
|
||||
-torch.sum((xy_grid - mean) ** 2., dim=-1) / (2 * variance))
|
||||
|
||||
# Make sure sum of values in gaussian kernel equals 1.
|
||||
gaussian_kernel = gaussian_kernel / torch.sum(gaussian_kernel)
|
||||
|
||||
return gaussian_kernel
|
||||
|
||||
def get_gaussian_kernel_1d(kernel_size, sigma):
|
||||
x_grid = torch.arange(kernel_size)
|
||||
mean = (kernel_size - 1) / 2.
|
||||
variance = sigma ** 2.
|
||||
gaussian_kernel = (1. / ((2. * np.pi) ** 0.5 * sigma)) * torch.exp(-(x_grid - mean) ** 2. / (2 * variance))
|
||||
gaussian_kernel = gaussian_kernel / torch.sum(gaussian_kernel)
|
||||
return gaussian_kernel
|
||||
|
||||
def get_hann_kernel_1d(kernel_size, periodic=False):
|
||||
# periodic=False gives symmetric kernel, otherwise equivalent to hann(kernel_size + 1)
|
||||
return torch.hann_window(kernel_size, periodic)
|
||||
|
||||
def get_triangle_kernel_1d(kernel_size):
|
||||
kernel = torch.zeros(kernel_size)
|
||||
for idx in range(kernel_size):
|
||||
kernel[idx] = 1 - abs((idx - (kernel_size - 1) / 2) / ((kernel_size - 1) / 2))
|
||||
return kernel
|
||||
|
||||
def add_gaussian_noise(tensor, mean=0, std=1):
|
||||
noise = torch.randn(tensor.size()) * std + mean
|
||||
noisy_tensor = tensor + noise
|
||||
return noisy_tensor
|
||||
@@ -0,0 +1,5 @@
|
||||
import os
|
||||
|
||||
os.environ["OMP_NUM_THREADS"] = "1"
|
||||
os.environ['TF_NUM_INTEROP_THREADS'] = '1'
|
||||
os.environ['TF_NUM_INTRAOP_THREADS'] = '1'
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user