| 289 | |
| 290 | |
| 291 | class Downsample1d(nn.Module): |
| 292 | def __init__(self, kernel="linear", pad_mode="reflect"): |
| 293 | super().__init__() |
| 294 | self.pad_mode = pad_mode |
| 295 | kernel_1d = torch.tensor(_kernels[kernel]) |
| 296 | self.pad = kernel_1d.shape[0] // 2 - 1 |
| 297 | self.register_buffer("kernel", kernel_1d) |
| 298 | |
| 299 | def forward(self, hidden_states): |
| 300 | hidden_states = F.pad(hidden_states, (self.pad,) * 2, self.pad_mode) |
| 301 | weight = hidden_states.new_zeros([hidden_states.shape[1], hidden_states.shape[1], self.kernel.shape[0]]) |
| 302 | indices = torch.arange(hidden_states.shape[1], device=hidden_states.device) |
| 303 | kernel = self.kernel.to(weight)[None, :].expand(hidden_states.shape[1], -1) |
| 304 | weight[indices, indices] = kernel |
| 305 | return F.conv1d(hidden_states, weight, stride=2) |
| 306 | |
| 307 | |
| 308 | class Upsample1d(nn.Module): |