MCPcopy Create free account
hub / github.com/Standard-Intelligence/hertz-dev / CausalConvTranspose1d

Class CausalConvTranspose1d

tokenizer.py:133–169  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

131
132
133class CausalConvTranspose1d(NonCausalConvTranspose1d):
134 def __init__(
135 self,
136 in_channels,
137 out_channels,
138 kernel_size,
139 stride,
140 bias=True,
141 pad_buffer=None,
142 ):
143 super(CausalConvTranspose1d, self).__init__(
144 in_channels=in_channels,
145 out_channels=out_channels,
146 kernel_size=kernel_size,
147 stride=stride,
148 padding=0,
149 output_padding=0,
150 bias=bias,
151 )
152 self.stride = stride
153 self.pad_length = (math.ceil(kernel_size/stride) - 1)
154 if pad_buffer is None:
155 pad_buffer = T.zeros(1, in_channels, self.pad_length)
156 self.register_buffer("pad_buffer", pad_buffer)
157
158 def forward(self, x):
159 pad = nn.ReplicationPad1d((self.pad_length, 0))
160 x = pad(x)
161 return self.deconv(x)[:, :, self.stride : -self.stride]
162
163 def inference(self, x):
164 x = T.cat((self.pad_buffer, x), -1)
165 self.pad_buffer = x[:, :, -self.pad_length:]
166 return self.deconv(x)[:, :, self.stride : -self.stride]
167
168 def reset_buffer(self):
169 self.pad_buffer.zero_()
170
171
172class NonCausalResUnit(nn.Module):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected