Skip to content

Commit b000a33

Browse files
[TTS][EasyMagpietts] Added CFG distillation
1 parent a95ea79 commit b000a33

6 files changed

Lines changed: 4392 additions & 12 deletions

File tree

.github/workflows/cicd-main-speech.yml

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -402,6 +402,9 @@ jobs:
402402
- runner: ${{ inputs.runner }}
403403
script: L2_TTS_Fast_dev_runs_EasyMagpietts_OnlinePO
404404
timeout: 20
405+
- runner: ${{ inputs.runner }}
406+
script: L2_TTS_Fast_dev_runs_EasyMagpietts_OnlineCFGDistillation
407+
timeout: 20
405408
- runner: ${{ inputs.runner }}
406409
script: L2_TTS_InferEvaluate_EasyMagpietts
407410
timeout: 20

examples/tts/easy_magpietts.py

100644100755
Lines changed: 13 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -16,11 +16,17 @@
1616
import torch.multiprocessing as mp
1717
from omegaconf import OmegaConf, open_dict
1818

19-
from nemo.collections.tts.models import EasyMagpieTTSModel, EasyMagpieTTSModelOnlinePO
19+
from nemo.collections.tts.models import EasyMagpieCFGDistillation, EasyMagpieTTSModel, EasyMagpieTTSModelOnlinePO
2020
from nemo.core.config import hydra_runner
2121
from nemo.utils import logging
2222
from nemo.utils.exp_manager import exp_manager
2323

24+
_TRAIN_MODES: list[str] = [
25+
"train",
26+
"online_cfg_distillation_train",
27+
"onlinepo_train",
28+
]
29+
2430

2531
@hydra_runner(config_path="conf/magpietts", config_name="easy_magpietts")
2632
def main(cfg):
@@ -43,8 +49,12 @@ def main(cfg):
4349
exp_manager(trainer, cfg.get("exp_manager", None))
4450

4551
mode = cfg.get('mode', 'train')
52+
train_modes_msg = ", ".join(_TRAIN_MODES)
53+
4654
if mode == 'train':
4755
model = EasyMagpieTTSModel(cfg=cfg.model, trainer=trainer)
56+
elif mode == "online_cfg_distillation_train":
57+
model = EasyMagpieCFGDistillation(cfg=cfg.model, trainer=trainer)
4858
elif mode == 'onlinepo_train':
4959
model_cfg = cfg.model
5060
with open_dict(model_cfg):
@@ -53,14 +63,14 @@ def main(cfg):
5363
elif mode == 'test':
5464
model = EasyMagpieTTSModel(cfg=cfg.model, trainer=trainer)
5565
else:
56-
raise NotImplementedError(f"Only train, onlinepo_train and test modes are supported. Got {mode}")
66+
raise NotImplementedError(f"Only {train_modes_msg} and test modes are supported. Got {mode}.")
5767

5868
if cfg.get("pretrained_model", None):
5969
model.restore_from_pretrained_checkpoint(cfg.pretrained_model)
6070

6171
model.maybe_init_from_pretrained_checkpoint(cfg=cfg)
6272

63-
if mode in ['train', 'onlinepo_train']:
73+
if mode in _TRAIN_MODES:
6474
trainer.fit(model)
6575
elif mode == 'test':
6676
trainer.test(model)

nemo/collections/tts/losses/magpietts_cfg_distillation.py

Lines changed: 49 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -15,7 +15,7 @@
1515
Losses used in CFG distillation of the MagpieTTS model.
1616
"""
1717

18-
from typing import Generator, Optional
18+
from typing import Callable, Generator, Optional
1919

2020
import torch
2121
from torch import Tensor, nn
@@ -29,8 +29,13 @@
2929
"NRMSELogitsLoss",
3030
]
3131

32+
_CODEBOOK_ORDERING_MODES: set[str] = {
33+
"frame-major",
34+
"codebook-major",
35+
}
3236

33-
def _iter_slices(
37+
38+
def _iter_slices_frame_major(
3439
num_codebooks: int,
3540
num_tokens_per_codebook: int,
3641
frame_stacking_factor: int,
@@ -48,6 +53,33 @@ def _iter_slices(
4853
yield fs_index, codebook, start, end, slice_mask, slice_len
4954

5055

56+
def _iter_slices_codebook_major(
57+
num_codebooks: int,
58+
num_tokens_per_codebook: int,
59+
frame_stacking_factor: int,
60+
mask: Tensor,
61+
) -> Generator[tuple[int, int, int, int, Tensor, Tensor], None, None]:
62+
for codebook in range(num_codebooks):
63+
for fs_index in range(frame_stacking_factor):
64+
slice_mask = mask[:, fs_index::frame_stacking_factor].float()
65+
slice_len = slice_mask.sum(dim=-1).clamp_min(1)
66+
67+
channel = codebook * frame_stacking_factor + fs_index
68+
start = channel * num_tokens_per_codebook
69+
end = start + num_tokens_per_codebook
70+
71+
yield fs_index, codebook, start, end, slice_mask, slice_len
72+
73+
74+
def _get_slice_iterator(mode: str) -> Callable:
75+
if mode not in _CODEBOOK_ORDERING_MODES:
76+
raise ValueError(
77+
f"Unsupported codebook ordering {mode!r}; expected one of {sorted(_CODEBOOK_ORDERING_MODES)}."
78+
)
79+
80+
return _iter_slices_frame_major if mode == "frame-major" else _iter_slices_codebook_major
81+
82+
5183
class KLDivergenceLoss(Loss):
5284
"""The Kullback-Leibler divergence loss."""
5385

@@ -79,11 +111,13 @@ def __init__(
79111
num_codebooks: int,
80112
num_tokens_per_codebook: int,
81113
frame_stacking_factor: int,
114+
codebook_ordering: str = "frame-major",
82115
) -> None:
83116
super().__init__()
84117
self.num_codebooks = num_codebooks
85118
self.num_tokens_per_codebook = num_tokens_per_codebook
86119
self.frame_stacking_factor = frame_stacking_factor
120+
self.iter_slices: Callable = _get_slice_iterator(mode=codebook_ordering)
87121
self.criterion = nn.KLDivLoss(reduction="none", log_target=False)
88122

89123
@typecheck()
@@ -111,17 +145,19 @@ def forward(
111145
Tensor: Scalar tensor representing the averaged masked KL divergence loss.
112146
"""
113147
loss = 0.0
114-
student_log_probs = student_logits.log_softmax(dim=-1)
115-
teacher_probs = teacher_logits.softmax(dim=-1)
116148

117-
for _, _, start, end, slice_mask, slice_len in _iter_slices(
149+
for _, _, start, end, slice_mask, slice_len in self.iter_slices(
118150
self.num_codebooks,
119151
self.num_tokens_per_codebook,
120152
self.frame_stacking_factor,
121153
mask,
122154
):
123-
teacher_probs_slice = teacher_probs[:, :, start:end]
124-
student_log_probs_slice = student_log_probs[:, :, start:end]
155+
# Normalize within this head only. Normalizing over the full
156+
# concatenated dimension would incorrectly make
157+
# independent codebook heads compete for probability mass.
158+
student_log_probs_slice = student_logits[:, :, start:end].log_softmax(dim=-1)
159+
teacher_probs_slice = teacher_logits[:, :, start:end].softmax(dim=-1)
160+
125161
slice_loss = self.criterion(input=student_log_probs_slice, target=teacher_probs_slice)
126162
slice_loss = slice_loss.sum(dim=-1)
127163
slice_loss = (slice_loss * slice_mask).sum(dim=-1) / slice_len
@@ -166,11 +202,13 @@ def __init__(
166202
num_codebooks: int,
167203
num_tokens_per_codebook: int,
168204
frame_stacking_factor: int,
205+
codebook_ordering: str = "frame-major",
169206
) -> None:
170207
super().__init__()
171208
self.num_codebooks = num_codebooks
172209
self.num_tokens_per_codebook = num_tokens_per_codebook
173210
self.frame_stacking_factor = frame_stacking_factor
211+
self.iter_slices: Callable = _get_slice_iterator(mode=codebook_ordering)
174212
self.criterion = nn.CrossEntropyLoss(reduction="none")
175213

176214
@typecheck()
@@ -199,7 +237,7 @@ def forward(
199237
"""
200238
loss = 0.0
201239

202-
for fs_index, codebook, start, end, slice_mask, slice_len in _iter_slices(
240+
for fs_index, codebook, start, end, slice_mask, slice_len in self.iter_slices(
203241
self.num_codebooks,
204242
self.num_tokens_per_codebook,
205243
self.frame_stacking_factor,
@@ -250,11 +288,13 @@ def __init__(
250288
num_codebooks: int,
251289
num_tokens_per_codebook: int,
252290
frame_stacking_factor: int,
291+
codebook_ordering: str = "frame-major",
253292
) -> None:
254293
super().__init__()
255294
self.num_codebooks = num_codebooks
256295
self.num_tokens_per_codebook = num_tokens_per_codebook
257296
self.frame_stacking_factor = frame_stacking_factor
297+
self.iter_slices: Callable = _get_slice_iterator(mode=codebook_ordering)
258298
self.eps = 1e-8
259299
self.criterion = nn.MSELoss(reduction="none")
260300

@@ -286,7 +326,7 @@ def forward(
286326
student_logits = student_logits.masked_fill(inf_mask, 0.0)
287327
loss = 0.0
288328

289-
for _, _, start, end, slice_mask, slice_len in _iter_slices(
329+
for _, _, start, end, slice_mask, slice_len in self.iter_slices(
290330
self.num_codebooks,
291331
self.num_tokens_per_codebook,
292332
self.frame_stacking_factor,

nemo/collections/tts/models/__init__.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -15,6 +15,7 @@
1515
from nemo.collections.tts.models.aligner import AlignerModel
1616
from nemo.collections.tts.models.audio_codec import AudioCodecModel
1717
from nemo.collections.tts.models.easy_magpietts import EasyMagpieTTSModel
18+
from nemo.collections.tts.models.easy_magpietts_cfg_distillation import EasyMagpieCFGDistillation
1819
from nemo.collections.tts.models.easy_magpietts_inference import EasyMagpieTTSInferenceModel
1920
from nemo.collections.tts.models.easy_magpietts_preference_optimization import EasyMagpieTTSModelOnlinePO
2021
from nemo.collections.tts.models.fastpitch import FastPitchModel
@@ -45,4 +46,5 @@
4546
"MagpieTTSModelOfflinePODataGen",
4647
"MagpieTTSModelOfflinePO",
4748
"MagpieTTSModelOnlinePO",
49+
"EasyMagpieCFGDistillation",
4850
]

0 commit comments

Comments
 (0)