9f86135c3c
* Refactor some binarizer modules * Training framework * Acoustic model training * Restore config values * Disable `shuffle_batches` * Support saving weights only * Fix metric candidates * Support console color and wrap empty objects * Support specifying log dir * `dask.compute` everything at once to improve perf * Support nested schedulers * Support auto resuming from latest checkpoint * Remove some default values * Check for weights_only * Support more flexible optimizer settings * Fix stuck in DDP (probably) * Use default monitor candidates * Support `ReduceLROnPlateau` scheduler * Require sync with validation for ReduceLROnPlateau * Print metrics * Add suggestion to change coverage check option * Support fine-tuning and parameter freezing * Move registries * Edit message * Fix high memory usage and epoch syncing * Add `build_xxx_dataset` methods * Simplify augmentation index * Fix typo: compact -> compat * Add rank in file pattern to avoid conflict * `torch.load` with weights_only=True * Change to rank_zero_info * sync_dist=True * Fix `rank_zero_only.rank` needs to be set before use * Try to optimize message * Try to optimize message * Try to optimize message * Fix config check failure * Fix overlapping points on TensorBoard when accumulate_grad_batches > 1 * Support EMA (experimental) * Rename file * Rename embedding * Add comments and type hints * Clean up and re-organize code * Rename file * Rename `used` to `enabled` * Remove redundant wrapper method * Fix voicing extraction * Add augmentation flag and check * Update tensorboard logging
89 lines
3.2 KiB
Python
89 lines
3.2 KiB
Python
import torch
|
|
from torch import nn
|
|
|
|
from lib.config.schema import LinguisticEncoderConfig, MelodyEncoderConfig
|
|
from lib.reflection import filter_kwargs_by_class
|
|
from .commons.common_layers import (
|
|
NormalInitEmbedding as Embedding,
|
|
XavierUniformInitLinear as Linear,
|
|
)
|
|
from .commons.tts_modules import FastSpeech2Encoder
|
|
|
|
__all__ = [
|
|
"LinguisticEncoder",
|
|
"MelodyEncoder",
|
|
]
|
|
|
|
ENCODERS = {
|
|
"fs2": FastSpeech2Encoder,
|
|
}
|
|
|
|
|
|
class LinguisticEncoder(nn.Module):
|
|
def __init__(self, config: LinguisticEncoderConfig):
|
|
super().__init__()
|
|
self.token_embedding = Embedding(config.vocab_size, config.hidden_size, padding_idx=0)
|
|
self.use_lang_id = config.use_lang_id
|
|
if self.use_lang_id:
|
|
self.language_embedding = Embedding(config.num_lang + 1, config.hidden_size, padding_idx=0)
|
|
self.duration_embedding = Linear(1, config.hidden_size)
|
|
self.encoder = (cls := ENCODERS[config.arch])(
|
|
hidden_size=config.hidden_size, **filter_kwargs_by_class(cls, config.kwargs)
|
|
)
|
|
|
|
def forward(self, tokens, durations, languages=None):
|
|
txt_embed = self.token_embedding(tokens)
|
|
dur_embed = self.duration_embedding(durations[:, :, None].float())
|
|
if self.use_lang_id:
|
|
lang_embed = self.language_embedding(languages)
|
|
extra_embed = dur_embed + lang_embed
|
|
else:
|
|
extra_embed = dur_embed
|
|
encoder_out = self.encoder(txt_embed, extra_embed, tokens == 0)
|
|
|
|
return encoder_out
|
|
|
|
|
|
class MelodyEncoder(nn.Module):
|
|
def __init__(self, config: MelodyEncoderConfig):
|
|
super().__init__()
|
|
# MIDI inputs
|
|
hidden_size = config.hidden_size
|
|
self.midi_embedding = Linear(1, hidden_size)
|
|
self.duration_embedding = Linear(1, hidden_size)
|
|
|
|
# ornament inputs
|
|
self.use_glide_embed = config.use_glide_id
|
|
self.glide_embed_scale = config.glide_embed_scale
|
|
if self.use_glide_embed:
|
|
# 0: none, 1: up, 2: down
|
|
self.glide_embedding = Embedding(config.num_glide + 1, hidden_size, padding_idx=0)
|
|
|
|
self.encoder = ENCODERS[config.arch](
|
|
hidden_size=config.hidden_size, **config.kwargs
|
|
)
|
|
self.out_proj = Linear(hidden_size, config.out_size)
|
|
|
|
def forward(self, note_midi, note_rest, note_dur, glide=None):
|
|
"""
|
|
:param note_midi: float32 [B, T_n], -1: padding
|
|
:param note_rest: bool [B, T_n]
|
|
:param note_dur: int64 [B, T_n]
|
|
:param glide: int64 [B, T_n]
|
|
:return: [B, T_n, H]
|
|
"""
|
|
midi_embed = self.midi_embedding(note_midi[:, :, None]) * ~note_rest[:, :, None]
|
|
dur_embed = self.duration_embedding(note_dur.float()[:, :, None])
|
|
ornament_embed = 0
|
|
ornament_embeds = []
|
|
if self.use_glide_embed:
|
|
ornament_embeds.append(self.glide_embedding(glide) * self.glide_embed_scale)
|
|
if ornament_embeds:
|
|
ornament_embed = torch.stack(ornament_embeds, dim=-1).sum(dim=-1)
|
|
encoder_out = self.encoder(
|
|
midi_embed, dur_embed + ornament_embed,
|
|
padding_mask=note_rest
|
|
)
|
|
encoder_out = self.out_proj(encoder_out)
|
|
return encoder_out
|