from flax import nnx class InspectableCNN(nnx.Module): def __init__(self, *, rngs: nnx.Rngs): self.conv1 = nnx.Conv(1, 32, kernel_size=(3, 3), rngs=rngs) self.linear = nnx.Linear(21632, 10, rngs=rngs) def __call__(self, x): x = self.conv1(x) x = nnx.relu(x) # Stash this activation for later inspection self.sow(nnx.Intermediate, 'conv1_activation', x) x = x.reshape(x.shape[0], -1) return self.linear(x) model = InspectableCNN(rngs=nnx.Rngs(0)) dummy_input = jnp.ones((4, 28, 28, 1)) output = model(dummy_input) # Access what got "sown" print(model.conv1_activation.value.shape) # (4, 26, 26, 32) __ __