#!/usr/bin/env python

##############################################################################
# This file implements the neural network training using Galerkin POD as
# described in the paper:
#
# Fully discrete analysis of the Galerkin POD neural network approximation
# with application to 3D acoustic wave scattering. 2025.
# Authors: J. Dölz, F. Henríquez
# Submitted to SIAM Journal on Scientific Computing
# Arxiv: https://arxiv.org/abs/2502.01859
#
# Call runit.sh to run it with the required arguments. 
#
# For all details and theoretical background we refer to the paper.
##############################################################################


##############################################################################
# Helper function
##############################################################################

def flushit(out):
    print(out, flush=True)

##############################################################################
# Import packages
##############################################################################

flushit("start")

import sys
import h5py
import numpy as np
from scipy import linalg
from pytictoc import TicToc
import torch
import torch.nn as nn
import torch.optim as optim

flushit("packages loaded")

##############################################################################
# Define abstract neural network architecture
##############################################################################

class TanhNeuralNetworkBatchNormalization(nn.Module):
    def __init__(self, input_size, output_size, width, depth):
        super(TanhNeuralNetworkBatchNormalization, self).__init__()
        
        self.layers = nn.ModuleList()
        self.layers.append(nn.Linear(input_size, width))
        self.layers.append(nn.BatchNorm1d(width))
        for _ in range(depth - 2):
            self.layers.append(nn.Linear(width, width))
            self.layers.append(nn.BatchNorm1d(width))
        self.layers.append(nn.Linear(width, output_size))
        self.layers.append(nn.BatchNorm1d(output_size))
        
        self.activation = nn.Tanh()
    
    def forward(self, x):
        for layer in self.layers[:-1]:
            x = self.activation(layer(x))
        x = self.layers[-1](x)
        return x

##############################################################################
# Load raw data
##############################################################################

path_data = sys.argv[1]
path_results = sys.argv[2]

assert(str(path_data) != str(path_results))

flushit("Data path: " + str(path_data))
flushit("Data results: " + str(path_results))

with h5py.File(path_data, 'r') as f:
    D_HaltonSamples = f['/HaltonSamples'][:]
    D_HaltonPoints = f['/HaltonPoints'][:]
    D_Dimension = f['/Dimension']

d_X = D_HaltonPoints.shape[0]
D_NumberOfSamples = D_HaltonPoints.shape[1]

flushit("Raw data loaded, number of samples: " + str(D_NumberOfSamples))

##############################################################################
# Neural network training using Galerkin POD
##############################################################################

# Check if GPU is available
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
flushit(f"Using device: {device}")

# data processing

number_of_samples = D_HaltonPoints.shape[1]
maximal_training_level = int(np.log(number_of_samples)/np.log(2) - 1)

L_levels = list()
L_number_of_training_samples = list()
L_t_svd = list()
L_t_training = list()
L_decay_POD_singular_values = list()
L_POD_generalization_error = list()
L_POD_tol = list()
L_POD_rank = list()
L_NN_epochs = list()
L_NN_training_error = list()
L_NN_generalization_error = list()
L_NN_n = list()
L_NN_width = list()
L_NN_depth = list()
L_total_generalization_error = list()

# iterate over number of training samples, in the paper this is N
for level in np.arange(5,maximal_training_level+1):
    number_of_training_samples = 2**level
    if number_of_training_samples == number_of_samples:
        break;
    flushit("===== Level " + str(level) + " ===============================")
    L_levels.append(level)
    flushit("Number of samples: " + str(number_of_training_samples))
    L_number_of_training_samples.append(number_of_training_samples)

    # compute SVD
    tsvd = TicToc()
    tsvd.tic()
    U, S, Vh = linalg.svd(D_HaltonSamples[:,:number_of_training_samples], full_matrices=False)
    tsvd.toc()
    L_t_svd.append(tsvd.tocvalue())
    L_decay_POD_singular_values.append(S)

    # Determine reduced basis rank
    tol = 0.01/np.sqrt(number_of_training_samples)
    flushit("POD tolerance: " + str(tol))
    L_POD_tol.append(tol)
    for POD_rank in np.arange(1,len(S)):
        if linalg.norm(S[POD_rank:])/linalg.norm(S[0]) < tol:
            break
    flushit("Reduced basis rank: " + str(POD_rank))
    L_POD_rank.append(POD_rank)

    # compute generalization error of POD
    POD_generalization_error = np.sqrt(np.mean(np.sum(np.abs(D_HaltonSamples - np.dot(U[:, :POD_rank], np.dot(np.conj(U[:, :POD_rank].T), D_HaltonSamples)))**2, axis=1)))
    flushit("POD generalization error (L2-norm): " + str(POD_generalization_error))
    L_POD_generalization_error.append(POD_generalization_error)
    
    # Setup pytorch training data
    Data_X = D_HaltonPoints[:,:number_of_training_samples].T
    Data_Y = np.dot(np.conj(U[:, :POD_rank].T), D_HaltonSamples[:,:number_of_training_samples])
    Data_Y = np.concatenate((np.real(Data_Y), np.imag(Data_Y)), axis=0).T
    Data_X_pytorch = torch.tensor(Data_X, dtype=torch.float32).to(device)
    Data_Y_pytorch = torch.tensor(Data_Y, dtype=torch.float32).to(device)

    # normalize Y values
    Data_Y_pytorch_min = Data_Y_pytorch.min(dim=0, keepdim=True).values.to(device)
    Data_Y_pytorch_max = Data_Y_pytorch.max(dim=0, keepdim=True).values.to(device)
    Data_Y_pytorch_norm = (2 * (Data_Y_pytorch - Data_Y_pytorch_min) / (Data_Y_pytorch_max - Data_Y_pytorch_min) - 1).to(device)

    # Setup pytorch test data
    Test_X = D_HaltonPoints[:,int(D_NumberOfSamples/2):].T
    Test_Y = np.dot(np.conj(U[:, :POD_rank].T), D_HaltonSamples[:,int(D_NumberOfSamples/2):])
    Test_Y = np.concatenate((np.real(Test_Y), np.imag(Test_Y)), axis=0).T
    Test_X_pytorch = torch.tensor(Test_X, dtype=torch.float32).to(device)
    Test_Y_pytorch = torch.tensor(Test_Y, dtype=torch.float32).to(device)

    L_NN_n.append(list())
    L_NN_width.append(list())
    L_NN_depth.append(list())
    L_NN_epochs.append(list())
    L_NN_training_error.append(list())
    L_NN_generalization_error.append(list())
    L_total_generalization_error.append(list())
    L_t_training.append(list())
    tol = 10./number_of_training_samples
    p = 2./4.5
    for n in np.full((1,), number_of_training_samples**(p/((1-p)*(2-p)))):
        n = int(n)
        torch.manual_seed(0)
        depth = np.max([int(0.5*np.log(n)/np.log(2)), 1])
        width = np.max([int(n**2),1])
        L_NN_n[-1].append(n)
        L_NN_width[-1].append(width)
        L_NN_depth[-1].append(depth)

        # define model and loss function
        model = TanhNeuralNetworkBatchNormalization(d_X, 2 * POD_rank, width, depth).to(device)
        loss_function = nn.MSELoss(reduction='mean')
        flushit("N="+str(number_of_training_samples)+", n="+str(n)+", width="+str(width)+", depth="+str(depth)+", tol="+str(tol))

        optimizer = optim.AdamW(model.parameters(), lr=1e-3, eps=1e-12)
        scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=0.1, patience=10, eps=1e-20)

        # start timer
        t_training = TicToc()
        t_training.tic()
        epoch = 0
        errors = np.zeros(int(1e5+1))
        outputs = model(Data_X_pytorch)
        loss = loss_function(outputs, Data_Y_pytorch_norm)
        errors[0] = loss.item()
        while loss.item() > tol and epoch < 1e5:

            current_lr = optimizer.param_groups[0]['lr']
            alpha = 1
            for param_group in optimizer.param_groups:
                param_group['weight_decay'] = alpha

            if epoch % 1e2 == 0:
                flushit("loss: " + str(loss.item()) + ", learning rate: " + str(current_lr) + ", alpha: " +str(alpha))

            model.train()
            optimizer.zero_grad()
            
            # Forward pass
            outputs = model(Data_X_pytorch)
            loss = loss_function(outputs, Data_Y_pytorch_norm)
            
            # Backward pass and optimization
            loss.backward()
            optimizer.step()
            
            # Step the scheduler
            scheduler.step(loss.item())
            
            epoch = epoch + 1

            errors[epoch] = loss.item()

        # stop timing
        t_training.toc()
        L_t_training[-1].append(t_training.tocvalue())

        L_NN_epochs[-1].append(epoch)
        L_NN_training_error[-1].append(errors)

        outputs = model(Test_X_pytorch)
        outputs = ((outputs + 1) / 2) * (Data_Y_pytorch_max - Data_Y_pytorch_min) + Data_Y_pytorch_min
        loss = loss_function(outputs, Test_Y_pytorch)
        L_NN_generalization_error[-1].append(loss.item())
        L_total_generalization_error[-1].append(np.sqrt(loss.item()+POD_generalization_error**2))
        flushit("DNN with n=" + str(n) + " has training error " + str(errors[epoch]) + " and generalization error " + str(loss.item()) + " after " + str(epoch) + " iterations (loss functional values=squared L2-norm)")
        flushit("Combined absolute L2-error of Galerkin-POD-NN " + str(np.sqrt(loss.item()+POD_generalization_error**2)))

##############################################################################
# Store results
##############################################################################

with h5py.File(path_results, 'w') as f:
    f.create_dataset(name='/d_X', data=d_X, dtype=int)
    f.create_dataset(name='/levels', data=np.asarray(L_levels), dtype=int)
    f.create_dataset(name='/number_of_training_samples', data=np.asarray(L_number_of_training_samples), dtype=int)

    f.create_dataset(name='/POD_tol', data=np.asarray(L_POD_tol), dtype=int)
    f.create_dataset(name='/POD_rank', data=np.asarray(L_POD_rank), dtype=int)
    for L in L_levels:
        f.create_dataset(name='/t_svd/' + str(L), data=L_t_svd[L-L_levels[0]], dtype=float)
    for L in L_levels:
        f.create_dataset(name='/decay_POD_singular_values/' + str(L), data=L_decay_POD_singular_values[L-L_levels[0]], dtype=float)
    for L in L_levels:
        f.create_dataset(name='/POD_generalization_error/' + str(L), data=L_POD_generalization_error[L-L_levels[0]], dtype=float)

    for L in L_levels:
        f.create_dataset(name='/t_training/' + str(L), data=np.asarray(L_t_training[L-L_levels[0]]), dtype=float)
    for L in L_levels:
        f.create_dataset(name='/NN_epochs_error/' + str(L), data=np.asarray(L_NN_epochs[L-L_levels[0]]), dtype=float)
    for L in L_levels:
        f.create_dataset(name='/NN_training_error/' + str(L), data=np.asarray(L_NN_training_error[L-L_levels[0]]), dtype=float)
    for L in L_levels:
        f.create_dataset(name='/NN_generalization_error/' + str(L), data=np.asarray(L_NN_generalization_error[L-L_levels[0]]), dtype=float)
    for L in L_levels:
        f.create_dataset(name='/NN_n/' + str(L), data=np.asarray(L_NN_n[L-L_levels[0]]), dtype=float)
    for L in L_levels:
        f.create_dataset(name='/NN_width/' + str(L), data=np.asarray(L_NN_width[L-L_levels[0]]), dtype=float)
    for L in L_levels:
        f.create_dataset(name='/NN_depth/' + str(L), data=np.asarray(L_NN_depth[L-L_levels[0]]), dtype=float)

    for L in L_levels:
        f.create_dataset(name='/total_generalization_error/' + str(L), data=np.asarray(L_total_generalization_error[L-L_levels[0]]), dtype=float)
