diff --git a/modules.py b/modules.py index cc0aecf..92e0f32 100644 --- a/modules.py +++ b/modules.py @@ -370,7 +370,22 @@ class ConvFlow(nn.Module): unnormalized_widths = h[..., :self.num_bins] / math.sqrt(self.filter_channels) unnormalized_heights = h[..., self.num_bins:2*self.num_bins] / math.sqrt(self.filter_channels) unnormalized_derivatives = h[..., 2 * self.num_bins:] - + + x1, logabsdet = piecewise_rational_quadratic_transform(x1, + unnormalized_widths, + unnormalized_heights, + unnormalized_derivatives, + inverse=reverse, + tails='linear', + tail_bound=self.tail_bound + ) + + x = torch.cat([x0, x1], 1) * x_mask + logdet = torch.sum(logabsdet * x_mask, [1,2]) + if not reverse: + return x, logdet + else: + return x class TransformerCouplingLayer(nn.Module): def __init__(self, channels,