# Access a specific layer's weights print(model.linear1.kernel.value.shape) # (3136, 256) print(model.linear1.bias.value.shape) # (256,) # Parameters are nnx.Param objects; .value gets the JAX array __ __