Files
yqzhishen 9f86135c3c Basic training framework and acoustic model training (#254)
* 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
2025-05-02 16:16:19 +08:00

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