import os import numpy as np import tensorflow as tf import tensorflow_datasets as tfds import matplotlib.pyplot as plt import jax import jax.numpy as jnp import optax from tqdm.auto import tqdm from flax import linen as nn from flax.training import train_state import dm_pix as pix # Image processing in JAX __ __