import os os.environ["KERAS_BACKEND"] = "jax" import keras import numpy as np from keras import layers keras.utils.set_random_seed(23) image_shape = (96, 96, 3) num_classes = 2 rng = np.random.default_rng(23) x_train = rng.uniform(0, 255, size=(8, *image_shape)).astype("float32") y_train = np.array([0, 1, 0, 1, 0, 1, 0, 1], dtype="int32") x_val = rng.uniform(0, 255, size=(4, *image_shape)).astype("float32") y_val = np.array([0, 1, 0, 1], dtype="int32")