MCPcopy Create free account
hub / github.com/FireRedTeam/FireRedTTS2 / UpConv

Class UpConv

fireredtts2/codec/model.py:123–148  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

121
122
123class UpConv(nn.Module):
124 def __init__(
125 self,
126 embed_dim: int = 768,
127 stride: int = 4,
128 ):
129 super().__init__()
130 self.embed_dim = embed_dim
131 self.stride = stride
132 self.in_proj = nn.Linear(embed_dim, self.stride * embed_dim)
133 # Simple transpose convolution layer to keep channel number consistent
134 self.up_conv = nn.ConvTranspose1d(
135 self.stride * embed_dim,
136 embed_dim,
137 kernel_size=stride,
138 stride=stride,
139 bias=False,
140 )
141
142 def forward(self, x: torch.Tensor, input_length: torch.Tensor):
143 x = self.in_proj(x)
144 x = x.transpose(1, 2)
145 res = self.up_conv(x)
146 res = res.transpose(1, 2)
147 output_length = input_length * self.stride
148 return res, output_length
149
150
151class RedCodec(nn.Module):

Callers 1

from_configMethod · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected