import numpy as np
import tensorflow as tf
from sklearn.datasets import make_moons
from sklearn.linear_model import LogisticRegression
from sklearn.metrics import balanced_accuracy_score, f1_score
from sklearn.model_selection import train_test_split
from sklearn.pipeline import make_pipeline
from sklearn.preprocessing import StandardScaler

tf.keras.utils.set_random_seed(7)
X, y = make_moons(n_samples=1000, noise=0.25, random_state=7)
X = X.astype("float32")
y = y.astype("float32")
X_dev, X_test, y_dev, y_test = train_test_split(
    X, y, test_size=0.2, random_state=42, stratify=y
)
X_train, X_val, y_train, y_val = train_test_split(
    X_dev, y_dev, test_size=0.25, random_state=42, stratify=y_dev
)
y_train = y_train.reshape(-1, 1)
y_val = y_val.reshape(-1, 1)

normalize = tf.keras.layers.Normalization(axis=-1)
normalize.adapt(X_train)
model = tf.keras.Sequential([
    tf.keras.Input(shape=(2,)),
    normalize,
    tf.keras.layers.Dense(16, activation="relu"),
    tf.keras.layers.Dense(16, activation="relu"),
    tf.keras.layers.Dense(1, activation="sigmoid"),
])
model.compile(
    optimizer=tf.keras.optimizers.Adam(learning_rate=0.001),
    loss=tf.keras.losses.BinaryCrossentropy(from_logits=False),
    metrics=[tf.keras.metrics.BinaryAccuracy(name="accuracy")],
)
assert sum(np.prod(weight.shape) for weight in model.trainable_weights) == 337
stop = tf.keras.callbacks.EarlyStopping(
    monitor="val_loss", patience=10, restore_best_weights=True
)
history = model.fit(
    X_train, y_train, validation_data=(X_val, y_val),
    epochs=100, batch_size=32, callbacks=[stop], verbose=0
)
print("Train/validation/test rows:", len(X_train), len(X_val), len(X_test))
print("Epochs run:", len(history.history["loss"]))
print("Best validation epoch:", int(np.argmin(history.history["val_loss"])) + 1)

# A fixed reference, trained on the same training rows.
baseline = make_pipeline(StandardScaler(), LogisticRegression(max_iter=1000))
baseline.fit(X_train, y_train.ravel())
probability = model(X_test, training=False).numpy().ravel()
prediction = (probability >= 0.5).astype(int)
assert probability.shape == y_test.shape
assert np.isfinite(probability).all()
assert ((probability >= 0) & (probability <= 1)).all()
print("Logistic test balanced accuracy:",
      balanced_accuracy_score(y_test, baseline.predict(X_test)))
print("Neural test balanced accuracy:", balanced_accuracy_score(y_test, prediction))
print("Neural test F1:", f1_score(y_test, prediction))

model.save("ml-week-4.keras")
restored = tf.keras.models.load_model("ml-week-4.keras")
reloaded_probability = restored(X_test, training=False).numpy().ravel()
np.testing.assert_allclose(probability, reloaded_probability, atol=1e-6, rtol=1e-6)
print("Saved and reloaded predictions match.")
