| 284 | |
| 285 | |
| 286 | class Downsample(nn.Module): |
| 287 | with_conv: bool |
| 288 | |
| 289 | @nn.compact |
| 290 | def __call__(self, hidden_states): |
| 291 | if self.with_conv: |
| 292 | hidden_states = jnp.pad( |
| 293 | hidden_states, |
| 294 | [(0, 0), (0, 1), (0, 1), (0, 0)] |
| 295 | ) |
| 296 | hidden_states = nn.Conv( |
| 297 | hidden_states.shape[-1], [3, 3], |
| 298 | strides=[2, 2], |
| 299 | padding="VALID" |
| 300 | )(hidden_states) |
| 301 | else: |
| 302 | hidden_states = nn.avg_pool(hidden_states, [2, 2], [2, 2]) |
| 303 | return hidden_states |
| 304 | |
| 305 | |
| 306 | class Upsample(nn.Module): |