#!/usr/bin/env python3
"""
simple_nn.py
============

教育用の1隠れ層ニューラルネットワーク。

構造:
    vin (Nin)
      |
      | W1, b1
      v
    Hidden (Nnode) -- ReLU
      |
      | W2, b2
      v
    vout (Nout) -- Linear

外部NNライブラリは使わず、NumPyだけで
  1. 順伝播
  2. 損失計算
  3. 逆伝播
  4. 勾配降下法によるパラメータ更新
を実装する。

使用例
------
学習:
    python simple_nn.py --mode train --nin 2 --nout 2 --nnode 32

予測:
    python simple_nn.py --mode predict --model nn_model.npz --vin 0.2 -0.4

学習対象の関数は target_function() を編集して変更できる。
"""

import argparse
import numpy as np


# ============================================================
# 学習させる教師関数 vout = f(vin)
# ============================================================
def target_function(vin, nout):
    """
    教師データを作るための例示関数。

    Parameters
    ----------
    vin : ndarray, shape = (Nsamples, Nin)
        入力ベクトル群。
    nout : int
        出力次元。

    Returns
    -------
    vout : ndarray, shape = (Nsamples, Nout)

    Notes
    -----
    教育用に、線形ではない滑らかな関数を使っている。
    実際の問題では、この関数だけを目的の f(vin) に置き換える。

    例:
        Nin=2, Nout=1 なら

        vout[:, 0] = sin(vin[:, 0]) + 0.5 * vin[:, 1]**2

    のように書いてよい。
    """
    vin = np.asarray(vin, dtype=float)

    if vin.ndim == 1:
        vin = vin.reshape(1, -1)

    nsample, nin = vin.shape
    vout = np.zeros((nsample, nout), dtype=float)

    # Nin, Nout が任意でも動くサンプル関数
    for k in range(nout):
        for i in range(nin):
            vout[:, k] += np.sin((i + k + 1) * vin[:, i]) / (i + 1)
            vout[:, k] += 0.15 * (k + 1) * vin[:, i] ** 2

        vout[:, k] /= nin

    return vout


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


def relu_derivative(x):
    """ReLU の微分。x > 0 のとき1、それ以外0。"""
    return (x > 0.0).astype(float)


# ============================================================
# 1隠れ層ニューラルネットワーク
# ============================================================
class SimpleNN:
    """
    1隠れ層の全結合ニューラルネットワーク。

    数式
    ----
    z1 = X W1 + b1
    h  = ReLU(z1)
    y  = h W2 + b2

    出力層は回帰問題を想定して線形活性化。
    """

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

        rng = np.random.default_rng(seed)

        # ReLU に適した He 初期化
        self.W1 = rng.normal(
            loc=0.0,
            scale=np.sqrt(2.0 / nin),
            size=(nin, nnode),
        )
        self.b1 = np.zeros(nnode, dtype=float)

        self.W2 = rng.normal(
            loc=0.0,
            scale=np.sqrt(2.0 / nnode),
            size=(nnode, nout),
        )
        self.b2 = np.zeros(nout, dtype=float)

        # forward() 時の中間値を保存する領域
        self.X = None
        self.z1 = None
        self.h = None
        self.y = None

    def forward(self, X):
        """
        順伝播。

        X        : (batch, Nin)
        z1       : (batch, Nnode)
        h        : (batch, Nnode)
        y        : (batch, Nout)
        """
        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

    @staticmethod
    def mse_loss(y, target):
        """平均二乗誤差 MSE。"""
        return np.mean((y - target) ** 2)

    def backward(self, target):
        """
        逆伝播により各パラメータの勾配を求める。

        Loss = mean((y - target)^2)

        逆伝播:
            dL/dy
              ↓
            dL/dW2, dL/db2
              ↓
            dL/dh
              ↓ ReLU'
            dL/dz1
              ↓
            dL/dW1, dL/db1
        """
        batch_size = target.shape[0]

        # MSE:
        # L = 1/(batch_size*Nout) * sum((y-target)^2)
        dY = 2.0 * (self.y - target) / (batch_size * self.nout)

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

        # 隠れ層へ誤差を伝播
        dH = dY @ self.W2.T

        # ReLU の微分を掛ける
        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, gradients, learning_rate):
        """単純な勾配降下法 Gradient Descent で更新する。"""
        dW1, db1, dW2, db2 = gradients

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

    def predict(self, X):
        """予測値を返す。"""
        return self.forward(X)

    def save(self, filename):
        """学習済みモデルを npz 形式で保存する。"""
        np.savez(
            filename,
            nin=self.nin,
            nnode=self.nnode,
            nout=self.nout,
            W1=self.W1,
            b1=self.b1,
            W2=self.W2,
            b2=self.b2,
        )

    @classmethod
    def load(cls, filename):
        """保存済みモデルを読み込む。"""
        data = np.load(filename)

        nin = int(data["nin"])
        nnode = int(data["nnode"])
        nout = int(data["nout"])

        model = cls(nin, nnode, nout)

        model.W1 = data["W1"]
        model.b1 = data["b1"]
        model.W2 = data["W2"]
        model.b2 = data["b2"]

        return model


# ============================================================
# 学習データ生成
# ============================================================
def make_training_data(nsample, nin, nout, xmin, xmax, seed):
    """
    vin を一様乱数で生成し、target_function() から vout を作る。
    """
    rng = np.random.default_rng(seed)

    vin = rng.uniform(
        low=xmin,
        high=xmax,
        size=(nsample, nin),
    )

    vout = target_function(vin, nout)

    return vin, vout


# ============================================================
# 学習
# ============================================================
def train(args):
    X, T = make_training_data(
        nsample=args.nsample,
        nin=args.nin,
        nout=args.nout,
        xmin=args.xmin,
        xmax=args.xmax,
        seed=args.seed,
    )

    model = SimpleNN(
        nin=args.nin,
        nnode=args.nnode,
        nout=args.nout,
        seed=args.seed,
    )

    rng = np.random.default_rng(args.seed + 1)

    print("===== Training =====")
    print(f"Nin        = {args.nin}")
    print(f"Nnode      = {args.nnode}")
    print(f"Nout       = {args.nout}")
    print(f"Nsamples   = {args.nsample}")
    print(f"Epochs     = {args.epochs}")
    print(f"Batch size = {args.batch}")
    print(f"LR         = {args.lr}")
    print()

    # --------------------------------------------------------
    # mini-batch 学習
    # --------------------------------------------------------
    for epoch in range(1, args.epochs + 1):

        # 毎epochでデータ順序をランダム化
        indices = rng.permutation(args.nsample)

        for start in range(0, args.nsample, args.batch):
            batch_idx = indices[start:start + args.batch]

            xb = X[batch_idx]
            tb = T[batch_idx]

            # 1. 順伝播
            yb = model.forward(xb)

            # 2. 逆伝播
            gradients = model.backward(tb)

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

        # 全学習データに対する損失を表示
        if epoch == 1 or epoch % args.print_every == 0 or epoch == args.epochs:
            y_all = model.predict(X)
            loss = model.mse_loss(y_all, T)

            print(
                f"epoch {epoch:6d} / {args.epochs:6d}"
                f"    MSE = {loss:.8e}"
            )

    model.save(args.model)

    print()
    print(f"Model saved: {args.model}")

    # --------------------------------------------------------
    # 学習結果の簡単な確認
    # --------------------------------------------------------
    ncheck = min(5, args.nsample)

    print()
    print("===== Training sample check =====")

    for i in range(ncheck):
        x = X[i:i + 1]
        target = T[i]
        pred = model.predict(x)[0]

        print(f"vin    = {x[0]}")
        print(f"target = {target}")
        print(f"pred   = {pred}")
        print()


# ============================================================
# 予測
# ============================================================
def predict(args):
    model = SimpleNN.load(args.model)

    vin = np.asarray(args.vin, dtype=float)

    if vin.size != model.nin:
        raise ValueError(
            f"入力数が一致しません: "
            f"モデル Nin={model.nin}, "
            f"--vin の要素数={vin.size}"
        )

    vin = vin.reshape(1, -1)

    vout = model.predict(vin)[0]

    print("===== Prediction =====")
    print(f"model = {args.model}")
    print(f"vin   = {vin[0]}")
    print(f"vout  = {vout}")

    # 教育用: 真値も計算可能なので比較する
    if args.show_target:
        target = target_function(vin, model.nout)[0]
        print(f"f(vin)= {target}")
        print(f"error = {vout - target}")


# ============================================================
# コマンドライン
# ============================================================
def parse_args():
    parser = argparse.ArgumentParser(
        description="NumPyだけで実装した教育用1隠れ層NN"
    )

    parser.add_argument(
        "--mode",
        required=True,
        choices=["train", "predict"],
        help="train: 学習, predict: 予測",
    )

    parser.add_argument(
        "--model",
        default="nn_model.npz",
        help="モデル保存/読込ファイル",
    )

    # 学習時
    parser.add_argument("--nin", type=int, default=2)
    parser.add_argument("--nout", type=int, default=1)
    parser.add_argument("--nnode", type=int, default=32)

    parser.add_argument("--nsample", type=int, default=2000)
    parser.add_argument("--epochs", type=int, default=3000)
    parser.add_argument("--batch", type=int, default=64)
    parser.add_argument("--lr", type=float, default=0.02)

    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("--print-every", type=int, default=100)

    # 予測時
    parser.add_argument(
        "--vin",
        nargs="+",
        type=float,
        default=None,
        help="predict時の入力値。例: --vin 0.2 -0.4",
    )

    parser.add_argument(
        "--show-target",
        action="store_true",
        help="predict時に教師関数 f(vin) の真値も表示",
    )

    return parser.parse_args()


def main():
    args = parse_args()

    if args.mode == "train":
        train(args)

    elif args.mode == "predict":
        if args.vin is None:
            raise ValueError(
                "predict モードでは --vin を指定してください。"
            )

        predict(args)


if __name__ == "__main__":
    main()
