import matplotlib.pyplot as plt
import numpy as np
import tensorflow as tf


CLASS_NAMES = [
    "T-shirt/top",
    "Trouser",
    "Pullover",
    "Dress",
    "Coat",
    "Sandal",
    "Shirt",
    "Sneaker",
    "Bag",
    "Ankle boot",
]


def main() -> None:
    tf.keras.utils.set_random_seed(42)

    # 初回実行時にFashion-MNISTがインターネットからダウンロードされる。
    (train_images, train_labels), (test_images, test_labels) = (
        tf.keras.datasets.fashion_mnist.load_data()
    )
    train_images = train_images.astype("float32") / 255.0
    test_images = test_images.astype("float32") / 255.0
    train_images = train_images[..., np.newaxis]
    test_images = test_images[..., np.newaxis]

    model = tf.keras.Sequential(
        [
            tf.keras.layers.Input(shape=(28, 28, 1)),
            tf.keras.layers.Conv2D(32, 3, activation="relu"),
            tf.keras.layers.MaxPooling2D(),
            tf.keras.layers.Conv2D(64, 3, activation="relu"),
            tf.keras.layers.MaxPooling2D(),
            tf.keras.layers.Flatten(),
            tf.keras.layers.Dense(64, activation="relu"),
            tf.keras.layers.Dropout(0.3),
            tf.keras.layers.Dense(10, activation="softmax"),
        ]
    )
    model.compile(
        optimizer="adam",
        loss="sparse_categorical_crossentropy",
        metrics=["accuracy"],
    )
    model.summary()
    model.fit(
        train_images,
        train_labels,
        epochs=5,
        batch_size=64,
        validation_split=0.1,
    )

    test_loss, test_accuracy = model.evaluate(test_images, test_labels, verbose=0)
    print(f"テストデータの損失: {test_loss:.4f}")
    print(f"テストデータの正解率: {test_accuracy:.2%}")

    probabilities = model.predict(test_images, verbose=0)
    predicted_labels = probabilities.argmax(axis=1)
    wrong_indices = np.flatnonzero(predicted_labels != test_labels)[:12]

    fig, axes = plt.subplots(3, 4, figsize=(9, 7))
    for axis, index in zip(axes.ravel(), wrong_indices):
        axis.imshow(test_images[index].squeeze(), cmap="gray")
        actual = CLASS_NAMES[test_labels[index]]
        predicted = CLASS_NAMES[predicted_labels[index]]
        confidence = probabilities[index, predicted_labels[index]]
        axis.set_title(
            f"actual: {actual}\npred: {predicted} ({confidence:.0%})",
            fontsize=9,
        )
        axis.axis("off")
    fig.suptitle("Misclassified test images")
    fig.tight_layout()
    fig.savefig("fashion_mnist_errors.png", dpi=150)
    model.save("fashion_mnist_model.keras")
    print("誤分類画像と学習済みモデルを保存しました")


if __name__ == "__main__":
    main()
