1515Losses used in CFG distillation of the MagpieTTS model.
1616"""
1717
18- from typing import Generator , Optional
18+ from typing import Callable , Generator , Optional
1919
2020import torch
2121from torch import Tensor , nn
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+
5183class 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 ,
0 commit comments