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

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

修正版:
- node寄与グラフと出力グラフを、1つのFigure内の別subplotに表示
  上段: 出力曲線
  下段: 各node寄与
- 環境によって別ウィンドウ表示が不安定な場合でも見やすい構成

実行例
------
python nn_relu_animation_subplot.py
python nn_relu_animation_subplot.py --nnode 16
python nn_relu_animation_subplot.py --plot-every 10
python nn_relu_animation_subplot.py --save-prefix relu_demo
"""

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)


class SimpleReLUNN1D:
    """
    1入力1出力, 1隠れ層 ReLU NN
    """

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

        rng = np.random.default_rng(seed)

        self.W1 = rng.normal(0.0, np.sqrt(2.0 / self.nin), size=(1, nnode))
        self.b1 = np.zeros(nnode, dtype=float)

        self.W2 = rng.normal(0.0, np.sqrt(2.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 = relu(self.z1)
        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):
        batch_size = y_true.shape[0]

        dY = 2.0 * (self.y - y_true) / batch_size

        dW2 = self.h.T @ dY
        db2 = np.sum(dY, axis=0)

        dH = dY @ self.W2.T
        dZ1 = dH * relu_derivative(self.z1)

        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):
        """
        各隠れノードの寄与:
            contribution_j(x) = W2[j,0] * ReLU(W1[0,j] * x + b1[j])
        """
        z1 = X @ self.W1 + self.b1
        h = relu(z1)
        contrib = h * self.W2[:, 0]
        return contrib


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):
    """
    1つのFigureに2つのsubplotを作る
      ax_out    : 出力曲線
      ax_node   : 各node寄与
    """
    fig, (ax_out, ax_node) = plt.subplots(2, 1, figsize=(8, 10))

    # -------------------------------
    # 上段: 出力曲線
    # -------------------------------
    ax_out.set_title("ReLU NN approximation to sin(pi x)")
    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("Contribution of each ReLU node")
    ax_node.set_xlabel("x")
    ax_node.set_ylabel("contribution")
    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],
        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 = SimpleReLUNN1D(nnode=args.nnode, 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 = []

    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"epoch = 0\nloss = {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):
        model.forward(x_train)
        grads = model.backward(y_train)
        model.update(grads, args.lr)

        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"epoch = {epoch}\n"
                f"loss = {loss:.6e}\n"
                f"nnode = {args.nnode}\n"
                f"lr = {args.lr}"
            )

            # 下段の縦軸自動調整
            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("Training loss history")
    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"final loss = {losses[-1]:.8e}")

    plt.ioff()
    plt.show()


def parse_args():
    parser = argparse.ArgumentParser(
        description="sin(pi x) を ReLU NN で近似し、出力とnode寄与を別subplotで表示する教育用プログラム"
    )

    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,
                        help="x最小")
    parser.add_argument("--xmax", type=float, default=1.0,
                        help="x最大")
    parser.add_argument("--seed", type=int, default=0,
                        help="乱数シード")
    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()
