#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
nn_activation_animation.py

教育用:
  y = sin(pi x) を 1隠れ層ニューラルネットワークで近似し、
  学習の進行をアニメーション表示する。

活性化関数:
    --activation relu
    --activation tanh
    --activation sigmoid

表示:
  上段: 教師関数・学習前出力・現在の出力
  下段: 各hidden nodeの出力層への寄与

実行例
------
python nn_activation_animation.py --activation relu
python nn_activation_animation.py --activation tanh
python nn_activation_animation.py --activation sigmoid

比較しやすい例:
python nn_activation_animation.py --activation relu    --nnode 12 --lr 0.03
python nn_activation_animation.py --activation tanh    --nnode 12 --lr 0.03
python nn_activation_animation.py --activation sigmoid --nnode 12 --lr 0.03
"""

import argparse
import numpy as np
import matplotlib.pyplot as plt


# ============================================================
# 教師関数
# ============================================================
def target_function(x):
    """教師関数 y = sin(pi x)"""
    x = np.asarray(x, dtype=float)
    return np.sin(np.pi * x)


# ============================================================
# 活性化関数
# ============================================================
def relu(x):
    return np.maximum(0.0, x)


def relu_derivative(x):
    return (x > 0.0).astype(float)


def sigmoid(x):
    """
    sigmoid(x) = 1 / (1 + exp(-x))

    exp() のoverflowを避けるため範囲を制限する。
    """
    x = np.clip(x, -60.0, 60.0)
    return 1.0 / (1.0 + np.exp(-x))


def sigmoid_derivative(x):
    s = sigmoid(x)
    return s * (1.0 - s)


def tanh(x):
    return np.tanh(x)


def tanh_derivative(x):
    t = np.tanh(x)
    return 1.0 - t * t


def activation(x, name):
    """指定された活性化関数を計算する。"""
    if name == "relu":
        return relu(x)
    elif name == "tanh":
        return tanh(x)
    elif name == "sigmoid":
        return sigmoid(x)
    else:
        raise ValueError(f"unknown activation: {name}")


def activation_derivative(x, name):
    """指定された活性化関数の微分を計算する。"""
    if name == "relu":
        return relu_derivative(x)
    elif name == "tanh":
        return tanh_derivative(x)
    elif name == "sigmoid":
        return sigmoid_derivative(x)
    else:
        raise ValueError(f"unknown activation: {name}")


# ============================================================
# 1入力1出力, 1隠れ層 NN
# ============================================================
class SimpleNN1D:
    """
    1入力1出力, 1隠れ層ニューラルネットワーク。

        z1 = x * W1 + b1
        h  = activation(z1)
        y  = h @ W2 + b2
    """

    def __init__(self, nnode=12, activation_name="relu", seed=0):
        self.nin = 1
        self.nout = 1
        self.nnode = nnode
        self.activation_name = activation_name

        rng = np.random.default_rng(seed)

        # 教育用の簡単な初期化。
        # ReLUではHe初期化、tanh/sigmoidではXavier系のスケール。
        if activation_name == "relu":
            scale1 = np.sqrt(2.0 / self.nin)
        else:
            scale1 = np.sqrt(1.0 / self.nin)

        self.W1 = rng.normal(0.0, scale1, size=(1, nnode))
        self.b1 = np.zeros(nnode, dtype=float)

        self.W2 = rng.normal(
            0.0,
            np.sqrt(1.0 / nnode),
            size=(nnode, 1)
        )
        self.b2 = np.zeros(1, dtype=float)

        self.X = None
        self.z1 = None
        self.h = None
        self.y = None

    def forward(self, X):
        """順伝播"""
        self.X = X
        self.z1 = X @ self.W1 + self.b1

        self.h = activation(
            self.z1,
            self.activation_name
        )

        self.y = self.h @ self.W2 + self.b2

        return self.y

    def mse_loss(self, y_true):
        return np.mean((self.y - y_true) ** 2)

    def backward(self, y_true):
        """
        逆伝播

        Loss = mean((y - y_true)^2)
        """
        batch_size = y_true.shape[0]

        # 出力層の誤差
        dY = 2.0 * (self.y - y_true) / batch_size

        # hidden -> output
        dW2 = self.h.T @ dY
        db2 = np.sum(dY, axis=0)

        # hidden layer へ誤差を伝播
        dH = dY @ self.W2.T

        # 活性化関数の微分
        dZ1 = dH * activation_derivative(
            self.z1,
            self.activation_name
        )

        # input -> hidden
        dW1 = self.X.T @ dZ1
        db1 = np.sum(dZ1, axis=0)

        return dW1, db1, dW2, db2

    def update(self, grads, lr):
        """勾配降下法"""
        dW1, db1, dW2, db2 = grads

        self.W1 -= lr * dW1
        self.b1 -= lr * db1
        self.W2 -= lr * dW2
        self.b2 -= lr * db2

    def predict(self, X):
        return self.forward(X)

    def contributions(self, X):
        """
        各hidden nodeが最終出力へ与える寄与。

        contribution_j(x)
          = W2[j] * activation(W1[j]*x + b1[j])

        最終出力:
          y(x) = sum_j contribution_j(x) + b2
        """
        z1 = X @ self.W1 + self.b1
        h = activation(z1, self.activation_name)

        return h * self.W2[:, 0]


# ============================================================
# データ
# ============================================================
def make_data(nsample=80, xmin=-1.0, xmax=1.0):
    x = np.linspace(xmin, xmax, nsample).reshape(-1, 1)
    y = target_function(x)
    return x, y


# ============================================================
# 描画
# ============================================================
def prepare_figure(
    x_train,
    y_train,
    x_plot,
    y_true_plot,
    y_init_plot,
    model
):
    fig, (ax_out, ax_node) = plt.subplots(
        2, 1,
        figsize=(8, 10)
    )

    # --------------------------------------------------------
    # 上段: 最終出力
    # --------------------------------------------------------
    ax_out.set_title(
        f"NN approximation to sin(pi x) "
        f"[activation = {model.activation_name}]"
    )
    ax_out.set_xlabel("x")
    ax_out.set_ylabel("y")
    ax_out.grid(True)

    ax_out.plot(
        x_plot[:, 0],
        y_true_plot[:, 0],
        label="target: sin(pi x)"
    )

    ax_out.scatter(
        x_train[:, 0],
        y_train[:, 0],
        s=18,
        label="training data"
    )

    line_init, = ax_out.plot(
        x_plot[:, 0],
        y_init_plot[:, 0],
        label="initial prediction"
    )

    line_pred, = ax_out.plot(
        x_plot[:, 0],
        y_init_plot[:, 0],
        label="current prediction"
    )

    text_info = ax_out.text(
        0.02,
        0.98,
        "",
        transform=ax_out.transAxes,
        va="top"
    )

    ax_out.legend(loc="lower left")

    # --------------------------------------------------------
    # 下段: node寄与
    # --------------------------------------------------------
    ax_node.set_title(
        f"Contribution of each hidden node "
        f"[{model.activation_name}]"
    )
    ax_node.set_xlabel("x")
    ax_node.set_ylabel("contribution to output")
    ax_node.grid(True)

    contrib0 = model.contributions(x_plot)

    contrib_lines = []

    for j in range(model.nnode):
        line_j, = ax_node.plot(
            x_plot[:, 0],
            contrib0[:, j],
            label=f"node {j}"
        )
        contrib_lines.append(line_j)

    line_sum, = ax_node.plot(
        x_plot[:, 0],
        np.sum(contrib0, axis=1) + model.b2[0],
        linewidth=2.5,
        label="sum + b2"
    )

    ax_node.legend(
        loc="upper right",
        ncol=2,
        fontsize=8
    )

    fig.tight_layout()

    return (
        fig,
        ax_out,
        ax_node,
        line_init,
        line_pred,
        text_info,
        contrib_lines,
        line_sum
    )


def save_current_figure(fig, prefix, epoch):
    fig.savefig(
        f"{prefix}_epoch_{epoch:05d}.png",
        dpi=150,
        bbox_inches="tight"
    )


# ============================================================
# 学習 + アニメーション
# ============================================================
def animate_training(args):
    x_train, y_train = make_data(
        args.nsample,
        args.xmin,
        args.xmax
    )

    x_plot = np.linspace(
        args.xmin,
        args.xmax,
        args.nplot
    ).reshape(-1, 1)

    y_true_plot = target_function(x_plot)

    model = SimpleNN1D(
        nnode=args.nnode,
        activation_name=args.activation,
        seed=args.seed
    )

    y_init_plot = model.predict(x_plot).copy()

    (
        fig,
        ax_out,
        ax_node,
        line_init,
        line_pred,
        text_info,
        contrib_lines,
        line_sum
    ) = prepare_figure(
        x_train,
        y_train,
        x_plot,
        y_true_plot,
        y_init_plot,
        model
    )

    plt.ion()
    plt.show()

    losses = []

    # --------------------------------------------------------
    # epoch = 0
    # --------------------------------------------------------
    y_now = model.predict(x_plot)

    loss0 = np.mean(
        (model.predict(x_train) - y_train) ** 2
    )

    losses.append(loss0)

    line_pred.set_ydata(y_now[:, 0])

    contrib_now = model.contributions(x_plot)

    for j, line_j in enumerate(contrib_lines):
        line_j.set_ydata(contrib_now[:, j])

    line_sum.set_ydata(
        np.sum(contrib_now, axis=1) + model.b2[0]
    )

    text_info.set_text(
        f"activation = {args.activation}\n"
        f"epoch = 0\n"
        f"loss = {loss0:.6e}"
    )

    fig.canvas.draw()
    fig.canvas.flush_events()

    if args.save_prefix:
        save_current_figure(
            fig,
            args.save_prefix,
            0
        )

    # --------------------------------------------------------
    # 学習ループ
    # --------------------------------------------------------
    for epoch in range(1, args.epochs + 1):

        # 1. 順伝播
        model.forward(x_train)

        # 2. 逆伝播
        grads = model.backward(y_train)

        # 3. パラメータ更新
        model.update(grads, args.lr)

        # 4. loss
        model.forward(x_train)
        loss = model.mse_loss(y_train)
        losses.append(loss)

        # ----------------------------------------------------
        # グラフ更新
        # ----------------------------------------------------
        if (
            epoch % args.plot_every == 0
            or epoch == 1
            or epoch == args.epochs
        ):
            y_plot = model.predict(x_plot)

            line_pred.set_ydata(y_plot[:, 0])

            contrib = model.contributions(x_plot)

            for j, line_j in enumerate(contrib_lines):
                line_j.set_ydata(contrib[:, j])

            sum_curve = (
                np.sum(contrib, axis=1)
                + model.b2[0]
            )

            line_sum.set_ydata(sum_curve)

            text_info.set_text(
                f"activation = {args.activation}\n"
                f"epoch = {epoch}\n"
                f"loss = {loss:.6e}\n"
                f"nnode = {args.nnode}\n"
                f"lr = {args.lr}"
            )

            # node寄与グラフのy範囲調整
            ymin = min(
                np.min(contrib),
                np.min(sum_curve)
            )
            ymax = max(
                np.max(contrib),
                np.max(sum_curve)
            )

            span = ymax - ymin

            if span < 1e-8:
                span = 1.0

            ax_node.set_ylim(
                ymin - 0.1 * span,
                ymax + 0.1 * span
            )

            fig.canvas.draw()
            fig.canvas.flush_events()

            plt.pause(args.pause)

            if args.save_prefix:
                save_current_figure(
                    fig,
                    args.save_prefix,
                    epoch
                )

    # --------------------------------------------------------
    # loss履歴
    # --------------------------------------------------------
    fig_loss, ax_loss = plt.subplots()

    ax_loss.set_title(
        f"Training loss [{args.activation}]"
    )
    ax_loss.set_xlabel("epoch")
    ax_loss.set_ylabel("MSE loss")
    ax_loss.grid(True)

    ax_loss.plot(
        np.arange(len(losses)),
        losses
    )

    print("Training finished.")
    print(f"activation = {args.activation}")
    print(f"final loss = {losses[-1]:.8e}")

    plt.ioff()
    plt.show()


# ============================================================
# CLI
# ============================================================
def parse_args():
    parser = argparse.ArgumentParser(
        description=(
            "sin(pi x) を1隠れ層NNで近似し、"
            "活性化関数を比較する教育用プログラム"
        )
    )

    parser.add_argument(
        "--activation",
        choices=[
            "relu",
            "tanh",
            "sigmoid"
        ],
        default="relu",
        help="隠れ層の活性化関数"
    )

    parser.add_argument(
        "--nnode",
        type=int,
        default=12,
        help="隠れ層ノード数"
    )

    parser.add_argument(
        "--nsample",
        type=int,
        default=80,
        help="学習点数"
    )

    parser.add_argument(
        "--nplot",
        type=int,
        default=400,
        help="描画用x分割数"
    )

    parser.add_argument(
        "--epochs",
        type=int,
        default=1000,
        help="学習エポック数"
    )

    parser.add_argument(
        "--lr",
        type=float,
        default=0.03,
        help="学習率"
    )

    parser.add_argument(
        "--xmin",
        type=float,
        default=-1.0
    )

    parser.add_argument(
        "--xmax",
        type=float,
        default=1.0
    )

    parser.add_argument(
        "--seed",
        type=int,
        default=0
    )

    parser.add_argument(
        "--plot-every",
        type=int,
        default=5,
        help="何epochごとに表示更新するか"
    )

    parser.add_argument(
        "--pause",
        type=float,
        default=0.01,
        help="表示更新ごとの待ち時間[秒]"
    )

    parser.add_argument(
        "--save-prefix",
        type=str,
        default="",
        help="指定すると更新時のPNGを保存"
    )

    return parser.parse_args()


def main():
    args = parse_args()
    animate_training(args)


if __name__ == "__main__":
    main()
