{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# import packages\n",
    "import h5py                     # hdf5 reader\n",
    "import numpy as np              # numpy\n",
    "import matplotlib.pyplot as plt # plotting\n",
    "\n",
    "plt.rcParams.update({\n",
    "    \"pgf.texsystem\": \"pdflatex\",  # Use pdflatex for compatibility\n",
    "    \"text.usetex\": True,           # Use LaTeX for text rendering\n",
    "    \"font.family\": \"serif\",        # Use serif fonts\n",
    "    \"pgf.preamble\": r\"\\usepackage{amsmath}\\usepackage{amssymb}\",  # Add any necessary packages\n",
    "    \"pgf.rcfonts\": False,          # Disable rc fonts to avoid extra font definitions\n",
    "    \"font.size\": 22,\n",
    "    \"figure.figsize\": (12,4.6)\n",
    "})"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "wave1_path = 'training-results/32ksamples_wave1-NN.hdf'\n",
    "wave4_path = 'training-results/32ksamples_wave4-NN.hdf'"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "\n",
    "def load_hdf5_list(file, levels, dataset):\n",
    "    l = list()\n",
    "    for L in levels:\n",
    "        l.append(file[dataset+'/'+str(L)][:])\n",
    "    return l\n",
    "\n",
    "#load data\n",
    "wave1_f = h5py.File(wave1_path, 'r')\n",
    "levels = wave1_f['/levels'][:]\n",
    "number_of_training_samples = wave1_f['/number_of_training_samples'][:]\n",
    "wave1_NN_training_error = load_hdf5_list(wave1_f, levels, '/NN_training_error')\n",
    "wave1_NN_epochs_error = load_hdf5_list(wave1_f, levels, '/NN_epochs_error')\n",
    "wave1_NN_generalization_error = load_hdf5_list(wave1_f, levels, '/NN_generalization_error')\n",
    "wave1_NN_width = load_hdf5_list(wave1_f, levels, '/NN_width')\n",
    "wave1_NN_depth = load_hdf5_list(wave1_f, levels, '/NN_depth')\n",
    "wave1_f.close()\n",
    "wave4_f = h5py.File(wave4_path, 'r')\n",
    "levels = wave4_f['/levels'][:]\n",
    "number_of_training_samples = wave4_f['/number_of_training_samples'][:]\n",
    "wave4_NN_training_error = load_hdf5_list(wave4_f, levels, '/NN_training_error')\n",
    "wave4_NN_epochs_error = load_hdf5_list(wave4_f, levels, '/NN_epochs_error')\n",
    "wave4_NN_generalization_error = load_hdf5_list(wave4_f, levels, '/NN_generalization_error')\n",
    "wave4_NN_width = load_hdf5_list(wave4_f, levels, '/NN_width')\n",
    "wave4_NN_depth = load_hdf5_list(wave4_f, levels, '/NN_depth')\n",
    "wave4_f.close()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# generate strings for legend in following plot\n",
    "labels_number_of_snapshots = list()\n",
    "for i in range(len(levels)):\n",
    "    labels_number_of_snapshots.append(str(number_of_training_samples[i]) + ' snapshots')"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "plt.subplots(constrained_layout=True)\n",
    "wave1_errors = np.zeros(len(levels))\n",
    "wave4_errors = np.zeros(len(levels))\n",
    "for i in range(len(levels)):\n",
    "    wave1_errors[i] = wave1_NN_training_error[i][0][int(wave1_NN_epochs_error[i][0])]\n",
    "    wave4_errors[i] = wave4_NN_training_error[i][0][int(wave4_NN_epochs_error[i][0])]\n",
    "plt.loglog(2**levels,wave1_errors, '+-')\n",
    "plt.loglog(2**levels,wave4_errors, '+-')\n",
    "plt.loglog(2**levels,10.*2.**(-levels),'--',color='black')\n",
    "plt.xlabel('number of training samples $N$')\n",
    "plt.ylabel('$L_{\\mathrm{MSE}}$')\n",
    "plt.title('training error')\n",
    "plt.legend(['$\\kappa=1$', '$\\kappa=4$', r'$10\\cdot N^{-\\alpha}$ (target training error)'])\n",
    "plt.savefig('plots/N-convergence-training.eps')\n",
    "plt.show()\n",
    "\n",
    "plt.subplots(constrained_layout=True)\n",
    "plt.loglog(2**levels,wave1_NN_generalization_error, '+-')\n",
    "plt.loglog(2**levels,wave4_NN_generalization_error, '+-')\n",
    "plt.loglog(2**levels,1e4*2.**(-levels),'--',color='black')\n",
    "plt.loglog(2**levels,1e2*2.**(-levels),'--',color='black')\n",
    "plt.xlabel('number of training samples $N$')\n",
    "plt.ylabel('$L_{\\mathrm{MSE}}$')\n",
    "plt.title('generalization error')\n",
    "plt.legend(['$\\kappa=1$', '$\\kappa=4$', r'rate $N^{-\\alpha}$'])\n",
    "plt.savefig('plots/N-convergence-generalization.eps')\n",
    "plt.show()"
   ]
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "Python 3",
   "language": "python",
   "name": "python3"
  },
  "language_info": {
   "codemirror_mode": {
    "name": "ipython",
    "version": 3
   },
   "file_extension": ".py",
   "mimetype": "text/x-python",
   "name": "python",
   "nbconvert_exporter": "python",
   "pygments_lexer": "ipython3",
   "version": "3.9.2"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 2
}
