File size: 665 Bytes
07d181a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
import jax
import jax.numpy as jnp
from flax import linen as nn

Array = jax.Array


class Cifar10CNN(nn.Module):
    num_classes: int = 10
    channels: tuple[int, ...] = (64, 64, 128, 128, 256, 256)
    groups: int = 8

    @nn.compact
    def __call__(self, x: Array) -> Array:
        for i, c in enumerate(self.channels):
            x = nn.Conv(c, kernel_size=(3, 3), padding="SAME")(x)
            x = nn.GroupNorm(num_groups=min(self.groups, c))(x)
            x = nn.relu(x)
            if i % 2 == 1:
                x = nn.max_pool(x, window_shape=(2, 2), strides=(2, 2))
        x = jnp.mean(x, axis=(1, 2))
        return nn.Dense(self.num_classes)(x)