Receives x input of dim [N,C], and repeats the vector to create tensor of shape [N, C, K] : repeats: int, the number of repetitions for the vector.
| 101 | return self.network(x) |
| 102 | |
| 103 | class RepeatVector(nn.Module): |
| 104 | """ |
| 105 | Receives x input of dim [N,C], and repeats the vector |
| 106 | to create tensor of shape [N, C, K] |
| 107 | : repeats: int, the number of repetitions for the vector. |
| 108 | """ |
| 109 | def __init__(self, repeats): |
| 110 | super(RepeatVector, self).__init__() |
| 111 | self.repeats = repeats |
| 112 | |
| 113 | def forward(self, x): |
| 114 | x = x.unsqueeze(-1).repeat(1, 1, self.repeats) # <------------ Mejorar? |
| 115 | return x |
| 116 | |
| 117 | class Chomp1d(nn.Module): |
| 118 | """ |