{ "cells": [ { "cell_type": "code", "execution_count": 14, "id": "17ebc670", "metadata": {}, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "The autoreload extension is already loaded. To reload it, use:\n", " %reload_ext autoreload\n" ] } ], "source": [ "%load_ext autoreload\n", "%autoreload 2" ] }, { "cell_type": "code", "execution_count": 15, "id": "989b4b5b", "metadata": {}, "outputs": [], "source": [ "import numpy as np\n", "import torch\n", "import matplotlib.pyplot as plt\n", "\n", "from plt_utils import plot, savefig\n", "\n", "import os\n", "\n", "task_name, data_name = \"var_copy\", \"data_100_5_30\"" ] }, { "cell_type": "code", "execution_count": 16, "id": "e0f9f46f", "metadata": {}, "outputs": [], "source": [ "dashed_task_name = '-'.join(task_name.split('_'))\n", "\n", "all_losses = {}\n", "all_accs = {}\n", "all_params = {}\n", "all_eval_losses = {}\n", "all_eval_accs = {}" ] }, { "cell_type": "code", "execution_count": 17, "id": "f708d8a0", "metadata": {}, "outputs": [], "source": [ "for run_name in os.listdir(\"results_ood/\" + task_name + \"/\" + data_name):\n", " if len(run_name.split(\"_\")) != 10: continue\n", "\n", " for run in os.listdir(\"results_ood/\" + task_name + \"/\" + data_name + \"/\" + run_name):\n", " if run.startswith(\".\"):\n", " continue\n", " \n", " # print(\"results_ood/\" + task_name + \"/\" + data_name + \"/\" + run_name + \"/\" + run)\n", " data = torch.load(\"results_ood/\" + task_name + \"/\" + data_name + \"/\" + run_name + \"/\" + run, weights_only=False)\n", " losses = data[\"losses\"]\n", " accs = data[\"accs\"]\n", " # args = data[\"args\"]\n", " params = data[\"param_count\"]\n", " eval_loss = data[\"eval_loss\"]\n", " eval_acc = data[\"eval_acc\"].item()\n", "\n", " if run_name not in all_losses.keys():\n", " all_losses[run_name] = losses.unsqueeze(0)\n", " all_eval_losses[run_name] = [eval_loss]\n", " all_accs[run_name] = accs.unsqueeze(0)\n", " all_eval_accs[run_name] = [eval_acc]\n", " all_params[run_name] = params\n", " else:\n", " all_accs[run_name] = torch.cat((all_accs[run_name], accs.unsqueeze(0)))\n", " all_eval_accs[run_name].append(eval_acc)\n", " all_eval_losses[run_name].append(eval_loss)\n", " all_losses[run_name] = torch.cat((all_losses[run_name], losses.unsqueeze(0)))\n", "\n", "for run_name in all_losses.keys():\n", " all_losses[run_name] = all_losses[run_name].detach().numpy()\n", " all_accs[run_name] = all_accs[run_name].detach().numpy()\n", " all_eval_losses[run_name] = np.array(all_eval_losses[run_name])\n", " all_eval_accs[run_name] = np.array(all_eval_accs[run_name])" ] }, { "cell_type": "code", "execution_count": 18, "id": "2833ca7a", "metadata": {}, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "run_var-copy_SSM_SSM_w100_d12_nh1_sd1_nn5_nv30 0.43282261 0.39877721514891495\n", "run_var-copy_SSM_SSM_w100_d15_nh1_sd1_nn5_nv30 0.61931986 0.5685863576152108\n", "run_var-copy_TF_SSM_w100_d2_nh1_sd1_nn5_nv30 0.060773052 0.10188358480280096\n", "run_var-copy_TF_SSM_w100_d12_nh1_sd1_nn5_nv30 0.21031484 0.24504086121239446\n", "run_var-copy_TF_TF_w100_d2_nh1_sd1_nn5_nv30 0.03534364 0.02636057697236538\n", "run_var-copy_TF_TF_w100_d15_nh1_sd1_nn5_nv30 0.41554186 0.5049803880127993\n", "run_var-copy_SSM_TF_w100_d6_nh1_sd1_nn5_nv30 0.11954223 0.1465694470839067\n", "run_var-copy_SSM_SSM_w100_d10_nh1_sd1_nn5_nv30 0.6125856 0.5405284464359283\n", "run_var-copy_SSM_TF_w100_d8_nh1_sd1_nn5_nv30 0.45219678 0.4328746416352012\n", "run_var-copy_SSM_SSM_w100_d2_nh1_sd1_nn5_nv30 0.0787262 0.11496004902503708\n", "run_var-copy_TF_SSM_w100_d4_nh1_sd1_nn5_nv30 0.1014059 0.13624397258866916\n", "run_var-copy_SSM_TF_w100_d20_nh1_sd1_nn5_nv30 0.67468256 0.6649037870493802\n", "run_var-copy_SSM_TF_w100_d15_nh1_sd1_nn5_nv30 0.57955116 0.5650092546235431\n", "run_var-copy_TF_TF_w100_d8_nh1_sd1_nn5_nv30 0.24238192 0.3181798972866752\n", "run_var-copy_SSM_SSM_w100_d4_nh1_sd1_nn5_nv30 0.38514006 0.3294621326706626\n", "run_var-copy_SSM_SSM_w100_d24_nh1_sd1_nn5_nv30 0.72951454 0.6708598299459978\n", "run_var-copy_TF_SSM_w100_d10_nh1_sd1_nn5_nv30 0.26325858 0.29859892143444583\n", "run_var-copy_SSM_TF_w100_d4_nh1_sd1_nn5_nv30 0.08546053 0.12569612569429658\n", "run_var-copy_TF_TF_w100_d20_nh1_sd1_nn5_nv30 0.53927666 0.6394596750086005\n", "run_var-copy_SSM_TF_w100_d12_nh1_sd1_nn5_nv30 0.55815214 0.45932409302754834\n", "run_var-copy_TF_TF_w100_d24_nh1_sd1_nn5_nv30 0.69330955 0.7885203632441434\n", "run_var-copy_TF_SSM_w100_d15_nh1_sd1_nn5_nv30 0.3340461 0.3105349088595672\n", "run_var-copy_SSM_TF_w100_d24_nh1_sd1_nn5_nv30 0.78975546 0.6099584522572431\n", "run_var-copy_SSM_SSM_w100_d20_nh1_sd1_nn5_nv30 0.029622627 0.019936622882431202\n", "run_var-copy_TF_SSM_w100_d6_nh1_sd1_nn5_nv30 0.1267125 0.14985936405983838\n", "run_var-copy_TF_TF_w100_d6_nh1_sd1_nn5_nv30 0.09642029 0.14757748219099912\n", "run_var-copy_TF_SSM_w100_d20_nh1_sd1_nn5_nv30 0.46733 0.3361467854543166\n", "run_var-copy_SSM_SSM_w100_d8_nh1_sd1_nn5_nv30 0.55762124 0.5054087720134042\n", "run_var-copy_TF_TF_w100_d4_nh1_sd1_nn5_nv30 0.06929611 0.08527342568744313\n", "run_var-copy_SSM_TF_w100_d2_nh1_sd1_nn5_nv30 0.06488876 0.10887749154459346\n", "run_var-copy_SSM_SSM_w100_d6_nh1_sd1_nn5_nv30 0.31115448 0.3066261898387562\n", "run_var-copy_TF_SSM_w100_d8_nh1_sd1_nn5_nv30 0.22431688 0.2580069218846885\n", "run_var-copy_TF_SSM_w100_d24_nh1_sd1_nn5_nv30 0.48085716 0.3478721175342798\n", "run_var-copy_TF_TF_w100_d10_nh1_sd1_nn5_nv30 0.26761094 0.3145741671323776\n", "run_var-copy_SSM_TF_w100_d10_nh1_sd1_nn5_nv30 0.6060383 0.5851903490044854\n", "run_var-copy_TF_TF_w100_d12_nh1_sd1_nn5_nv30 0.2751147 0.3406459783965891\n" ] } ], "source": [ "for run_name in all_losses.keys():\n", " layer1 = run_name.split(\"_\")[2]\n", " layer2 = run_name.split(\"_\")[3]\n", " # if layer1 == \"SSM\" and layer2 == \"SSM\":\n", " print(run_name, np.mean(all_accs[run_name][:,-1]), np.mean(all_eval_accs[run_name]))" ] }, { "cell_type": "code", "execution_count": null, "id": "2f1f9878", "metadata": {}, "outputs": [], "source": [] } ], "metadata": { "kernelspec": { "display_name": "hybrid", "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.12.11" } }, "nbformat": 4, "nbformat_minor": 5 }