From 323eb0298d5c935c95e0db1fcf4358d440aff5bb Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Stardust=C2=B7=E5=87=8F?= <2225664821@qq.com> Date: Sun, 20 Aug 2023 15:00:29 +0800 Subject: [PATCH] Update attentions.py --- attentions.py | 17 ++++++++++++++++- 1 file changed, 16 insertions(+), 1 deletion(-) diff --git a/attentions.py b/attentions.py index d752f24..634897f 100644 --- a/attentions.py +++ b/attentions.py @@ -7,7 +7,22 @@ from torch.nn import functional as F import commons import modules -from modules import LayerNorm + +class LayerNorm(nn.Module): + def __init__(self, channels, eps=1e-5): + super().__init__() + self.channels = channels + self.eps = eps + + self.gamma = nn.Parameter(torch.ones(channels)) + self.beta = nn.Parameter(torch.zeros(channels)) + + def forward(self, x): + x = x.transpose(1, -1) + x = F.layer_norm(x, (self.channels,), self.gamma, self.beta, self.eps) + return x.transpose(1, -1) + + @torch.jit.script def fused_add_tanh_sigmoid_multiply(input_a, input_b, n_channels):