From 6b9f4d11d2c0be4c13aa7642d102f43af3c464e9 Mon Sep 17 00:00:00 2001 From: ylzz1997 Date: Wed, 30 Aug 2023 02:02:52 +0800 Subject: [PATCH] =?UTF-8?q?debug=20=E5=8F=8D=E9=87=8F=E5=8C=96=20=E6=88=91?= =?UTF-8?q?=E5=8F=91=E7=8E=B0flow=E9=87=8C=E9=9D=A2=E6=9C=89=E5=8F=8D?= =?UTF-8?q?=E9=87=8F=E5=8C=96=E7=9A=84=E4=BB=A3=E7=A0=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- models.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/models.py b/models.py index 37cfe2c..62a2759 100644 --- a/models.py +++ b/models.py @@ -665,12 +665,13 @@ class SynthesizerTrn(nn.Module): attn = monotonic_align.maximum_path(neg_cent, attn_mask.squeeze(1)).unsqueeze(1).detach() w = attn.sum(2) - # 反量化 - w -= torch.rand_like(w) l_length_sdp = self.sdp(x, x_mask, w, g=g) l_length_sdp = l_length_sdp / torch.sum(x_mask) + # 反量化 + w -= torch.rand_like(w) + logw_ = torch.log(w + 1e-6) * x_mask logw = self.dp(x, x_mask, g=g) l_length_dp = torch.sum((logw - logw_) ** 2, [1, 2]) / torch.sum(x_mask) # for averaging