import os os.environ["TF_CPP_MIN_LOG_LEVEL"] = "2" import tensorflow as tf tf.get_logger().setLevel("ERROR") tf.keras.utils.set_random_seed(7) features = tf.random.stateless_normal((512, 4), seed=(7, 11)) scores = 1.5 * features[:, 0] - features[:, 1] + 0.5 * features[:, 2] labels = tf.cast(scores > 0, tf.float32)[:, None] train_dataset = ( tf.data.Dataset.from_tensor_slices((features[:448], labels[:448])) .shuffle(448, seed=7, reshuffle_each_iteration=True) .batch(32) ) validation_dataset = tf.data.Dataset.from_tensor_slices( (features[448:], labels[448:]) ).batch(32) model = tf.keras.Sequential( [ tf.keras.layers.Input(shape=(4,)), tf.keras.layers.Dense( 8, activation="relu", kernel_regularizer=tf.keras.regularizers.L2(1e-4), ), tf.keras.layers.Dense(1, activation="sigmoid"), ] ) loss_function = tf.keras.losses.BinaryCrossentropy() optimizer = tf.keras.optimizers.Adam(learning_rate=0.02) train_loss = tf.keras.metrics.Mean(name="train_loss") train_accuracy = tf.keras.metrics.BinaryAccuracy(name="train_accuracy") validation_loss = tf.keras.metrics.Mean(name="validation_loss") validation_accuracy = tf.keras.metrics.BinaryAccuracy(name="validation_accuracy") @tf.function def train_step(batch_features, batch_labels): with tf.GradientTape() as tape: predictions = model(batch_features, training=True) loss = loss_function(batch_labels, predictions) loss += tf.add_n(model.losses) gradients = tape.gradient(loss, model.trainable_variables) optimizer.apply_gradients(zip(gradients, model.trainable_variables)) train_loss.update_state(loss) train_accuracy.update_state(batch_labels, predictions) @tf.function def validation_step(batch_features, batch_labels): predictions = model(batch_features, training=False) loss = loss_function(batch_labels, predictions) loss += tf.add_n(model.losses) validation_loss.update_state(loss) validation_accuracy.update_state(batch_labels, predictions) initial_validation_loss = None for epoch in range(1, 6): train_loss.reset_state() train_accuracy.reset_state() validation_loss.reset_state() validation_accuracy.reset_state() for batch_features, batch_labels in train_dataset: train_step(batch_features, batch_labels) for batch_features, batch_labels in validation_dataset: validation_step(batch_features, batch_labels) if initial_validation_loss is None: initial_validation_loss = tf.identity(validation_loss.result()) print( f"epoch={epoch} " f"train_loss={train_loss.result():.4f} " f"train_accuracy={train_accuracy.result():.4f} " f"validation_loss={validation_loss.result():.4f} " f"validation_accuracy={validation_accuracy.result():.4f}" ) tf.debugging.assert_less(validation_loss.result(), initial_validation_loss) tf.debugging.assert_greater(validation_accuracy.result(), 0.90) print( "custom_loop_check=passed " f"optimizer_steps={int(optimizer.iterations)} " f"validation_accuracy={validation_accuracy.result():.4f}" )