| 159 | |
| 160 | |
| 161 | class ChannelAvgPoolFlat(AvgPool1d): |
| 162 | def forward(self, input): |
| 163 | if len(input.size()) == 4: |
| 164 | n, c, w, h = input.size() |
| 165 | pool = lambda x: F.avg_pool1d( |
| 166 | x, |
| 167 | self.kernel_size, |
| 168 | self.stride, |
| 169 | self.padding, |
| 170 | self.ceil_mode, |
| 171 | ) |
| 172 | out = rearrange( |
| 173 | pool(rearrange(input, "n c w h -> n (w h) c")), |
| 174 | "n (w h) c -> n c w h", |
| 175 | n=n, |
| 176 | w=w, |
| 177 | h=h, |
| 178 | ) |
| 179 | return out.squeeze() |
| 180 | elif len(input.size()) == 3: |
| 181 | n, c, l = input.size() |
| 182 | pool = lambda x: F.avg_pool1d( |
| 183 | x, |
| 184 | self.kernel_size, |
| 185 | self.stride, |
| 186 | self.padding, |
| 187 | self.ceil_mode, |
| 188 | ) |
| 189 | out = rearrange( |
| 190 | pool(rearrange(input, "n c l -> n l c")), |
| 191 | "n l c -> n c l", |
| 192 | n=n, |
| 193 | l=l |
| 194 | ) |
| 195 | return out.squeeze() |
| 196 | else: |
| 197 | raise NotImplementedError |
| 198 | |
| 199 | |
| 200 | class AttributeTransformer2(nn.Module): |