From 34fa303aea7496face28145731489b088247fab7 Mon Sep 17 00:00:00 2001 From: mgarbowski Date: Thu, 2 Oct 2025 00:12:16 +0200 Subject: [PATCH 01/14] code for creating plots for experiment 05 --- notebooks/12-plotting_from_wandb_api.ipynb | 712 +++++++++++++++++++++ 1 file changed, 712 insertions(+) create mode 100644 notebooks/12-plotting_from_wandb_api.ipynb diff --git a/notebooks/12-plotting_from_wandb_api.ipynb b/notebooks/12-plotting_from_wandb_api.ipynb new file mode 100644 index 0000000..b1be1d4 --- /dev/null +++ b/notebooks/12-plotting_from_wandb_api.ipynb @@ -0,0 +1,712 @@ +{ + "cells": [ + { + "cell_type": "code", + "execution_count": null, + "id": "a5185de0", + "metadata": {}, + "outputs": [], + "source": [ + "import wandb\n", + "from pathlib import Path\n", + "from pprint import pprint\n", + "import pandas as pd\n", + "import matplotlib.pyplot as plt\n", + "import numpy as np" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "16b3fdc4", + "metadata": {}, + "outputs": [], + "source": [ + "def extract_run_data(runs):\n", + " \"\"\"\n", + " Extract parameters, metrics, and metadata from runs.\n", + " \n", + " Args:\n", + " runs: List of wandb Run objects\n", + " \n", + " Returns:\n", + " pandas.DataFrame with all run data\n", + " \"\"\"\n", + " data = []\n", + " \n", + " for run in runs:\n", + " row = {\n", + " 'run_id': run.id,\n", + " 'run_name': run.name,\n", + " 'state': run.state, # finished, failed, running, etc.\n", + " 'created_at': run.created_at,\n", + " 'runtime': run.summary[\"_runtime\"],\n", + " }\n", + " \n", + " # Add config parameters (hyperparameters)\n", + " for key, value in run.config.items():\n", + " row[f'config_{key}'] = value\n", + " \n", + " # Add summary metrics (final values)\n", + " for key, value in run.summary.items():\n", + " if not key.startswith('_'): # Skip internal wandb fields\n", + " row[f'summary_{key}'] = value\n", + " \n", + " # Add history metrics (you can get specific values)\n", + " history = run.history()\n", + " if not history.empty:\n", + " # Get final values\n", + " for col in history.columns:\n", + " if not col.startswith('_'):\n", + " row[f'final_{col}'] = history[col].iloc[-1] if len(history) > 0 else None\n", + " \n", + " # Get maximum values for accuracy metrics\n", + " accuracy_cols = [col for col in history.columns if 'accuracy' in col.lower()]\n", + " for col in accuracy_cols:\n", + " row[f'max_{col}'] = history[col].max()\n", + " \n", + " data.append(row)\n", + " \n", + " return pd.DataFrame(data)" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "2f7d7b69", + "metadata": {}, + "outputs": [], + "source": [ + "def get_run_history(run, metrics=None):\n", + " \"\"\"\n", + " Get full history for specific metrics from a run.\n", + " \n", + " Args:\n", + " run: wandb Run object\n", + " metrics: List of metric names to retrieve (None for all)\n", + " \n", + " Returns:\n", + " pandas.DataFrame with timestamped metrics\n", + " \"\"\"\n", + " history = run.history()\n", + " \n", + " if metrics:\n", + " # Filter to specific metrics (plus _step and _timestamp)\n", + " available_metrics = [m for m in metrics if m in history.columns]\n", + " cols_to_keep = ['_step', '_timestamp'] + available_metrics\n", + " history = history[cols_to_keep]\n", + " \n", + " history['run_id'] = run.id\n", + " history['run_name'] = run.name\n", + " \n", + " return history" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "6cec58a3", + "metadata": {}, + "outputs": [], + "source": [ + "output_dir = Path(\"results/plots\")\n", + "output_dir.mkdir(parents=True, exist_ok=True)\n", + "\n", + "api = wandb.Api()\n", + "project_name = \"thesis\"\n", + "experiment_05_sweep_id = \"bwom9hlj\"" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "07d208ff", + "metadata": {}, + "outputs": [], + "source": [ + "sweep = api.sweep(f\"{project_name}/{experiment_05_sweep_id}\")\n", + "runs = sweep.runs\n", + "\n", + "run_data = extract_run_data(runs)" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "49f63a0e", + "metadata": {}, + "outputs": [], + "source": [ + "run_data = run_data[run_data['state'] == 'finished']" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "36150ba5", + "metadata": {}, + "outputs": [], + "source": [ + "run_histories = [\n", + " get_run_history(\n", + " run,\n", + " metrics=[\n", + " \"training/train_loss\",\n", + " \"training/val_loss\",\n", + " \"training/train_accuracy\",\n", + " \"training/val_accuracy\",\n", + " ],\n", + " )\n", + " for run in runs\n", + " if run.state == \"finished\"\n", + "]" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "1174193f", + "metadata": {}, + "outputs": [], + "source": [ + "run_histories[0]" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "514ad825", + "metadata": {}, + "outputs": [], + "source": [ + "run_data.columns" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "e6d4cfba", + "metadata": {}, + "outputs": [], + "source": [ + "example_run = runs[5]\n", + "example_run_history = get_run_history(example_run, metrics=[\"training/train_loss\", \"training/val_loss\", \"training/train_accuracy\", \"training/val_accuracy\"])" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "056da49e", + "metadata": {}, + "outputs": [], + "source": [ + "example_run.summary" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "93fa5d02", + "metadata": {}, + "outputs": [], + "source": [ + "example_run_history" + ] + }, + { + "cell_type": "markdown", + "id": "3416f7e9", + "metadata": {}, + "source": [ + "## Experiment 05 plots" + ] + }, + { + "cell_type": "markdown", + "id": "567f7c13", + "metadata": {}, + "source": [ + "### Learning rate" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "59c345bf", + "metadata": {}, + "outputs": [], + "source": [ + "def make_lr_plot(run_data: pd.DataFrame, output_dir: Path, filename: str, metric_key: str, metric_name: str, show: bool = False):\n", + " plot_data = run_data.copy()\n", + "\n", + " plt.figure(figsize=(9, 6))\n", + "\n", + " head_only_data = plot_data[plot_data['config_only_head'] == True]\n", + " full_model_data = plot_data[plot_data['config_only_head'] == False]\n", + "\n", + " plt.scatter(\n", + " head_only_data['config_learning_rate'], \n", + " head_only_data[metric_key],\n", + " alpha=0.7,\n", + " s=60,\n", + " c='steelblue',\n", + " edgecolors='black',\n", + " linewidth=0.5,\n", + " label='only_head=True'\n", + " )\n", + "\n", + " plt.scatter(\n", + " full_model_data['config_learning_rate'], \n", + " full_model_data[metric_key],\n", + " alpha=0.7,\n", + " s=60,\n", + " c='red',\n", + " edgecolors='black',\n", + " linewidth=0.5,\n", + " label='only_head=False'\n", + " )\n", + "\n", + " plt.xscale('log')\n", + " plt.xlabel('Współczynnik uczenia (skala log)', fontsize=12)\n", + " plt.ylabel(metric_name, fontsize=12)\n", + " plt.grid(True, alpha=0.3)\n", + " plt.legend(fontsize=11, loc='lower left')\n", + " plt.gca().xaxis.set_major_formatter(plt.FuncFormatter(lambda x, p: f'{x:.1e}'))\n", + " plt.tight_layout()\n", + "\n", + " plt.savefig(output_dir / filename, dpi=300, bbox_inches='tight')\n", + " if show:\n", + " plt.show()\n", + " plt.close()\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "ef7ed257", + "metadata": {}, + "outputs": [], + "source": [ + "experiment_05_plots_dir = output_dir / \"experiment-05\"\n", + "experiment_05_plots_dir.mkdir(parents=True, exist_ok=True)\n", + "\n", + "plot_args = [\n", + " (\"lr-vs-lfw.png\", \"summary_benchmark/lfw_accuracy\", \"Dokładność LFW\"),\n", + " (\"lr-vs-rof-m.png\", \"summary_benchmark/rof_masked_accuracy\", \"Dokładność ROF-m\"),\n", + " (\"lr-vs-rof-s.png\", \"summary_benchmark/rof_sunglasses_accuracy\", \"Dokładność ROF-s\"),\n", + " (\"lr-vs-train-accuracy.png\", \"max_training/train_accuracy\", \"Dokładność na zbiorze treningowym\"),\n", + " (\"lr-vs-val-accuracy.png\", \"max_training/val_accuracy\", \"Dokładność na zbiorze walidacyjnym\"),\n", + "]\n", + "\n", + "for filename, metric_key, metric_name in plot_args:\n", + " make_lr_plot(\n", + " run_data=run_data,\n", + " output_dir=experiment_05_plots_dir,\n", + " filename=filename,\n", + " metric_key=metric_key,\n", + " metric_name=metric_name,\n", + " show=True\n", + " )" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "3cf607f6", + "metadata": {}, + "outputs": [], + "source": [ + "def make_margin_plot(run_data: pd.DataFrame, output_dir: Path, filename: str, metric_key: str, metric_name: str, show: bool = False):\n", + " plot_data = run_data.copy()\n", + "\n", + " plt.figure(figsize=(9, 6))\n", + "\n", + " plt.scatter(\n", + " plot_data['config_margin'], \n", + " plot_data[metric_key],\n", + " alpha=0.7,\n", + " s=60,\n", + " c='green',\n", + " edgecolors='black',\n", + " linewidth=0.5,\n", + " )\n", + "\n", + " plt.xlabel('Margines', fontsize=12)\n", + " plt.ylabel(metric_name, fontsize=12)\n", + " plt.grid(True, alpha=0.3)\n", + " plt.tight_layout()\n", + " \n", + " plt.savefig(output_dir / filename, dpi=300, bbox_inches='tight')\n", + " if show:\n", + " plt.show()\n", + " plt.close()\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "5a1beaab", + "metadata": {}, + "outputs": [], + "source": [ + "plot_args = [\n", + " (\"margin-vs-lfw.png\", \"summary_benchmark/lfw_accuracy\", \"Dokładność LFW\"),\n", + " (\"margin-vs-rof-m.png\", \"summary_benchmark/rof_masked_accuracy\", \"Dokładność ROF-m\"),\n", + " (\"margin-vs-rof-s.png\", \"summary_benchmark/rof_sunglasses_accuracy\", \"Dokładność ROF-s\"),\n", + " (\"margin-vs-train-accuracy.png\", \"max_training/train_accuracy\", \"Dokładność na zbiorze treningowym\"),\n", + " (\"margin-vs-val-accuracy.png\", \"max_training/val_accuracy\", \"Dokładność na zbiorze walidacyjnym\"),\n", + "]\n", + "\n", + "for filename, metric_key, metric_name in plot_args:\n", + " make_margin_plot(\n", + " run_data=run_data,\n", + " output_dir=experiment_05_plots_dir,\n", + " filename=filename,\n", + " metric_key=metric_key,\n", + " metric_name=metric_name\n", + " )" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "a0129437", + "metadata": {}, + "outputs": [], + "source": [ + "def make_augmentation_plot(run_data: pd.DataFrame, output_dir: Path, filename: str, metric_key: str, metric_name: str, show: bool = False):\n", + " plot_data = run_data.copy()\n", + " plot_data['config_augmentation'] = plot_data['config_augmentation'].fillna('None')\n", + "\n", + " plt.figure(figsize=(9, 6))\n", + "\n", + " augmentations = ['None', 'AddRandomRectangleAverageColor']\n", + "\n", + " data_by_augmentation = [plot_data[plot_data['config_augmentation'] == aug][metric_key].values for aug in augmentations]\n", + "\n", + " box_plot = plt.boxplot(\n", + " data_by_augmentation,\n", + " tick_labels=[\"\", \"\"],\n", + " patch_artist=True,\n", + " showmeans=True,\n", + " vert=False,\n", + " )\n", + " \n", + " box_plot[\"boxes\"][0].set_facecolor(\"lightblue\")\n", + " box_plot[\"boxes\"][0].set_alpha(0.7)\n", + " box_plot[\"boxes\"][1].set_facecolor(\"lightcoral\")\n", + " box_plot[\"boxes\"][1].set_alpha(0.7)\n", + "\n", + " # Get the left edge of the entire plot area with some margin\n", + " x_min = plt.xlim()[0]\n", + " x_range = plt.xlim()[1] - plt.xlim()[0]\n", + " x_margin = x_min + (0.05 * x_range) # 5% margin from left edge\n", + "\n", + " for i, label in enumerate(augmentations):\n", + " # Get the top edge of the box for positioning above\n", + " box_top = box_plot[\"boxes\"][i].get_path().vertices[:, 1].max()\n", + "\n", + " plt.text(\n", + " x_margin,\n", + " box_top + 0.15,\n", + " label,\n", + " horizontalalignment=\"left\", # Left aligned to plot area\n", + " verticalalignment=\"bottom\",\n", + " fontsize=10,\n", + " )\n", + "\n", + " plt.xlabel(metric_name, fontsize=12)\n", + " plt.grid(True, alpha=0.3, axis='x')\n", + " plt.tight_layout()\n", + "\n", + " plt.savefig(output_dir / filename, dpi=300, bbox_inches='tight')\n", + " if show:\n", + " plt.show()\n", + " plt.close()\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "4d634d5c", + "metadata": {}, + "outputs": [], + "source": [ + "plot_args = [\n", + " (\"augmentation-vs-lfw.png\", \"summary_benchmark/lfw_accuracy\", \"Dokładność LFW\"),\n", + " (\"augmentation-vs-rof-m.png\", \"summary_benchmark/rof_masked_accuracy\", \"Dokładność ROF-m\"),\n", + " (\"augmentation-vs-rof-s.png\", \"summary_benchmark/rof_sunglasses_accuracy\", \"Dokładność ROF-s\"),\n", + " (\"augmentation-vs-train-accuracy.png\", \"max_training/train_accuracy\", \"Dokładność na zbiorze treningowym\"),\n", + " (\"augmentation-vs-val-accuracy.png\", \"max_training/val_accuracy\", \"Dokładność na zbiorze walidacyjnym\"),\n", + "]\n", + "\n", + "for filename, metric_key, metric_name in plot_args:\n", + " make_augmentation_plot(\n", + " run_data=run_data,\n", + " output_dir=experiment_05_plots_dir,\n", + " filename=filename,\n", + " metric_key=metric_key,\n", + " metric_name=metric_name,\n", + " show=True,\n", + " )" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "ec637f9e", + "metadata": {}, + "outputs": [], + "source": [ + "def make_only_head_plot(\n", + " run_data: pd.DataFrame,\n", + " output_dir: Path,\n", + " filename: str,\n", + " metric_key: str,\n", + " metric_name: str,\n", + " show: bool = False,\n", + "):\n", + " plot_data = run_data.copy()\n", + " plot_data[\"config_only_head\"] = plot_data[\"config_only_head\"].map(\n", + " {True: \"True\", False: \"False\"}\n", + " )\n", + "\n", + " plt.figure(figsize=(9, 6))\n", + "\n", + " data_by_only_head = [\n", + " plot_data[plot_data[\"config_only_head\"] == val][metric_key].values\n", + " for val in [\"True\", \"False\"]\n", + " ]\n", + "\n", + " labels = [\"only_head=True\", \"only_head=False\"]\n", + "\n", + " box_plot = plt.boxplot(\n", + " data_by_only_head,\n", + " tick_labels=[\"\", \"\"],\n", + " patch_artist=True,\n", + " showmeans=True,\n", + " vert=False,\n", + " )\n", + "\n", + " box_plot[\"boxes\"][0].set_facecolor(\"lightblue\")\n", + " box_plot[\"boxes\"][0].set_alpha(0.7)\n", + " box_plot[\"boxes\"][1].set_facecolor(\"lightcoral\")\n", + " box_plot[\"boxes\"][1].set_alpha(0.7)\n", + "\n", + " # Get the left edge of the entire plot area with some margin\n", + " x_min = plt.xlim()[0]\n", + " x_range = plt.xlim()[1] - plt.xlim()[0]\n", + " x_margin = x_min + (0.05 * x_range) # 5% margin from left edge\n", + "\n", + " for i, label in enumerate(labels):\n", + " # Get the top edge of the box for positioning above\n", + " box_top = box_plot[\"boxes\"][i].get_path().vertices[:, 1].max()\n", + "\n", + " plt.text(\n", + " x_margin,\n", + " box_top + 0.15,\n", + " label,\n", + " horizontalalignment=\"left\", # Left aligned to plot area\n", + " verticalalignment=\"bottom\",\n", + " fontsize=10,\n", + " )\n", + "\n", + " plt.xlabel(metric_name, fontsize=12)\n", + " plt.grid(True, alpha=0.3, axis=\"x\")\n", + " plt.tick_params(axis=\"both\", which=\"major\", labelsize=9)\n", + " plt.tight_layout()\n", + "\n", + " plt.savefig(output_dir / filename, dpi=300, bbox_inches=\"tight\")\n", + " if show:\n", + " plt.show()\n", + " plt.close()" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "a49c220e", + "metadata": {}, + "outputs": [], + "source": [ + "plot_args = [\n", + " (\"only-head-vs-lfw.png\", \"summary_benchmark/lfw_accuracy\", \"Dokładność LFW\"),\n", + " (\"only-head-vs-rof-m.png\", \"summary_benchmark/rof_masked_accuracy\", \"Dokładność ROF-m\"),\n", + " (\"only-head-vs-rof-s.png\", \"summary_benchmark/rof_sunglasses_accuracy\", \"Dokładność ROF-s\"),\n", + " (\"only-head-vs-train-accuracy.png\", \"max_training/train_accuracy\", \"Dokładność na zbiorze treningowym\"),\n", + " (\"only-head-vs-val-accuracy.png\", \"max_training/val_accuracy\", \"Dokładność na zbiorze walidacyjnym\"),\n", + " (\"only-head-vs-runtime.png\", \"runtime\", \"Czas obliczeń [s]\"),\n", + "]\n", + "\n", + "for filename, metric_key, metric_name in plot_args:\n", + " make_only_head_plot(\n", + " run_data=run_data,\n", + " output_dir=experiment_05_plots_dir,\n", + " filename=filename,\n", + " metric_key=metric_key,\n", + " metric_name=metric_name,\n", + " show=True,\n", + " )" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "e6919153", + "metadata": {}, + "outputs": [], + "source": [ + "def make_mini_batch_size_plot(\n", + " run_data: pd.DataFrame,\n", + " output_dir: Path,\n", + " filename: str,\n", + " metric_key: str,\n", + " metric_name: str,\n", + " show: bool = False,\n", + "):\n", + " plot_data = run_data.copy()\n", + "\n", + " plt.figure(figsize=(9, 6))\n", + "\n", + " mini_batch_sizes = sorted(plot_data[\"config_batch_size\"].unique())\n", + "\n", + " data_by_mini_batch_size = [\n", + " plot_data[plot_data[\"config_batch_size\"] == size][metric_key].values\n", + " for size in mini_batch_sizes\n", + " ]\n", + "\n", + " box_plot = plt.boxplot(\n", + " data_by_mini_batch_size,\n", + " tick_labels=[str(size) for size in mini_batch_sizes],\n", + " patch_artist=True,\n", + " showmeans=True,\n", + " vert=False,\n", + " )\n", + "\n", + " colors = plt.cm.viridis(np.linspace(0, 1, len(mini_batch_sizes)))\n", + " for i, box in enumerate(box_plot[\"boxes\"]):\n", + " box.set_facecolor(colors[i])\n", + " box.set_alpha(0.7)\n", + "\n", + "\n", + " plt.xlabel(metric_name, fontsize=12)\n", + " plt.ylabel(\"Rozmiar mini-pakietu\", fontsize=12)\n", + " plt.grid(True, alpha=0.3, axis=\"x\")\n", + " plt.tight_layout()\n", + "\n", + " plt.savefig(output_dir / filename, dpi=300, bbox_inches=\"tight\")\n", + " if show:\n", + " plt.show()\n", + " plt.close()" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "011085e3", + "metadata": {}, + "outputs": [], + "source": [ + "plot_args = [\n", + " (\"mini-batch-size-vs-lfw.png\", \"summary_benchmark/lfw_accuracy\", \"Dokładność LFW\"),\n", + " (\"mini-batch-size-vs-rof-m.png\", \"summary_benchmark/rof_masked_accuracy\", \"Dokładność ROF-m\"),\n", + " (\"mini-batch-size-vs-rof-s.png\", \"summary_benchmark/rof_sunglasses_accuracy\", \"Dokładność ROF-s\"),\n", + " (\"mini-batch-size-vs-train-accuracy.png\", \"max_training/train_accuracy\", \"Dokładność na zbiorze treningowym\"),\n", + " (\"mini-batch-size-vs-val-accuracy.png\", \"max_training/val_accuracy\", \"Dokładność na zbiorze walidacyjnym\"),\n", + "]\n", + "\n", + "for filename, metric_key, metric_name in plot_args:\n", + " make_mini_batch_size_plot(\n", + " run_data=run_data,\n", + " output_dir=experiment_05_plots_dir,\n", + " filename=filename,\n", + " metric_key=metric_key,\n", + " metric_name=metric_name,\n", + " show=True,\n", + " )" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "63740ae0", + "metadata": {}, + "outputs": [], + "source": [ + "def make_loss_or_acc_plot(\n", + " run_histories: list[pd.DataFrame],\n", + " output_dir: Path,\n", + " filename: str,\n", + " metric_key: str,\n", + " metric_name: str,\n", + " show: bool = False,\n", + "):\n", + " plt.figure(figsize=(10, 6))\n", + "\n", + " for history in run_histories:\n", + " plt.plot(\n", + " history[\"_step\"],\n", + " history[metric_key],\n", + " alpha=0.3,\n", + " linewidth=1,\n", + " )\n", + "\n", + " plt.xlabel(\"Epoka\", fontsize=12)\n", + " plt.ylabel(metric_name, fontsize=12)\n", + " plt.grid(True, alpha=0.3)\n", + " plt.tight_layout()\n", + "\n", + " plt.savefig(output_dir / filename, dpi=300, bbox_inches=\"tight\")\n", + " if show:\n", + " plt.show()\n", + " plt.close()" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "43c1d78a", + "metadata": {}, + "outputs": [], + "source": [ + "plot_args = [\n", + " (\"train-loss-over-epochs.png\", \"training/train_loss\", \"Strata na zbiorze treningowym\"),\n", + " (\"val-loss-over-epochs.png\", \"training/val_loss\", \"Strata na zbiorze walidacyjnym\"),\n", + " (\"train-accuracy-over-epochs.png\", \"training/train_accuracy\", \"Dokładność na zbiorze treningowym\"),\n", + " (\"val-accuracy-over-epochs.png\", \"training/val_accuracy\", \"Dokładność na zbiorze walidacyjnym\"),\n", + "]\n", + "\n", + "for filename, metric_key, metric_name in plot_args:\n", + " make_loss_or_acc_plot(\n", + " run_histories=run_histories,\n", + " output_dir=experiment_05_plots_dir,\n", + " filename=filename,\n", + " metric_key=metric_key,\n", + " metric_name=metric_name,\n", + " show=True,\n", + " )" + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": "thesis-3.12", + "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.3" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} From 32818843dab8e64ca79821d3e7a5aa87aeda79ed Mon Sep 17 00:00:00 2001 From: mgarbowski Date: Thu, 2 Oct 2025 00:13:00 +0200 Subject: [PATCH 02/14] add todo --- notebooks/12-plotting_from_wandb_api.ipynb | 12 +++--------- 1 file changed, 3 insertions(+), 9 deletions(-) diff --git a/notebooks/12-plotting_from_wandb_api.ipynb b/notebooks/12-plotting_from_wandb_api.ipynb index b1be1d4..4b83b8d 100644 --- a/notebooks/12-plotting_from_wandb_api.ipynb +++ b/notebooks/12-plotting_from_wandb_api.ipynb @@ -218,15 +218,9 @@ "id": "3416f7e9", "metadata": {}, "source": [ - "## Experiment 05 plots" - ] - }, - { - "cell_type": "markdown", - "id": "567f7c13", - "metadata": {}, - "source": [ - "### Learning rate" + "## Experiment 05 \n", + "\n", + "TODO: parameter importance data has to be scraped manually to recreate plots" ] }, { From 93afaa704ec07c7dd75a5fd13a95386493f56d67 Mon Sep 17 00:00:00 2001 From: mgarbowski Date: Thu, 2 Oct 2025 17:47:44 +0200 Subject: [PATCH 03/14] migrate plotting code to python script --- src/scripts/plots_for_publication.py | 670 +++++++++++++++++++++++++++ 1 file changed, 670 insertions(+) create mode 100644 src/scripts/plots_for_publication.py diff --git a/src/scripts/plots_for_publication.py b/src/scripts/plots_for_publication.py new file mode 100644 index 0000000..ef9c4b7 --- /dev/null +++ b/src/scripts/plots_for_publication.py @@ -0,0 +1,670 @@ +from dataclasses import dataclass +from typing import Any +import wandb +from pathlib import Path +from pprint import pprint +import pandas as pd +import json +import matplotlib.pyplot as plt +import numpy as np + + +@dataclass +class ScriptConfiguration: + output_dir: str = "results/plots" + project_name: str = "thesis" + cache_dir: str = "results/cache" + + +@dataclass +class SweepData: + sweep_id: str + sweep_runs_data: pd.DataFrame + run_histories: list[pd.DataFrame] + + +class WandbClient: + def __init__(self, project_name: str, cache_dir: Path): + self.api = wandb.Api() + self.project_name = project_name + self.cache_dir = cache_dir + self.cache_dir.mkdir(parents=True, exist_ok=True) + + def get_sweep_data(self, sweep_id: str) -> SweepData: + if self._get_cache_filepath(sweep_id).exists(): + return self._load_sweep_data_from_json(sweep_id) + else: + sweep_data = self.fetch_sweep_data_from_api(sweep_id) + self._save_sweep_data_to_json(sweep_data) + return sweep_data + + def fetch_sweep_data_from_api(self, sweep_id: str): + sweep = self._get_sweep_object(sweep_id) + sweep_runs_data = self._extract_sweep_run_data(sweep) + sweep_runs_data = sweep_runs_data[sweep_runs_data["state"] == "finished"] + run_histories = [self._extract_run_history(run) for run in sweep.runs] + return SweepData( + sweep_id=sweep_id, + sweep_runs_data=sweep_runs_data, + run_histories=run_histories, + ) + + def _get_cache_filepath(self, sweep_id: str) -> Path: + return self.cache_dir / f"sweep_{sweep_id}_data.json" + + def _save_sweep_data_to_json(self, sweep_data: SweepData): + with open(self._get_cache_filepath(sweep_data.sweep_id), "w") as f: + json.dump( + { + "sweep_runs_data": sweep_data.sweep_runs_data.to_dict( + orient="records" + ), + "run_histories": [ + rh.to_dict(orient="records") for rh in sweep_data.run_histories + ], + }, + f, + indent=4, + default=str, + ) + + def _load_sweep_data_from_json(self, sweep_id: str) -> SweepData: + with open(self._get_cache_filepath(sweep_id), "r") as f: + data = json.load(f) + sweep_runs_data = pd.DataFrame(data["sweep_runs_data"]) + run_histories = [pd.DataFrame(rh) for rh in data["run_histories"]] + return SweepData( + sweep_id=sweep_id, + sweep_runs_data=sweep_runs_data, + run_histories=run_histories, + ) + + def _extract_sweep_run_data(self, sweep) -> pd.DataFrame: + data = [] + + for run in sweep.runs: + row = { + "run_id": run.id, + "run_name": run.name, + "state": run.state, # finished, failed, running, etc. + "created_at": run.created_at, + "runtime": run.summary["_runtime"], + } + + # Add config parameters (hyperparameters) + for key, value in run.config.items(): + row[f"config_{key}"] = value + + # Add summary metrics (final values) + for key, value in run.summary.items(): + if not key.startswith("_"): # Skip internal wandb fields + row[f"summary_{key}"] = value + + # Add history metrics (you can get specific values) + history = run.history() + if not history.empty: + # Get final values + for col in history.columns: + if not col.startswith("_"): + row[f"final_{col}"] = ( + history[col].iloc[-1] if len(history) > 0 else None + ) + + # Get maximum values for accuracy metrics + accuracy_cols = [ + col for col in history.columns if "accuracy" in col.lower() + ] + for col in accuracy_cols: + row[f"max_{col}"] = history[col].max() + + data.append(row) + + return pd.DataFrame(data) + + def _extract_run_history(self, run) -> pd.DataFrame: + history = run.history() + metrics = [ + "training/train_loss", + "training/val_loss", + "training/train_accuracy", + "training/val_accuracy", + ] + + # Filter to specific metrics (plus _step and _timestamp) + available_metrics = [m for m in metrics if m in history.columns] + cols_to_keep = ["_step", "_timestamp"] + available_metrics + history = history[cols_to_keep] + + history["run_id"] = run.id + history["run_name"] = run.name + + return history + + def _get_sweep_object(self, sweep_id: str): + return self.api.sweep(f"{self.project_name}/{sweep_id}") + + +def make_lr_plot( + run_data: pd.DataFrame, + output_dir: Path, + filename: str, + metric_key: str, + metric_name: str, + show: bool = False, +): + plot_data = run_data.copy() + + plt.figure(figsize=(9, 6)) + + head_only_data = plot_data[plot_data["config_only_head"] == True] + full_model_data = plot_data[plot_data["config_only_head"] == False] + + plt.scatter( + head_only_data["config_learning_rate"], + head_only_data[metric_key], + alpha=0.7, + s=60, + c="steelblue", + edgecolors="black", + linewidth=0.5, + label="only_head=True", + ) + + plt.scatter( + full_model_data["config_learning_rate"], + full_model_data[metric_key], + alpha=0.7, + s=60, + c="red", + edgecolors="black", + linewidth=0.5, + label="only_head=False", + ) + + plt.xscale("log") + plt.xlabel("Współczynnik uczenia (skala log)", fontsize=12) + plt.ylabel(metric_name, fontsize=12) + plt.grid(True, alpha=0.3) + plt.legend(fontsize=11, loc="lower left") + plt.gca().xaxis.set_major_formatter(plt.FuncFormatter(lambda x, p: f"{x:.1e}")) + plt.tight_layout() + + plt.savefig(output_dir / filename, dpi=300, bbox_inches="tight") + if show: + plt.show() + plt.close() + + +def make_margin_plot( + run_data: pd.DataFrame, + output_dir: Path, + filename: str, + metric_key: str, + metric_name: str, + show: bool = False, +): + plot_data = run_data.copy() + + plt.figure(figsize=(9, 6)) + + plt.scatter( + plot_data["config_margin"], + plot_data[metric_key], + alpha=0.7, + s=60, + c="green", + edgecolors="black", + linewidth=0.5, + ) + + plt.xlabel("Margines", fontsize=12) + plt.ylabel(metric_name, fontsize=12) + plt.grid(True, alpha=0.3) + plt.tight_layout() + + plt.savefig(output_dir / filename, dpi=300, bbox_inches="tight") + if show: + plt.show() + plt.close() + + +def make_augmentation_plot( + run_data: pd.DataFrame, + output_dir: Path, + filename: str, + metric_key: str, + metric_name: str, + show: bool = False, +): + plot_data = run_data.copy() + plot_data["config_augmentation"] = plot_data["config_augmentation"].fillna("None") + + plt.figure(figsize=(9, 6)) + + augmentations = ["None", "AddRandomRectangleAverageColor"] + + data_by_augmentation = [ + plot_data[plot_data["config_augmentation"] == aug][metric_key].values + for aug in augmentations + ] + + box_plot = plt.boxplot( + data_by_augmentation, + tick_labels=["", ""], + patch_artist=True, + showmeans=True, + vert=False, + ) + + box_plot["boxes"][0].set_facecolor("lightblue") + box_plot["boxes"][0].set_alpha(0.7) + box_plot["boxes"][1].set_facecolor("lightcoral") + box_plot["boxes"][1].set_alpha(0.7) + + # Get the left edge of the entire plot area with some margin + x_min = plt.xlim()[0] + x_range = plt.xlim()[1] - plt.xlim()[0] + x_margin = x_min + (0.05 * x_range) # 5% margin from left edge + + for i, label in enumerate(augmentations): + # Get the top edge of the box for positioning above + box_top = box_plot["boxes"][i].get_path().vertices[:, 1].max() + + plt.text( + x_margin, + box_top + 0.15, + label, + horizontalalignment="left", # Left aligned to plot area + verticalalignment="bottom", + fontsize=10, + ) + + plt.xlabel(metric_name, fontsize=12) + plt.grid(True, alpha=0.3, axis="x") + plt.tight_layout() + + plt.savefig(output_dir / filename, dpi=300, bbox_inches="tight") + if show: + plt.show() + plt.close() + + +def make_only_head_plot( + run_data: pd.DataFrame, + output_dir: Path, + filename: str, + metric_key: str, + metric_name: str, + show: bool = False, +): + plot_data = run_data.copy() + plot_data["config_only_head"] = plot_data["config_only_head"].map( + {True: "True", False: "False"} + ) + + plt.figure(figsize=(9, 6)) + + data_by_only_head = [ + plot_data[plot_data["config_only_head"] == val][metric_key].values + for val in ["True", "False"] + ] + + labels = ["only_head=True", "only_head=False"] + + box_plot = plt.boxplot( + data_by_only_head, + tick_labels=["", ""], + patch_artist=True, + showmeans=True, + vert=False, + ) + + box_plot["boxes"][0].set_facecolor("lightblue") + box_plot["boxes"][0].set_alpha(0.7) + box_plot["boxes"][1].set_facecolor("lightcoral") + box_plot["boxes"][1].set_alpha(0.7) + + # Get the left edge of the entire plot area with some margin + x_min = plt.xlim()[0] + x_range = plt.xlim()[1] - plt.xlim()[0] + x_margin = x_min + (0.05 * x_range) # 5% margin from left edge + + for i, label in enumerate(labels): + # Get the top edge of the box for positioning above + box_top = box_plot["boxes"][i].get_path().vertices[:, 1].max() + + plt.text( + x_margin, + box_top + 0.15, + label, + horizontalalignment="left", # Left aligned to plot area + verticalalignment="bottom", + fontsize=10, + ) + + plt.xlabel(metric_name, fontsize=12) + plt.grid(True, alpha=0.3, axis="x") + plt.tick_params(axis="both", which="major", labelsize=9) + plt.tight_layout() + + plt.savefig(output_dir / filename, dpi=300, bbox_inches="tight") + if show: + plt.show() + plt.close() + + +def make_mini_batch_size_plot( + run_data: pd.DataFrame, + output_dir: Path, + filename: str, + metric_key: str, + metric_name: str, + show: bool = False, +): + plot_data = run_data.copy() + + plt.figure(figsize=(9, 6)) + + mini_batch_sizes = sorted(plot_data["config_batch_size"].unique()) + + data_by_mini_batch_size = [ + plot_data[plot_data["config_batch_size"] == size][metric_key].values + for size in mini_batch_sizes + ] + + box_plot = plt.boxplot( + data_by_mini_batch_size, + tick_labels=[str(size) for size in mini_batch_sizes], + patch_artist=True, + showmeans=True, + vert=False, + ) + + colors = plt.cm.viridis(np.linspace(0, 1, len(mini_batch_sizes))) + for i, box in enumerate(box_plot["boxes"]): + box.set_facecolor(colors[i]) + box.set_alpha(0.7) + + plt.xlabel(metric_name, fontsize=12) + plt.ylabel("Rozmiar mini-pakietu", fontsize=12) + plt.grid(True, alpha=0.3, axis="x") + plt.tight_layout() + + plt.savefig(output_dir / filename, dpi=300, bbox_inches="tight") + if show: + plt.show() + plt.close() + + +def make_loss_or_acc_plot( + run_histories: list[pd.DataFrame], + output_dir: Path, + filename: str, + metric_key: str, + metric_name: str, + show: bool = False, +): + plt.figure(figsize=(10, 6)) + + for history in run_histories: + plt.plot( + history["_step"], + history[metric_key], + alpha=0.3, + linewidth=1, + ) + + plt.xlabel("Epoka", fontsize=12) + plt.ylabel(metric_name, fontsize=12) + plt.grid(True, alpha=0.3) + plt.tight_layout() + + plt.savefig(output_dir / filename, dpi=300, bbox_inches="tight") + if show: + plt.show() + plt.close() + + +def make_experiment_05_plots(sweep_data: SweepData, output_dir: Path): + plot_args = [ + ("lr-vs-lfw.png", "summary_benchmark/lfw_accuracy", "Dokładność LFW"), + ( + "lr-vs-rof-m.png", + "summary_benchmark/rof_masked_accuracy", + "Dokładność ROF-m", + ), + ( + "lr-vs-rof-s.png", + "summary_benchmark/rof_sunglasses_accuracy", + "Dokładność ROF-s", + ), + ( + "lr-vs-train-accuracy.png", + "max_training/train_accuracy", + "Dokładność na zbiorze treningowym", + ), + ( + "lr-vs-val-accuracy.png", + "max_training/val_accuracy", + "Dokładność na zbiorze walidacyjnym", + ), + ] + + dir = output_dir / "lr" + dir.mkdir(parents=True, exist_ok=True) + for filename, metric_key, metric_name in plot_args: + make_lr_plot( + run_data=sweep_data.sweep_runs_data, + output_dir=dir, + filename=filename, + metric_key=metric_key, + metric_name=metric_name, + show=False, + ) + + plot_args = [ + ("margin-vs-lfw.png", "summary_benchmark/lfw_accuracy", "Dokładność LFW"), + ( + "margin-vs-rof-m.png", + "summary_benchmark/rof_masked_accuracy", + "Dokładność ROF-m", + ), + ( + "margin-vs-rof-s.png", + "summary_benchmark/rof_sunglasses_accuracy", + "Dokładność ROF-s", + ), + ( + "margin-vs-train-accuracy.png", + "max_training/train_accuracy", + "Dokładność na zbiorze treningowym", + ), + ( + "margin-vs-val-accuracy.png", + "max_training/val_accuracy", + "Dokładność na zbiorze walidacyjnym", + ), + ] + + dir = output_dir / "margin" + dir.mkdir(parents=True, exist_ok=True) + for filename, metric_key, metric_name in plot_args: + make_margin_plot( + run_data=sweep_data.sweep_runs_data, + output_dir=dir, + filename=filename, + metric_key=metric_key, + metric_name=metric_name, + ) + + plot_args = [ + ("augmentation-vs-lfw.png", "summary_benchmark/lfw_accuracy", "Dokładność LFW"), + ( + "augmentation-vs-rof-m.png", + "summary_benchmark/rof_masked_accuracy", + "Dokładność ROF-m", + ), + ( + "augmentation-vs-rof-s.png", + "summary_benchmark/rof_sunglasses_accuracy", + "Dokładność ROF-s", + ), + ( + "augmentation-vs-train-accuracy.png", + "max_training/train_accuracy", + "Dokładność na zbiorze treningowym", + ), + ( + "augmentation-vs-val-accuracy.png", + "max_training/val_accuracy", + "Dokładność na zbiorze walidacyjnym", + ), + ] + + dir = output_dir / "augmentation" + dir.mkdir(parents=True, exist_ok=True) + for filename, metric_key, metric_name in plot_args: + make_augmentation_plot( + run_data=sweep_data.sweep_runs_data, + output_dir=dir, + filename=filename, + metric_key=metric_key, + metric_name=metric_name, + show=False, + ) + + plot_args = [ + ("only-head-vs-lfw.png", "summary_benchmark/lfw_accuracy", "Dokładność LFW"), + ( + "only-head-vs-rof-m.png", + "summary_benchmark/rof_masked_accuracy", + "Dokładność ROF-m", + ), + ( + "only-head-vs-rof-s.png", + "summary_benchmark/rof_sunglasses_accuracy", + "Dokładność ROF-s", + ), + ( + "only-head-vs-train-accuracy.png", + "max_training/train_accuracy", + "Dokładność na zbiorze treningowym", + ), + ( + "only-head-vs-val-accuracy.png", + "max_training/val_accuracy", + "Dokładność na zbiorze walidacyjnym", + ), + ("only-head-vs-runtime.png", "runtime", "Czas obliczeń [s]"), + ] + + dir = output_dir / "only_head" + dir.mkdir(parents=True, exist_ok=True) + for filename, metric_key, metric_name in plot_args: + make_only_head_plot( + run_data=sweep_data.sweep_runs_data, + output_dir=dir, + filename=filename, + metric_key=metric_key, + metric_name=metric_name, + show=False, + ) + + plot_args = [ + ( + "mini-batch-size-vs-lfw.png", + "summary_benchmark/lfw_accuracy", + "Dokładność LFW", + ), + ( + "mini-batch-size-vs-rof-m.png", + "summary_benchmark/rof_masked_accuracy", + "Dokładność ROF-m", + ), + ( + "mini-batch-size-vs-rof-s.png", + "summary_benchmark/rof_sunglasses_accuracy", + "Dokładność ROF-s", + ), + ( + "mini-batch-size-vs-train-accuracy.png", + "max_training/train_accuracy", + "Dokładność na zbiorze treningowym", + ), + ( + "mini-batch-size-vs-val-accuracy.png", + "max_training/val_accuracy", + "Dokładność na zbiorze walidacyjnym", + ), + ] + + dir = output_dir / "mini-batch-size" + dir.mkdir(parents=True, exist_ok=True) + for filename, metric_key, metric_name in plot_args: + make_mini_batch_size_plot( + run_data=sweep_data.sweep_runs_data, + output_dir=dir, + filename=filename, + metric_key=metric_key, + metric_name=metric_name, + show=False, + ) + + plot_args = [ + ( + "train-loss-over-epochs.png", + "training/train_loss", + "Strata na zbiorze treningowym", + ), + ( + "val-loss-over-epochs.png", + "training/val_loss", + "Strata na zbiorze walidacyjnym", + ), + ( + "train-accuracy-over-epochs.png", + "training/train_accuracy", + "Dokładność na zbiorze treningowym", + ), + ( + "val-accuracy-over-epochs.png", + "training/val_accuracy", + "Dokładność na zbiorze walidacyjnym", + ), + ] + + dir = output_dir / "training" + dir.mkdir(parents=True, exist_ok=True) + for filename, metric_key, metric_name in plot_args: + make_loss_or_acc_plot( + run_histories=sweep_data.run_histories, + output_dir=dir, + filename=filename, + metric_key=metric_key, + metric_name=metric_name, + show=False, + ) + + + + +def main(): + config = ScriptConfiguration() + + output_dir = Path(config.output_dir) + output_dir.mkdir(parents=True, exist_ok=True) + + wandb_client = WandbClient(config.project_name, Path(config.cache_dir)) + + experiment_05_sweep_id = "bwom9hlj" + + sweep_data = wandb_client.get_sweep_data(experiment_05_sweep_id) + experiment_05_output_dir = output_dir / "experiment_05" + experiment_05_output_dir.mkdir(parents=True, exist_ok=True) + make_experiment_05_plots(sweep_data, experiment_05_output_dir) + + experiment_06_sweep_id = "1z7f1o6h" + + +if __name__ == "__main__": + main() From def6be6e4c5c046bec722b2c3faa13aeb12eed73 Mon Sep 17 00:00:00 2001 From: mgarbowski Date: Thu, 2 Oct 2025 17:55:26 +0200 Subject: [PATCH 04/14] plots for experiments 5 and 6 --- src/scripts/plots_for_publication.py | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/src/scripts/plots_for_publication.py b/src/scripts/plots_for_publication.py index ef9c4b7..c8111a5 100644 --- a/src/scripts/plots_for_publication.py +++ b/src/scripts/plots_for_publication.py @@ -663,8 +663,12 @@ def main(): experiment_05_output_dir.mkdir(parents=True, exist_ok=True) make_experiment_05_plots(sweep_data, experiment_05_output_dir) - experiment_06_sweep_id = "1z7f1o6h" - + experiment_06_sweep_id = "d5qbai3t" + sweep_data = wandb_client.get_sweep_data(experiment_06_sweep_id) + experiment_06_output_dir = output_dir / "experiment_06" + experiment_06_output_dir.mkdir(parents=True, exist_ok=True) + # Same as 05 + make_experiment_05_plots(sweep_data, experiment_06_output_dir) if __name__ == "__main__": main() From 2d31dec03617f48a924c4810bb82c03a0dca8c36 Mon Sep 17 00:00:00 2001 From: mgarbowski Date: Thu, 2 Oct 2025 17:56:53 +0200 Subject: [PATCH 05/14] rename function --- src/scripts/plots_for_publication.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/src/scripts/plots_for_publication.py b/src/scripts/plots_for_publication.py index c8111a5..68b41cd 100644 --- a/src/scripts/plots_for_publication.py +++ b/src/scripts/plots_for_publication.py @@ -425,7 +425,8 @@ def make_loss_or_acc_plot( plt.close() -def make_experiment_05_plots(sweep_data: SweepData, output_dir: Path): +def maker_hyperparam_search_experiment_plots(sweep_data: SweepData, output_dir: Path): + """Generate plots for experiments 05 and 06 (hyperparameter search).""" plot_args = [ ("lr-vs-lfw.png", "summary_benchmark/lfw_accuracy", "Dokładność LFW"), ( @@ -661,14 +662,14 @@ def main(): sweep_data = wandb_client.get_sweep_data(experiment_05_sweep_id) experiment_05_output_dir = output_dir / "experiment_05" experiment_05_output_dir.mkdir(parents=True, exist_ok=True) - make_experiment_05_plots(sweep_data, experiment_05_output_dir) + maker_hyperparam_search_experiment_plots(sweep_data, experiment_05_output_dir) experiment_06_sweep_id = "d5qbai3t" sweep_data = wandb_client.get_sweep_data(experiment_06_sweep_id) experiment_06_output_dir = output_dir / "experiment_06" experiment_06_output_dir.mkdir(parents=True, exist_ok=True) # Same as 05 - make_experiment_05_plots(sweep_data, experiment_06_output_dir) + maker_hyperparam_search_experiment_plots(sweep_data, experiment_06_output_dir) if __name__ == "__main__": main() From d09358dafca795a6f6a7c24558c712287edd954e Mon Sep 17 00:00:00 2001 From: mgarbowski Date: Thu, 2 Oct 2025 18:23:08 +0200 Subject: [PATCH 06/14] aug comparison plot --- src/scripts/plots_for_publication.py | 131 ++++++++++++++++++++++++--- 1 file changed, 119 insertions(+), 12 deletions(-) diff --git a/src/scripts/plots_for_publication.py b/src/scripts/plots_for_publication.py index 68b41cd..f1985f9 100644 --- a/src/scripts/plots_for_publication.py +++ b/src/scripts/plots_for_publication.py @@ -1,5 +1,5 @@ from dataclasses import dataclass -from typing import Any +from typing import Any, Literal import wandb from pathlib import Path from pprint import pprint @@ -424,8 +424,53 @@ def make_loss_or_acc_plot( plt.show() plt.close() +def make_aug_comparison_plot( + run_data: pd.DataFrame, + output_dir: Path, + filename: str, + metric_key: str, + metric_name: str, + aggregate_best: str = "max", # "max" or "min" + xmin: float = 0.0, + xmax: float = 1.0, + show: bool = False, +): + plot_data = run_data.copy() + plot_data["config_augmentation"] = plot_data["config_augmentation"].fillna("None") + + plt.figure(figsize=(9, 6)) + + aug_names = plot_data["config_augmentation"].unique() + best_result_by_aug = plot_data.groupby("config_augmentation").agg( + {metric_key: aggregate_best} + ).reset_index().sort_values(by=metric_key, ascending=True) # Changed to True for horizontal bars + + # Create horizontal bar plot + y_pos = np.arange(len(best_result_by_aug)) + plt.barh( + y_pos, + best_result_by_aug[metric_key], + color=plt.cm.viridis(np.linspace(0, 1, len(aug_names))), + alpha=0.7, + edgecolor="black", + ) + + # Set y-axis labels to augmentation names + plt.yticks(y_pos, best_result_by_aug["config_augmentation"]) + + plt.xlim(xmin, xmax) + plt.xlabel(metric_name, fontsize=12) + plt.grid(True, alpha=0.3, axis="x") # Changed to x-axis grid + plt.tight_layout() + + plt.savefig(output_dir / filename, dpi=300, bbox_inches="tight") + if show: + plt.show() + plt.close() + def maker_hyperparam_search_experiment_plots(sweep_data: SweepData, output_dir: Path): + # TODO parameter importance plots """Generate plots for experiments 05 and 06 (hyperparameter search).""" plot_args = [ ("lr-vs-lfw.png", "summary_benchmark/lfw_accuracy", "Dokładność LFW"), @@ -647,8 +692,63 @@ def maker_hyperparam_search_experiment_plots(sweep_data: SweepData, output_dir: ) +def make_aug_eval_experiment_plots(sweep_data: SweepData, output_dir: Path): + """Generate plots for experiments 07 and 08 (augmentation evaluation).""" + + plot_args = [ + ( + "aug-comparison-lfw.png", + "summary_benchmark/lfw_accuracy", + "Dokładność LFW", + "max", + 0.90, + 1.0, + ), + ( + "aug-comparison-rof-m.png", + "summary_benchmark/rof_masked_accuracy", + "Dokładność ROF-m", + "max", 0.80, + 0.9, + ), + ( + "aug-comparison-rof-s.png", + "summary_benchmark/rof_sunglasses_accuracy", + "Dokładność ROF-s", + "max", + 0.80, 0.9 + ), + ( + "aug-comparison-eer.png", + "summary_rococo-evaluation/eer", + "EER", + "min", 0.0, 0.5 + ), + ( + "aug-comparison-frr-at-far-zero.png", + "summary_rococo-evaluation/frr_at_far_zero", + "FRR @ FAR=0", + "min", 0.0, 1.0 + ), + ] + dir = output_dir / "augmentation-comparison" + dir.mkdir(parents=True, exist_ok=True) + for filename, metric_key, metric_name, aggregate_best, xmin, xmax in plot_args: + make_aug_comparison_plot( + run_data=sweep_data.sweep_runs_data, + output_dir=dir, + filename=filename, + metric_key=metric_key, + metric_name=metric_name, + aggregate_best=aggregate_best, + xmin=xmin, + xmax=xmax, + show=False, + ) + + def main(): config = ScriptConfiguration() @@ -657,19 +757,26 @@ def main(): wandb_client = WandbClient(config.project_name, Path(config.cache_dir)) - experiment_05_sweep_id = "bwom9hlj" + # experiment_05_sweep_id = "bwom9hlj" + + # sweep_data = wandb_client.get_sweep_data(experiment_05_sweep_id) + # experiment_05_output_dir = output_dir / "experiment_05" + # experiment_05_output_dir.mkdir(parents=True, exist_ok=True) + # maker_hyperparam_search_experiment_plots(sweep_data, experiment_05_output_dir) + + # experiment_06_sweep_id = "d5qbai3t" + # sweep_data = wandb_client.get_sweep_data(experiment_06_sweep_id) + # experiment_06_output_dir = output_dir / "experiment_06" + # experiment_06_output_dir.mkdir(parents=True, exist_ok=True) + # # Same as 05 + # maker_hyperparam_search_experiment_plots(sweep_data, experiment_06_output_dir) - sweep_data = wandb_client.get_sweep_data(experiment_05_sweep_id) - experiment_05_output_dir = output_dir / "experiment_05" - experiment_05_output_dir.mkdir(parents=True, exist_ok=True) - maker_hyperparam_search_experiment_plots(sweep_data, experiment_05_output_dir) + experiment_07_sweep_id = "afr8ycgh" + sweep_data = wandb_client.get_sweep_data(experiment_07_sweep_id) + experiment_07_output_dir = output_dir / "experiment_07" + experiment_07_output_dir.mkdir(parents=True, exist_ok=True) + make_aug_eval_experiment_plots(sweep_data, experiment_07_output_dir) - experiment_06_sweep_id = "d5qbai3t" - sweep_data = wandb_client.get_sweep_data(experiment_06_sweep_id) - experiment_06_output_dir = output_dir / "experiment_06" - experiment_06_output_dir.mkdir(parents=True, exist_ok=True) - # Same as 05 - maker_hyperparam_search_experiment_plots(sweep_data, experiment_06_output_dir) if __name__ == "__main__": main() From 0901c848e691e9364dca0061c43908c0afd47a2f Mon Sep 17 00:00:00 2001 From: mgarbowski Date: Thu, 2 Oct 2025 18:28:41 +0200 Subject: [PATCH 07/14] add baseline to aug comparison plot --- src/scripts/plots_for_publication.py | 21 +++++++++++++++++---- 1 file changed, 17 insertions(+), 4 deletions(-) diff --git a/src/scripts/plots_for_publication.py b/src/scripts/plots_for_publication.py index f1985f9..2daf1fe 100644 --- a/src/scripts/plots_for_publication.py +++ b/src/scripts/plots_for_publication.py @@ -430,6 +430,7 @@ def make_aug_comparison_plot( filename: str, metric_key: str, metric_name: str, + base: float, aggregate_best: str = "max", # "max" or "min" xmin: float = 0.0, xmax: float = 1.0, @@ -461,6 +462,12 @@ def make_aug_comparison_plot( plt.xlim(xmin, xmax) plt.xlabel(metric_name, fontsize=12) plt.grid(True, alpha=0.3, axis="x") # Changed to x-axis grid + + + # Add vertical line for base score + plt.axvline(x=base, color="red", linestyle="--", label="Baseline") + + plt.tight_layout() plt.savefig(output_dir / filename, dpi=300, bbox_inches="tight") @@ -703,6 +710,7 @@ def make_aug_eval_experiment_plots(sweep_data: SweepData, output_dir: Path): "max", 0.90, 1.0, + 0.971, ), ( "aug-comparison-rof-m.png", @@ -710,32 +718,36 @@ def make_aug_eval_experiment_plots(sweep_data: SweepData, output_dir: Path): "Dokładność ROF-m", "max", 0.80, 0.9, + 0.859 ), ( "aug-comparison-rof-s.png", "summary_benchmark/rof_sunglasses_accuracy", "Dokładność ROF-s", "max", - 0.80, 0.9 + 0.80, 0.9, + 0.872 ), ( "aug-comparison-eer.png", "summary_rococo-evaluation/eer", "EER", - "min", 0.0, 0.5 + "min", 0.0, 0.5, + 0.102 ), ( "aug-comparison-frr-at-far-zero.png", "summary_rococo-evaluation/frr_at_far_zero", "FRR @ FAR=0", - "min", 0.0, 1.0 + "min", 0.0, 1.0, + 0.763 ), ] dir = output_dir / "augmentation-comparison" dir.mkdir(parents=True, exist_ok=True) - for filename, metric_key, metric_name, aggregate_best, xmin, xmax in plot_args: + for filename, metric_key, metric_name, aggregate_best, xmin, xmax, base in plot_args: make_aug_comparison_plot( run_data=sweep_data.sweep_runs_data, output_dir=dir, @@ -745,6 +757,7 @@ def make_aug_eval_experiment_plots(sweep_data: SweepData, output_dir: Path): aggregate_best=aggregate_best, xmin=xmin, xmax=xmax, + base=base, show=False, ) From 9c76b95d32f2d0d492f71daaed260641e4e8bba2 Mon Sep 17 00:00:00 2001 From: mgarbowski Date: Thu, 2 Oct 2025 18:31:06 +0200 Subject: [PATCH 08/14] experiment 7 plots --- src/scripts/plots_for_publication.py | 34 ++++++++++++++++++++++++++++ 1 file changed, 34 insertions(+) diff --git a/src/scripts/plots_for_publication.py b/src/scripts/plots_for_publication.py index 2daf1fe..d9240dd 100644 --- a/src/scripts/plots_for_publication.py +++ b/src/scripts/plots_for_publication.py @@ -761,6 +761,40 @@ def make_aug_eval_experiment_plots(sweep_data: SweepData, output_dir: Path): show=False, ) + plot_args = [ + ( + "train-loss-over-epochs.png", + "training/train_loss", + "Strata na zbiorze treningowym", + ), + ( + "val-loss-over-epochs.png", + "training/val_loss", + "Strata na zbiorze walidacyjnym", + ), + ( + "train-accuracy-over-epochs.png", + "training/train_accuracy", + "Dokładność na zbiorze treningowym", + ), + ( + "val-accuracy-over-epochs.png", + "training/val_accuracy", + "Dokładność na zbiorze walidacyjnym", + ), + ] + + dir = output_dir / "training" + dir.mkdir(parents=True, exist_ok=True) + for filename, metric_key, metric_name in plot_args: + make_loss_or_acc_plot( + run_histories=sweep_data.run_histories, + output_dir=dir, + filename=filename, + metric_key=metric_key, + metric_name=metric_name, + show=False, + ) def main(): config = ScriptConfiguration() From 2625fa97448a0d6ae881c36d4ee63802abb5c124 Mon Sep 17 00:00:00 2001 From: mgarbowski Date: Thu, 2 Oct 2025 18:35:18 +0200 Subject: [PATCH 09/14] plots for experiment 08 --- src/scripts/plots_for_publication.py | 65 +++++++++++++++++++--------- 1 file changed, 44 insertions(+), 21 deletions(-) diff --git a/src/scripts/plots_for_publication.py b/src/scripts/plots_for_publication.py index d9240dd..f4fe259 100644 --- a/src/scripts/plots_for_publication.py +++ b/src/scripts/plots_for_publication.py @@ -424,6 +424,7 @@ def make_loss_or_acc_plot( plt.show() plt.close() + def make_aug_comparison_plot( run_data: pd.DataFrame, output_dir: Path, @@ -431,7 +432,7 @@ def make_aug_comparison_plot( metric_key: str, metric_name: str, base: float, - aggregate_best: str = "max", # "max" or "min" + aggregate_best: str = "max", # "max" or "min" xmin: float = 0.0, xmax: float = 1.0, show: bool = False, @@ -442,10 +443,13 @@ def make_aug_comparison_plot( plt.figure(figsize=(9, 6)) aug_names = plot_data["config_augmentation"].unique() - best_result_by_aug = plot_data.groupby("config_augmentation").agg( - {metric_key: aggregate_best} - ).reset_index().sort_values(by=metric_key, ascending=True) # Changed to True for horizontal bars - + best_result_by_aug = ( + plot_data.groupby("config_augmentation") + .agg({metric_key: aggregate_best}) + .reset_index() + .sort_values(by=metric_key, ascending=True) + ) # Changed to True for horizontal bars + # Create horizontal bar plot y_pos = np.arange(len(best_result_by_aug)) plt.barh( @@ -463,11 +467,9 @@ def make_aug_comparison_plot( plt.xlabel(metric_name, fontsize=12) plt.grid(True, alpha=0.3, axis="x") # Changed to x-axis grid - # Add vertical line for base score plt.axvline(x=base, color="red", linestyle="--", label="Baseline") - plt.tight_layout() plt.savefig(output_dir / filename, dpi=300, bbox_inches="tight") @@ -701,7 +703,7 @@ def maker_hyperparam_search_experiment_plots(sweep_data: SweepData, output_dir: def make_aug_eval_experiment_plots(sweep_data: SweepData, output_dir: Path): """Generate plots for experiments 07 and 08 (augmentation evaluation).""" - + plot_args = [ ( "aug-comparison-lfw.png", @@ -716,38 +718,51 @@ def make_aug_eval_experiment_plots(sweep_data: SweepData, output_dir: Path): "aug-comparison-rof-m.png", "summary_benchmark/rof_masked_accuracy", "Dokładność ROF-m", - "max", 0.80, - 0.9, - 0.859 + "max", + 0.70, + 0.90, + 0.859, ), ( "aug-comparison-rof-s.png", "summary_benchmark/rof_sunglasses_accuracy", "Dokładność ROF-s", - "max", - 0.80, 0.9, - 0.872 + "max", + 0.70, + 0.90, + 0.872, ), ( "aug-comparison-eer.png", "summary_rococo-evaluation/eer", "EER", - "min", 0.0, 0.5, - 0.102 + "min", + 0.0, + 0.5, + 0.102, ), ( "aug-comparison-frr-at-far-zero.png", "summary_rococo-evaluation/frr_at_far_zero", "FRR @ FAR=0", - "min", 0.0, 1.0, - 0.763 + "min", + 0.0, + 1.0, + 0.763, ), ] - dir = output_dir / "augmentation-comparison" dir.mkdir(parents=True, exist_ok=True) - for filename, metric_key, metric_name, aggregate_best, xmin, xmax, base in plot_args: + for ( + filename, + metric_key, + metric_name, + aggregate_best, + xmin, + xmax, + base, + ) in plot_args: make_aug_comparison_plot( run_data=sweep_data.sweep_runs_data, output_dir=dir, @@ -760,7 +775,7 @@ def make_aug_eval_experiment_plots(sweep_data: SweepData, output_dir: Path): base=base, show=False, ) - + plot_args = [ ( "train-loss-over-epochs.png", @@ -796,6 +811,7 @@ def make_aug_eval_experiment_plots(sweep_data: SweepData, output_dir: Path): show=False, ) + def main(): config = ScriptConfiguration() @@ -824,6 +840,13 @@ def main(): experiment_07_output_dir.mkdir(parents=True, exist_ok=True) make_aug_eval_experiment_plots(sweep_data, experiment_07_output_dir) + experiment_08_sweep_id = "lgn9xwtm" + sweep_data = wandb_client.get_sweep_data(experiment_08_sweep_id) + experiment_08_output_dir = output_dir / "experiment_08" + experiment_08_output_dir.mkdir(parents=True, exist_ok=True) + # Same as 07 + make_aug_eval_experiment_plots(sweep_data, experiment_08_output_dir) + if __name__ == "__main__": main() From 93fdb7657e082da1d2760fc3d78adde95f66c019 Mon Sep 17 00:00:00 2001 From: mgarbowski Date: Thu, 2 Oct 2025 18:50:23 +0200 Subject: [PATCH 10/14] complete plotting script --- src/scripts/plots_for_publication.py | 77 ++++++++++++++++++++++------ 1 file changed, 61 insertions(+), 16 deletions(-) diff --git a/src/scripts/plots_for_publication.py b/src/scripts/plots_for_publication.py index f4fe259..b597ca7 100644 --- a/src/scripts/plots_for_publication.py +++ b/src/scripts/plots_for_publication.py @@ -1,8 +1,6 @@ from dataclasses import dataclass -from typing import Any, Literal import wandb from pathlib import Path -from pprint import pprint import pandas as pd import json import matplotlib.pyplot as plt @@ -32,8 +30,10 @@ def __init__(self, project_name: str, cache_dir: Path): def get_sweep_data(self, sweep_id: str) -> SweepData: if self._get_cache_filepath(sweep_id).exists(): + print(f"Loading sweep data from cache: {sweep_id}") return self._load_sweep_data_from_json(sweep_id) else: + print(f"Fetching sweep data from API: {sweep_id}") sweep_data = self.fetch_sweep_data_from_api(sweep_id) self._save_sweep_data_to_json(sweep_data) return sweep_data @@ -811,6 +811,41 @@ def make_aug_eval_experiment_plots(sweep_data: SweepData, output_dir: Path): show=False, ) +def make_experiment_09_plots(sweep_data: SweepData, output_dir: Path): + plot_args = [ + ( + "train-loss-over-epochs.png", + "training/train_loss", + "Strata na zbiorze treningowym", + ), + ( + "val-loss-over-epochs.png", + "training/val_loss", + "Strata na zbiorze walidacyjnym", + ), + ( + "train-accuracy-over-epochs.png", + "training/train_accuracy", + "Dokładność na zbiorze treningowym", + ), + ( + "val-accuracy-over-epochs.png", + "training/val_accuracy", + "Dokładność na zbiorze walidacyjnym", + ), + ] + + dir = output_dir / "training" + dir.mkdir(parents=True, exist_ok=True) + for filename, metric_key, metric_name in plot_args: + make_loss_or_acc_plot( + run_histories=sweep_data.run_histories, + output_dir=dir, + filename=filename, + metric_key=metric_key, + metric_name=metric_name, + show=False, + ) def main(): config = ScriptConfiguration() @@ -820,26 +855,29 @@ def main(): wandb_client = WandbClient(config.project_name, Path(config.cache_dir)) - # experiment_05_sweep_id = "bwom9hlj" - - # sweep_data = wandb_client.get_sweep_data(experiment_05_sweep_id) - # experiment_05_output_dir = output_dir / "experiment_05" - # experiment_05_output_dir.mkdir(parents=True, exist_ok=True) - # maker_hyperparam_search_experiment_plots(sweep_data, experiment_05_output_dir) - - # experiment_06_sweep_id = "d5qbai3t" - # sweep_data = wandb_client.get_sweep_data(experiment_06_sweep_id) - # experiment_06_output_dir = output_dir / "experiment_06" - # experiment_06_output_dir.mkdir(parents=True, exist_ok=True) - # # Same as 05 - # maker_hyperparam_search_experiment_plots(sweep_data, experiment_06_output_dir) - + print("Generating plots for experiment 05...") + experiment_05_sweep_id = "bwom9hlj" + sweep_data = wandb_client.get_sweep_data(experiment_05_sweep_id) + experiment_05_output_dir = output_dir / "experiment_05" + experiment_05_output_dir.mkdir(parents=True, exist_ok=True) + maker_hyperparam_search_experiment_plots(sweep_data, experiment_05_output_dir) + + print("Generating plots for experiment 06...") + experiment_06_sweep_id = "d5qbai3t" + sweep_data = wandb_client.get_sweep_data(experiment_06_sweep_id) + experiment_06_output_dir = output_dir / "experiment_06" + experiment_06_output_dir.mkdir(parents=True, exist_ok=True) + # Same as 05 + maker_hyperparam_search_experiment_plots(sweep_data, experiment_06_output_dir) + + print("Generating plots for experiment 07...") experiment_07_sweep_id = "afr8ycgh" sweep_data = wandb_client.get_sweep_data(experiment_07_sweep_id) experiment_07_output_dir = output_dir / "experiment_07" experiment_07_output_dir.mkdir(parents=True, exist_ok=True) make_aug_eval_experiment_plots(sweep_data, experiment_07_output_dir) + print("Generating plots for experiment 08...") experiment_08_sweep_id = "lgn9xwtm" sweep_data = wandb_client.get_sweep_data(experiment_08_sweep_id) experiment_08_output_dir = output_dir / "experiment_08" @@ -847,6 +885,13 @@ def main(): # Same as 07 make_aug_eval_experiment_plots(sweep_data, experiment_08_output_dir) + print("Generating plots for experiment 09...") + experiment_09_id = "72yrgb2d" + sweep_data = wandb_client.get_sweep_data(experiment_09_id) + experiment_09_output_dir = output_dir / "experiment_09" + experiment_09_output_dir.mkdir(parents=True, exist_ok=True) + make_experiment_09_plots(sweep_data, experiment_09_output_dir) + if __name__ == "__main__": main() From ebbb38f58e22df53dd5b24ee2539b133997f8a10 Mon Sep 17 00:00:00 2001 From: mgarbowski Date: Thu, 2 Oct 2025 19:17:27 +0200 Subject: [PATCH 11/14] implement parameter importance plotting --- src/scripts/plots_for_publication.py | 262 +++++++++++++++++++++++---- 1 file changed, 230 insertions(+), 32 deletions(-) diff --git a/src/scripts/plots_for_publication.py b/src/scripts/plots_for_publication.py index b597ca7..10532d1 100644 --- a/src/scripts/plots_for_publication.py +++ b/src/scripts/plots_for_publication.py @@ -14,6 +14,125 @@ class ScriptConfiguration: cache_dir: str = "results/cache" +experiment_05_parameter_importance = { + "lfw": { + "importance": { + "learning_rate": 0.676, + "only_head": 0.250, + "batch_size": 0.043, + "margin": 0.020, + }, + "correlation": { + "learning_rate": -0.464, + "only_head": 0.406, + "batch_size": 0.170, + "margin": -0.116, + }, + }, + "rof-m": { + "importance": { + "learning_rate": 0.695, + "only_head": 0.229, + "batch_size": 0.047, + "margin": 0.017, + }, + "correlation": { + "learning_rate": -0.491, + "only_head": 0.402, + "batch_size": 0.187, + "margin": -0.014, + }, + }, + "rof-s": { + "importance": { + "learning_rate": 0.662, + "only_head": 0.228, + "batch_size": 0.077, + "margin": 0.029, + }, + "correlation": { + "learning_rate": -0.508, + "only_head": 0.427, + "batch_size": 0.195, + "margin": -0.037, + }, + }, + "val_accuracy": { + "importance": { + "learning_rate": 0.484, + "only_head": 0.127, + "batch_size": 0.130, + "margin": 0.107, + }, + "correlation": { + "learning_rate": -0.533, + "only_head": 0.305, + "batch_size": -0.312, + "margin": -0.319, + }, + }, +} + +experiment_06_parameter_importance = { + "lfw": { + "importance": { + "learning_rate": 0.733, + "only_head": 0.206, + "batch_size": 0.012, + "margin": 0.031, + }, + "correlation": { + "learning_rate": -0.490, + "only_head": 0.424, + "batch_size": 0.062, + "margin": -0.141, + }, + }, + "rof-m": { + "importance": { + "learning_rate": 0.755, + "only_head": 0.175, + "batch_size": 0.015, + "margin": 0.033, + }, + "correlation": { + "learning_rate": -0.481, + "only_head": 0.412, + "batch_size": 0.062, + "margin": -0.153, + }, + }, + "rof-s": { + "importance": { + "learning_rate": 0.676, + "only_head": 0.249, + "batch_size": 0.011, + "margin": 0.040, + }, + "correlation": { + "learning_rate": -0.419, + "only_head": 0.499, + "batch_size": -0.009, + "margin": -0.104, + }, + }, + "val_accuracy": { + "importance": { + "learning_rate": 0.370, + "only_head": 0.435, + "batch_size": 0.019, + "margin": 0.042, + }, + "correlation": { + "learning_rate": -0.290, + "only_head": -0.706, + "batch_size": 0.103, + "margin": -0.160, + }, + }, +} + + @dataclass class SweepData: sweep_id: str @@ -477,8 +596,76 @@ def make_aug_comparison_plot( plt.show() plt.close() +def make_parameter_analysis_plot( + importance_dict: dict, + correlation_dict: dict, + output_dir: Path, + filename: str, + title: str, + show: bool = False +): + parameters = list(importance_dict.keys()) + importances = [importance_dict[param] for param in parameters] + correlations = [correlation_dict[param] for param in parameters] + + # Get absolute values for correlation bar heights + abs_correlations = [abs(corr) for corr in correlations] + + # Create colors based on original correlation sign (positive=green, negative=red) + corr_colors = ['green' if corr > 0 else 'red' for corr in correlations] + + fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(9, 6)) + + # Left subplot: Parameter Importance + bars1 = ax1.bar( + parameters, importances, color="skyblue", edgecolor="black", alpha=0.7 + ) + ax1.set_ylabel("Ważność", fontsize=12) + ax1.set_ylim(0, max(importances) * 1.1) + ax1.grid(axis="y", alpha=0.3) + + # Add value labels on top of bars for importance + for bar in bars1: + height = bar.get_height() + ax1.text( + bar.get_x() + bar.get_width() / 2, + height, + f"{height:.3f}", + ha="center", + va="bottom", + fontsize=10, + ) + + # Right subplot: Parameter Correlation (absolute values with color coding) + bars2 = ax2.bar( + parameters, abs_correlations, color=corr_colors, edgecolor="black", alpha=0.7 + ) + ax2.set_ylabel("Korelacja", fontsize=12) + ax2.set_ylim(0, max(abs_correlations) * 1.1) + ax2.grid(axis="y", alpha=0.3) + + # Add value labels on top of bars for correlation (showing original values) + for bar, original_corr in zip(bars2, correlations): + height = bar.get_height() + ax2.text( + bar.get_x() + bar.get_width() / 2, + height, + f"{original_corr:.3f}", + ha="center", + va="bottom", + fontsize=10, + ) + + # Set overall title + fig.suptitle(title, fontsize=16) + + plt.tight_layout() + plt.savefig(output_dir / filename, dpi=300, bbox_inches="tight") + if show: + plt.show() + plt.close() -def maker_hyperparam_search_experiment_plots(sweep_data: SweepData, output_dir: Path): +def make_hyperparam_search_experiment_plots(sweep_data: SweepData, output_dir: Path): # TODO parameter importance plots """Generate plots for experiments 05 and 06 (hyperparameter search).""" plot_args = [ @@ -811,6 +998,7 @@ def make_aug_eval_experiment_plots(sweep_data: SweepData, output_dir: Path): show=False, ) + def make_experiment_09_plots(sweep_data: SweepData, output_dir: Path): plot_args = [ ( @@ -847,6 +1035,7 @@ def make_experiment_09_plots(sweep_data: SweepData, output_dir: Path): show=False, ) + def main(): config = ScriptConfiguration() @@ -860,37 +1049,46 @@ def main(): sweep_data = wandb_client.get_sweep_data(experiment_05_sweep_id) experiment_05_output_dir = output_dir / "experiment_05" experiment_05_output_dir.mkdir(parents=True, exist_ok=True) - maker_hyperparam_search_experiment_plots(sweep_data, experiment_05_output_dir) - - print("Generating plots for experiment 06...") - experiment_06_sweep_id = "d5qbai3t" - sweep_data = wandb_client.get_sweep_data(experiment_06_sweep_id) - experiment_06_output_dir = output_dir / "experiment_06" - experiment_06_output_dir.mkdir(parents=True, exist_ok=True) - # Same as 05 - maker_hyperparam_search_experiment_plots(sweep_data, experiment_06_output_dir) - - print("Generating plots for experiment 07...") - experiment_07_sweep_id = "afr8ycgh" - sweep_data = wandb_client.get_sweep_data(experiment_07_sweep_id) - experiment_07_output_dir = output_dir / "experiment_07" - experiment_07_output_dir.mkdir(parents=True, exist_ok=True) - make_aug_eval_experiment_plots(sweep_data, experiment_07_output_dir) - - print("Generating plots for experiment 08...") - experiment_08_sweep_id = "lgn9xwtm" - sweep_data = wandb_client.get_sweep_data(experiment_08_sweep_id) - experiment_08_output_dir = output_dir / "experiment_08" - experiment_08_output_dir.mkdir(parents=True, exist_ok=True) - # Same as 07 - make_aug_eval_experiment_plots(sweep_data, experiment_08_output_dir) - - print("Generating plots for experiment 09...") - experiment_09_id = "72yrgb2d" - sweep_data = wandb_client.get_sweep_data(experiment_09_id) - experiment_09_output_dir = output_dir / "experiment_09" - experiment_09_output_dir.mkdir(parents=True, exist_ok=True) - make_experiment_09_plots(sweep_data, experiment_09_output_dir) + make_hyperparam_search_experiment_plots(sweep_data, experiment_05_output_dir) + + + make_parameter_analysis_plot( + importance_dict=experiment_05_parameter_importance["lfw"]["importance"], + correlation_dict=experiment_05_parameter_importance["lfw"]["correlation"], + output_dir=experiment_05_output_dir, + filename="parameter-analysis-lfw.png", + title="LFW", + ) + + # print("Generating plots for experiment 06...") + # experiment_06_sweep_id = "d5qbai3t" + # sweep_data = wandb_client.get_sweep_data(experiment_06_sweep_id) + # experiment_06_output_dir = output_dir / "experiment_06" + # experiment_06_output_dir.mkdir(parents=True, exist_ok=True) + # # Same as 05 + # make_hyperparam_search_experiment_plots(sweep_data, experiment_06_output_dir) + + # print("Generating plots for experiment 07...") + # experiment_07_sweep_id = "afr8ycgh" + # sweep_data = wandb_client.get_sweep_data(experiment_07_sweep_id) + # experiment_07_output_dir = output_dir / "experiment_07" + # experiment_07_output_dir.mkdir(parents=True, exist_ok=True) + # make_aug_eval_experiment_plots(sweep_data, experiment_07_output_dir) + + # print("Generating plots for experiment 08...") + # experiment_08_sweep_id = "lgn9xwtm" + # sweep_data = wandb_client.get_sweep_data(experiment_08_sweep_id) + # experiment_08_output_dir = output_dir / "experiment_08" + # experiment_08_output_dir.mkdir(parents=True, exist_ok=True) + # # Same as 07 + # make_aug_eval_experiment_plots(sweep_data, experiment_08_output_dir) + + # print("Generating plots for experiment 09...") + # experiment_09_id = "72yrgb2d" + # sweep_data = wandb_client.get_sweep_data(experiment_09_id) + # experiment_09_output_dir = output_dir / "experiment_09" + # experiment_09_output_dir.mkdir(parents=True, exist_ok=True) + # make_experiment_09_plots(sweep_data, experiment_09_output_dir) if __name__ == "__main__": From 05ef46909f8e678297c6fc54706a9e97d251b436 Mon Sep 17 00:00:00 2001 From: mgarbowski Date: Thu, 2 Oct 2025 20:19:40 +0200 Subject: [PATCH 12/14] refactor plotting, create a library package for plotting --- pyproject.toml | 1 + src/plots/__init__.py | 4 + src/plots/individual_plots.py | 411 ++++++++++ src/plots/parameter_importance.py | 124 +++ src/plots/plot_groups.py | 355 ++++++++ src/plots/sweep_data.py | 23 + src/plots/wandb_client.py | 146 ++++ src/scripts/plots_for_publication.py | 1139 ++------------------------ 8 files changed, 1142 insertions(+), 1061 deletions(-) create mode 100644 src/plots/__init__.py create mode 100644 src/plots/individual_plots.py create mode 100644 src/plots/parameter_importance.py create mode 100644 src/plots/plot_groups.py create mode 100644 src/plots/sweep_data.py create mode 100644 src/plots/wandb_client.py diff --git a/pyproject.toml b/pyproject.toml index 5496121..bb8710b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -47,6 +47,7 @@ benchmark = "python -m src.scripts.benchmark" process_video = "python -m src.scripts.process_video" process_video_clean = "rm -rf data/rococo2/camera data/rococo2/faces" rococo_evaluation = "python -m src.scripts.rococo_evaluation" +plots = { shell = "python -m src.scripts.plots_for_publication", help = "Generate plots for the thesis" } # Testing test = "pytest test" diff --git a/src/plots/__init__.py b/src/plots/__init__.py new file mode 100644 index 0000000..c3b760d --- /dev/null +++ b/src/plots/__init__.py @@ -0,0 +1,4 @@ +"""Code for creating plots for the thesis. + +Visualizing experiment data downloaded from Weights & Biases API. +""" diff --git a/src/plots/individual_plots.py b/src/plots/individual_plots.py new file mode 100644 index 0000000..2b599fd --- /dev/null +++ b/src/plots/individual_plots.py @@ -0,0 +1,411 @@ +"""Reusable functions for creating plots from experiment data.""" + +from pathlib import Path + +import matplotlib.pyplot as plt +import numpy as np +import pandas as pd + + +def make_lr_plot( + run_data: pd.DataFrame, + output_dir: Path, + filename: str, + metric_key: str, + metric_name: str, + show: bool = False, +): + plot_data = run_data.copy() + + plt.figure(figsize=(9, 6)) + + head_only_data = plot_data[plot_data["config_only_head"] == True] + full_model_data = plot_data[plot_data["config_only_head"] == False] + + plt.scatter( + head_only_data["config_learning_rate"], + head_only_data[metric_key], + alpha=0.7, + s=60, + c="steelblue", + edgecolors="black", + linewidth=0.5, + label="only_head=True", + ) + + plt.scatter( + full_model_data["config_learning_rate"], + full_model_data[metric_key], + alpha=0.7, + s=60, + c="red", + edgecolors="black", + linewidth=0.5, + label="only_head=False", + ) + + plt.xscale("log") + plt.xlabel("Współczynnik uczenia (skala log)", fontsize=12) + plt.ylabel(metric_name, fontsize=12) + plt.grid(True, alpha=0.3) + plt.legend(fontsize=11, loc="lower left") + plt.gca().xaxis.set_major_formatter(plt.FuncFormatter(lambda x, p: f"{x:.1e}")) + plt.tight_layout() + + plt.savefig(output_dir / filename, dpi=300, bbox_inches="tight") + if show: + plt.show() + plt.close() + + +def make_margin_plot( + run_data: pd.DataFrame, + output_dir: Path, + filename: str, + metric_key: str, + metric_name: str, + show: bool = False, +): + plot_data = run_data.copy() + + plt.figure(figsize=(9, 6)) + + plt.scatter( + plot_data["config_margin"], + plot_data[metric_key], + alpha=0.7, + s=60, + c="green", + edgecolors="black", + linewidth=0.5, + ) + + plt.xlabel("Margines", fontsize=12) + plt.ylabel(metric_name, fontsize=12) + plt.grid(True, alpha=0.3) + plt.tight_layout() + + plt.savefig(output_dir / filename, dpi=300, bbox_inches="tight") + if show: + plt.show() + plt.close() + + +def make_augmentation_plot( + run_data: pd.DataFrame, + output_dir: Path, + filename: str, + metric_key: str, + metric_name: str, + show: bool = False, +): + plot_data = run_data.copy() + plot_data["config_augmentation"] = plot_data["config_augmentation"].fillna("None") + + plt.figure(figsize=(9, 6)) + + augmentations = ["None", "AddRandomRectangleAverageColor"] + + data_by_augmentation = [ + plot_data[plot_data["config_augmentation"] == aug][metric_key].values + for aug in augmentations + ] + + box_plot = plt.boxplot( + data_by_augmentation, + tick_labels=["", ""], + patch_artist=True, + showmeans=True, + vert=False, + ) + + box_plot["boxes"][0].set_facecolor("lightblue") + box_plot["boxes"][0].set_alpha(0.7) + box_plot["boxes"][1].set_facecolor("lightcoral") + box_plot["boxes"][1].set_alpha(0.7) + + # Get the left edge of the entire plot area with some margin + x_min = plt.xlim()[0] + x_range = plt.xlim()[1] - plt.xlim()[0] + x_margin = x_min + (0.05 * x_range) # 5% margin from left edge + + for i, label in enumerate(augmentations): + # Get the top edge of the box for positioning above + box_top = box_plot["boxes"][i].get_path().vertices[:, 1].max() + + plt.text( + x_margin, + box_top + 0.15, + label, + horizontalalignment="left", # Left aligned to plot area + verticalalignment="bottom", + fontsize=10, + ) + + plt.xlabel(metric_name, fontsize=12) + plt.grid(True, alpha=0.3, axis="x") + plt.tight_layout() + + plt.savefig(output_dir / filename, dpi=300, bbox_inches="tight") + if show: + plt.show() + plt.close() + + +def make_only_head_plot( + run_data: pd.DataFrame, + output_dir: Path, + filename: str, + metric_key: str, + metric_name: str, + show: bool = False, +): + plot_data = run_data.copy() + plot_data["config_only_head"] = plot_data["config_only_head"].map( + {True: "True", False: "False"} + ) + + plt.figure(figsize=(9, 6)) + + data_by_only_head = [ + plot_data[plot_data["config_only_head"] == val][metric_key].values + for val in ["True", "False"] + ] + + labels = ["only_head=True", "only_head=False"] + + box_plot = plt.boxplot( + data_by_only_head, + tick_labels=["", ""], + patch_artist=True, + showmeans=True, + vert=False, + ) + + box_plot["boxes"][0].set_facecolor("lightblue") + box_plot["boxes"][0].set_alpha(0.7) + box_plot["boxes"][1].set_facecolor("lightcoral") + box_plot["boxes"][1].set_alpha(0.7) + + # Get the left edge of the entire plot area with some margin + x_min = plt.xlim()[0] + x_range = plt.xlim()[1] - plt.xlim()[0] + x_margin = x_min + (0.05 * x_range) # 5% margin from left edge + + for i, label in enumerate(labels): + # Get the top edge of the box for positioning above + box_top = box_plot["boxes"][i].get_path().vertices[:, 1].max() + + plt.text( + x_margin, + box_top + 0.15, + label, + horizontalalignment="left", # Left aligned to plot area + verticalalignment="bottom", + fontsize=10, + ) + + plt.xlabel(metric_name, fontsize=12) + plt.grid(True, alpha=0.3, axis="x") + plt.tick_params(axis="both", which="major", labelsize=9) + plt.tight_layout() + + plt.savefig(output_dir / filename, dpi=300, bbox_inches="tight") + if show: + plt.show() + plt.close() + + +def make_mini_batch_size_plot( + run_data: pd.DataFrame, + output_dir: Path, + filename: str, + metric_key: str, + metric_name: str, + show: bool = False, +): + plot_data = run_data.copy() + + plt.figure(figsize=(9, 6)) + + mini_batch_sizes = sorted(plot_data["config_batch_size"].unique()) + + data_by_mini_batch_size = [ + plot_data[plot_data["config_batch_size"] == size][metric_key].values + for size in mini_batch_sizes + ] + + box_plot = plt.boxplot( + data_by_mini_batch_size, + tick_labels=[str(size) for size in mini_batch_sizes], + patch_artist=True, + showmeans=True, + vert=False, + ) + + colors = plt.cm.viridis(np.linspace(0, 1, len(mini_batch_sizes))) + for i, box in enumerate(box_plot["boxes"]): + box.set_facecolor(colors[i]) + box.set_alpha(0.7) + + plt.xlabel(metric_name, fontsize=12) + plt.ylabel("Rozmiar mini-pakietu", fontsize=12) + plt.grid(True, alpha=0.3, axis="x") + plt.tight_layout() + + plt.savefig(output_dir / filename, dpi=300, bbox_inches="tight") + if show: + plt.show() + plt.close() + + +def make_loss_or_acc_plot( + run_histories: list[pd.DataFrame], + output_dir: Path, + filename: str, + metric_key: str, + metric_name: str, + show: bool = False, +): + plt.figure(figsize=(9, 6)) + + for history in run_histories: + plt.plot( + history["_step"], + history[metric_key], + alpha=0.3, + linewidth=1, + ) + + plt.xlabel("Epoka", fontsize=12) + plt.ylabel(metric_name, fontsize=12) + plt.grid(True, alpha=0.3) + plt.tight_layout() + + plt.savefig(output_dir / filename, dpi=300, bbox_inches="tight") + if show: + plt.show() + plt.close() + + +def make_aug_comparison_plot( + run_data: pd.DataFrame, + output_dir: Path, + filename: str, + metric_key: str, + metric_name: str, + base: float, + aggregate_best: str = "max", # "max" or "min" + xmin: float = 0.0, + xmax: float = 1.0, + show: bool = False, +): + plot_data = run_data.copy() + plot_data["config_augmentation"] = plot_data["config_augmentation"].fillna("None") + + plt.figure(figsize=(9, 6)) + + aug_names = plot_data["config_augmentation"].unique() + best_result_by_aug = ( + plot_data.groupby("config_augmentation") + .agg({metric_key: aggregate_best}) + .reset_index() + .sort_values(by=metric_key, ascending=True) + ) # Changed to True for horizontal bars + + # Create horizontal bar plot + y_pos = np.arange(len(best_result_by_aug)) + plt.barh( + y_pos, + best_result_by_aug[metric_key], + color=plt.cm.viridis(np.linspace(0, 1, len(aug_names))), + alpha=0.7, + edgecolor="black", + ) + + # Set y-axis labels to augmentation names + plt.yticks(y_pos, best_result_by_aug["config_augmentation"]) + + plt.xlim(xmin, xmax) + plt.xlabel(metric_name, fontsize=12) + plt.grid(True, alpha=0.3, axis="x") # Changed to x-axis grid + + # Add vertical line for base score + plt.axvline(x=base, color="red", linestyle="--", label="Baseline") + + plt.tight_layout() + + plt.savefig(output_dir / filename, dpi=300, bbox_inches="tight") + if show: + plt.show() + plt.close() + + +def make_parameter_analysis_plot( + importance_dict: dict, + correlation_dict: dict, + output_dir: Path, + filename: str, + title: str, + show: bool = False, +): + parameters = list(importance_dict.keys()) + importances = [importance_dict[param] for param in parameters] + correlations = [correlation_dict[param] for param in parameters] + + # Get absolute values for correlation bar heights + abs_correlations = [abs(corr) for corr in correlations] + + # Create colors based on original correlation sign (positive=green, negative=red) + corr_colors = ["green" if corr > 0 else "red" for corr in correlations] + + fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(9, 6)) + + # Left subplot: Parameter Importance + bars1 = ax1.bar( + parameters, importances, color="skyblue", edgecolor="black", alpha=0.7 + ) + ax1.set_ylabel("Ważność", fontsize=12) + ax1.set_ylim(0, max(importances) * 1.1) + ax1.grid(axis="y", alpha=0.3) + + # Add value labels on top of bars for importance + for bar in bars1: + height = bar.get_height() + ax1.text( + bar.get_x() + bar.get_width() / 2, + height, + f"{height:.3f}", + ha="center", + va="bottom", + fontsize=10, + ) + + # Right subplot: Parameter Correlation (absolute values with color coding) + bars2 = ax2.bar( + parameters, abs_correlations, color=corr_colors, edgecolor="black", alpha=0.7 + ) + ax2.set_ylabel("Korelacja", fontsize=12) + ax2.set_ylim(0, max(abs_correlations) * 1.1) + ax2.grid(axis="y", alpha=0.3) + + # Add value labels on top of bars for correlation (showing original values) + for bar, original_corr in zip(bars2, correlations): + height = bar.get_height() + ax2.text( + bar.get_x() + bar.get_width() / 2, + height, + f"{original_corr:.3f}", + ha="center", + va="bottom", + fontsize=10, + ) + + # Set overall title + fig.suptitle(title, fontsize=16) + + plt.tight_layout() + plt.savefig(output_dir / filename, dpi=300, bbox_inches="tight") + if show: + plt.show() + plt.close() diff --git a/src/plots/parameter_importance.py b/src/plots/parameter_importance.py new file mode 100644 index 0000000..89f2530 --- /dev/null +++ b/src/plots/parameter_importance.py @@ -0,0 +1,124 @@ +"""Data on parameter importance and linear correlation with respect to metrics. + +This data is calculated by wandb but cannot be downloaded via the API. +I manually copied it from the web interface. +It is used for recreating the plots in the thesis. +""" + +experiment_05_parameter_importance = { + "lfw": { + "importance": { + "learning_rate": 0.676, + "only_head": 0.250, + "batch_size": 0.043, + "margin": 0.020, + }, + "correlation": { + "learning_rate": -0.464, + "only_head": 0.406, + "batch_size": 0.170, + "margin": -0.116, + }, + }, + "rof-m": { + "importance": { + "learning_rate": 0.695, + "only_head": 0.229, + "batch_size": 0.047, + "margin": 0.017, + }, + "correlation": { + "learning_rate": -0.491, + "only_head": 0.402, + "batch_size": 0.187, + "margin": -0.014, + }, + }, + "rof-s": { + "importance": { + "learning_rate": 0.662, + "only_head": 0.228, + "batch_size": 0.077, + "margin": 0.029, + }, + "correlation": { + "learning_rate": -0.508, + "only_head": 0.427, + "batch_size": 0.195, + "margin": -0.037, + }, + }, + "val_accuracy": { + "importance": { + "learning_rate": 0.484, + "only_head": 0.127, + "batch_size": 0.130, + "margin": 0.107, + }, + "correlation": { + "learning_rate": -0.533, + "only_head": 0.305, + "batch_size": -0.312, + "margin": -0.319, + }, + }, +} + +experiment_06_parameter_importance = { + "lfw": { + "importance": { + "learning_rate": 0.733, + "only_head": 0.206, + "batch_size": 0.012, + "margin": 0.031, + }, + "correlation": { + "learning_rate": -0.490, + "only_head": 0.424, + "batch_size": 0.062, + "margin": -0.141, + }, + }, + "rof-m": { + "importance": { + "learning_rate": 0.755, + "only_head": 0.175, + "batch_size": 0.015, + "margin": 0.033, + }, + "correlation": { + "learning_rate": -0.481, + "only_head": 0.412, + "batch_size": 0.062, + "margin": -0.153, + }, + }, + "rof-s": { + "importance": { + "learning_rate": 0.676, + "only_head": 0.249, + "batch_size": 0.011, + "margin": 0.040, + }, + "correlation": { + "learning_rate": -0.419, + "only_head": 0.499, + "batch_size": -0.009, + "margin": -0.104, + }, + }, + "val_accuracy": { + "importance": { + "learning_rate": 0.370, + "only_head": 0.435, + "batch_size": 0.019, + "margin": 0.042, + }, + "correlation": { + "learning_rate": -0.290, + "only_head": -0.706, + "batch_size": 0.103, + "margin": -0.160, + }, + }, +} diff --git a/src/plots/plot_groups.py b/src/plots/plot_groups.py new file mode 100644 index 0000000..27a84df --- /dev/null +++ b/src/plots/plot_groups.py @@ -0,0 +1,355 @@ +"""Functions for creating groups of plots for the thesis. + +There are groups of similar plots used together in the thesis. +(e.g. plots of learning rate vs metric X for different metrics). +""" + +from pathlib import Path +from typing import Any + +from src.plots.individual_plots import ( + make_aug_comparison_plot, + make_augmentation_plot, + make_loss_or_acc_plot, + make_lr_plot, + make_margin_plot, + make_mini_batch_size_plot, + make_only_head_plot, + make_parameter_analysis_plot, +) +from src.plots.sweep_data import SweepData + + +class PlotGroups: + def __init__(self, sweep_data: SweepData, output_dir: Path) -> None: + self.sweep_data = sweep_data + self.output_dir = output_dir + self.output_dir.mkdir(parents=True, exist_ok=True) + + def make_output_subdir(self, subdir_name: str) -> Path: + subdir = self.output_dir / subdir_name + subdir.mkdir(parents=True, exist_ok=True) + return subdir + + def learning_rate_v_metrics(self) -> None: + """Create plots showing the effect of learning rate on various metrics.""" + plot_args = [ + ("lr-vs-lfw.png", "summary_benchmark/lfw_accuracy", "Dokładność LFW"), + ( + "lr-vs-rof-m.png", + "summary_benchmark/rof_masked_accuracy", + "Dokładność ROF-m", + ), + ( + "lr-vs-rof-s.png", + "summary_benchmark/rof_sunglasses_accuracy", + "Dokładność ROF-s", + ), + ( + "lr-vs-train-accuracy.png", + "max_training/train_accuracy", + "Dokładność na zbiorze treningowym", + ), + ( + "lr-vs-val-accuracy.png", + "max_training/val_accuracy", + "Dokładność na zbiorze walidacyjnym", + ), + ] + + for filename, metric_key, metric_name in plot_args: + make_lr_plot( + run_data=self.sweep_data.sweep_runs_data, + output_dir=self.make_output_subdir("lr"), + filename=filename, + metric_key=metric_key, + metric_name=metric_name, + show=False, + ) + + def margin_v_metrics(self) -> None: + """Create plots showing the effect of margin on various metrics.""" + plot_args = [ + ("margin-vs-lfw.png", "summary_benchmark/lfw_accuracy", "Dokładność LFW"), + ( + "margin-vs-rof-m.png", + "summary_benchmark/rof_masked_accuracy", + "Dokładność ROF-m", + ), + ( + "margin-vs-rof-s.png", + "summary_benchmark/rof_sunglasses_accuracy", + "Dokładność ROF-s", + ), + ( + "margin-vs-train-accuracy.png", + "max_training/train_accuracy", + "Dokładność na zbiorze treningowym", + ), + ( + "margin-vs-val-accuracy.png", + "max_training/val_accuracy", + "Dokładność na zbiorze walidacyjnym", + ), + ] + + for filename, metric_key, metric_name in plot_args: + make_margin_plot( + run_data=self.sweep_data.sweep_runs_data, + output_dir=self.make_output_subdir("margin"), + filename=filename, + metric_key=metric_key, + metric_name=metric_name, + ) + + def augmentation_v_metrics(self) -> None: + """Create plots showing the effect of data augmentation on various metrics.""" + plot_args = [ + ( + "augmentation-vs-lfw.png", + "summary_benchmark/lfw_accuracy", + "Dokładność LFW", + ), + ( + "augmentation-vs-rof-m.png", + "summary_benchmark/rof_masked_accuracy", + "Dokładność ROF-m", + ), + ( + "augmentation-vs-rof-s.png", + "summary_benchmark/rof_sunglasses_accuracy", + "Dokładność ROF-s", + ), + ( + "augmentation-vs-train-accuracy.png", + "max_training/train_accuracy", + "Dokładność na zbiorze treningowym", + ), + ( + "augmentation-vs-val-accuracy.png", + "max_training/val_accuracy", + "Dokładność na zbiorze walidacyjnym", + ), + ] + + for filename, metric_key, metric_name in plot_args: + make_augmentation_plot( + run_data=self.sweep_data.sweep_runs_data, + output_dir=self.make_output_subdir("augmentation"), + filename=filename, + metric_key=metric_key, + metric_name=metric_name, + show=False, + ) + + def only_head_v_metrics(self) -> None: + """Create plots showing the effect of only_head on various metrics and runtime.""" + plot_args = [ + ( + "only-head-vs-lfw.png", + "summary_benchmark/lfw_accuracy", + "Dokładność LFW", + ), + ( + "only-head-vs-rof-m.png", + "summary_benchmark/rof_masked_accuracy", + "Dokładność ROF-m", + ), + ( + "only-head-vs-rof-s.png", + "summary_benchmark/rof_sunglasses_accuracy", + "Dokładność ROF-s", + ), + ( + "only-head-vs-train-accuracy.png", + "max_training/train_accuracy", + "Dokładność na zbiorze treningowym", + ), + ( + "only-head-vs-val-accuracy.png", + "max_training/val_accuracy", + "Dokładność na zbiorze walidacyjnym", + ), + ("only-head-vs-runtime.png", "runtime", "Czas obliczeń [s]"), + ] + + for filename, metric_key, metric_name in plot_args: + make_only_head_plot( + run_data=self.sweep_data.sweep_runs_data, + output_dir=self.make_output_subdir("only-head"), + filename=filename, + metric_key=metric_key, + metric_name=metric_name, + show=False, + ) + + def batch_size_v_metrics(self) -> None: + """Create plots showing the effect of batch size on various metrics.""" + plot_args = [ + ( + "mini-batch-size-vs-lfw.png", + "summary_benchmark/lfw_accuracy", + "Dokładność LFW", + ), + ( + "mini-batch-size-vs-rof-m.png", + "summary_benchmark/rof_masked_accuracy", + "Dokładność ROF-m", + ), + ( + "mini-batch-size-vs-rof-s.png", + "summary_benchmark/rof_sunglasses_accuracy", + "Dokładność ROF-s", + ), + ( + "mini-batch-size-vs-train-accuracy.png", + "max_training/train_accuracy", + "Dokładność na zbiorze treningowym", + ), + ( + "mini-batch-size-vs-val-accuracy.png", + "max_training/val_accuracy", + "Dokładność na zbiorze walidacyjnym", + ), + ] + + for filename, metric_key, metric_name in plot_args: + make_mini_batch_size_plot( + run_data=self.sweep_data.sweep_runs_data, + output_dir=self.make_output_subdir("mini-batch-size"), + filename=filename, + metric_key=metric_key, + metric_name=metric_name, + show=False, + ) + + def parameter_importance_and_correlation( + self, experiment_data_dict: dict[str, Any] + ) -> None: + """Create plots showing parameter importance and correlation with respect to metrics for selected metrics.""" + plot_args = [ + ("parameter-importance-lfw.png", "lfw", "Dokładność LFW"), + ("parameter-importance-rof-m.png", "rof-m", "Dokładność ROF-m"), + ("parameter-importance-rof-s.png", "rof-s", "Dokładność ROF-s"), + ( + "parameter-importance-val-accuracy.png", + "val_accuracy", + "Dokładność walidacyjna", + ), + ] + + for filename, metric_key, metric_name in plot_args: + make_parameter_analysis_plot( + importance_dict=experiment_data_dict[metric_key]["importance"], + correlation_dict=experiment_data_dict[metric_key]["correlation"], + output_dir=self.make_output_subdir("parameter-analysis"), + filename=filename, + title=metric_name, + ) + + def training_curves(self) -> None: + """Create training curves for all runs in the sweep.""" + plot_args = [ + ( + "train-loss-over-epochs.png", + "training/train_loss", + "Strata na zbiorze treningowym", + ), + ( + "val-loss-over-epochs.png", + "training/val_loss", + "Strata na zbiorze walidacyjnym", + ), + ( + "train-accuracy-over-epochs.png", + "training/train_accuracy", + "Dokładność na zbiorze treningowym", + ), + ( + "val-accuracy-over-epochs.png", + "training/val_accuracy", + "Dokładność na zbiorze walidacyjnym", + ), + ] + + for filename, metric_key, metric_name in plot_args: + make_loss_or_acc_plot( + run_histories=self.sweep_data.run_histories, + output_dir=self.make_output_subdir("training"), + filename=filename, + metric_key=metric_key, + metric_name=metric_name, + show=False, + ) + + def augmentation_comparison(self) -> None: + """Group of plots comparing different augmentations with respect to a metric for all metrics.""" + plot_args = [ + ( + "aug-comparison-lfw.png", + "summary_benchmark/lfw_accuracy", + "Dokładność LFW", + "max", + 0.90, + 1.0, + 0.971, + ), + ( + "aug-comparison-rof-m.png", + "summary_benchmark/rof_masked_accuracy", + "Dokładność ROF-m", + "max", + 0.70, + 0.90, + 0.859, + ), + ( + "aug-comparison-rof-s.png", + "summary_benchmark/rof_sunglasses_accuracy", + "Dokładność ROF-s", + "max", + 0.70, + 0.90, + 0.872, + ), + ( + "aug-comparison-eer.png", + "summary_rococo-evaluation/eer", + "EER", + "min", + 0.0, + 0.5, + 0.102, + ), + ( + "aug-comparison-frr-at-far-zero.png", + "summary_rococo-evaluation/frr_at_far_zero", + "FRR @ FAR=0", + "min", + 0.0, + 1.0, + 0.763, + ), + ] + + for ( + filename, + metric_key, + metric_name, + aggregate_best, + xmin, + xmax, + base, + ) in plot_args: + make_aug_comparison_plot( + run_data=self.sweep_data.sweep_runs_data, + output_dir=self.make_output_subdir("augmentation-comparison"), + filename=filename, + metric_key=metric_key, + metric_name=metric_name, + aggregate_best=aggregate_best, + xmin=xmin, + xmax=xmax, + base=base, + show=False, + ) diff --git a/src/plots/sweep_data.py b/src/plots/sweep_data.py new file mode 100644 index 0000000..0a68fd3 --- /dev/null +++ b/src/plots/sweep_data.py @@ -0,0 +1,23 @@ +from dataclasses import dataclass + +from pandas import DataFrame + + +@dataclass +class SweepData: + """Data from a single experiment sweep for plotting. + + Attributes + ---------- + sweep_id : str + The ID of the sweep. + sweep_runs_data : DataFrame + DataFrame containing summary data for all runs in the sweep. + run_histories : list[DataFrame] + List of DataFrames, each containing the history of metrics for a single run. + Used for plotting training curves for individual runs. + """ + + sweep_id: str + sweep_runs_data: DataFrame + run_histories: list[DataFrame] diff --git a/src/plots/wandb_client.py b/src/plots/wandb_client.py new file mode 100644 index 0000000..27fbe21 --- /dev/null +++ b/src/plots/wandb_client.py @@ -0,0 +1,146 @@ +import json +from pathlib import Path + +from pandas import DataFrame + +import wandb +from src.plots.sweep_data import SweepData + + +class WandbClient: + """Wrapper for Weights & Biases API client. + + Handles fetching and caching of sweep data. + Since the API request can be slow and I am worried about rate limits, + fetched data is cached in a local JSON file. + """ + + def __init__(self, project_name: str, cache_dir: Path): + self.api = wandb.Api() + self.project_name = project_name + self.cache_dir = cache_dir + self.cache_dir.mkdir(parents=True, exist_ok=True) + + def get_sweep_data(self, sweep_id: str) -> SweepData: + """Fetch sweep data from cache or API.""" + + if self._get_cache_filepath(sweep_id).exists(): + print(f"Loading sweep data from cache: {sweep_id}") + return self._load_sweep_data_from_json(sweep_id) + else: + print(f"Fetching sweep data from API: {sweep_id}") + sweep_data = self._fetch_sweep_data_from_api(sweep_id) + self._save_sweep_data_to_json(sweep_data) + return sweep_data + + def _fetch_sweep_data_from_api(self, sweep_id: str): + sweep = self._get_sweep_object(sweep_id) + sweep_runs_data = self._extract_sweep_run_data(sweep) + sweep_runs_data = sweep_runs_data[sweep_runs_data["state"] == "finished"] + run_histories = [self._extract_run_history(run) for run in sweep.runs] + return SweepData( + sweep_id=sweep_id, + sweep_runs_data=sweep_runs_data, + run_histories=run_histories, + ) + + def _get_cache_filepath(self, sweep_id: str) -> Path: + return self.cache_dir / f"sweep_{sweep_id}_data.json" + + def _save_sweep_data_to_json(self, sweep_data: SweepData): + with open(self._get_cache_filepath(sweep_data.sweep_id), "w") as f: + json.dump( + { + "sweep_runs_data": sweep_data.sweep_runs_data.to_dict( + orient="records" + ), + "run_histories": [ + rh.to_dict(orient="records") for rh in sweep_data.run_histories + ], + }, + f, + indent=4, + default=str, + ) + + def _load_sweep_data_from_json(self, sweep_id: str) -> SweepData: + with open(self._get_cache_filepath(sweep_id), "r") as f: + data = json.load(f) + sweep_runs_data = DataFrame(data["sweep_runs_data"]) + run_histories = [DataFrame(rh) for rh in data["run_histories"]] + return SweepData( + sweep_id=sweep_id, + sweep_runs_data=sweep_runs_data, + run_histories=run_histories, + ) + + def _extract_sweep_run_data(self, sweep) -> DataFrame: + """Transform Sweep object into a DataFrame. + + Each row corresponds to a single run in the sweep with metrics and config as columns. + """ + data = [] + + for run in sweep.runs: + row = { + "run_id": run.id, + "run_name": run.name, + "state": run.state, + "created_at": run.created_at, + "runtime": run.summary["_runtime"], + } + + # Config parameters + for key, value in run.config.items(): + row[f"config_{key}"] = value + + # Summary metrics + for key, value in run.summary.items(): + if not key.startswith("_"): # Skip internal wandb fields + row[f"summary_{key}"] = value + + # History metric values (final and max) + history = run.history() + if not history.empty: + for col in history.columns: + if not col.startswith("_"): + row[f"final_{col}"] = ( + history[col].iloc[-1] if len(history) > 0 else None + ) + + accuracy_cols = [ + col for col in history.columns if "accuracy" in col.lower() + ] + for col in accuracy_cols: + row[f"max_{col}"] = history[col].max() + + data.append(row) + + return DataFrame(data) + + def _extract_run_history(self, run) -> DataFrame: + """Transform a Run object into a DataFrame of its history. + + History includes metrics logged during training. + Each row corresponds to a single logging step (epoch). + """ + history = run.history() + metrics = [ + "training/train_loss", + "training/val_loss", + "training/train_accuracy", + "training/val_accuracy", + ] + + # Filter to specific metrics (plus _step and _timestamp) + available_metrics = [m for m in metrics if m in history.columns] + cols_to_keep = ["_step", "_timestamp"] + available_metrics + history = history[cols_to_keep] + + history["run_id"] = run.id + history["run_name"] = run.name + + return history + + def _get_sweep_object(self, sweep_id: str): + return self.api.sweep(f"{self.project_name}/{sweep_id}") diff --git a/src/scripts/plots_for_publication.py b/src/scripts/plots_for_publication.py index 10532d1..0c5fa2a 100644 --- a/src/scripts/plots_for_publication.py +++ b/src/scripts/plots_for_publication.py @@ -1,10 +1,12 @@ from dataclasses import dataclass -import wandb from pathlib import Path -import pandas as pd -import json -import matplotlib.pyplot as plt -import numpy as np + +from src.plots.parameter_importance import ( + experiment_05_parameter_importance, + experiment_06_parameter_importance, +) +from src.plots.plot_groups import PlotGroups +from src.plots.wandb_client import WandbClient @dataclass @@ -14,1026 +16,80 @@ class ScriptConfiguration: cache_dir: str = "results/cache" -experiment_05_parameter_importance = { - "lfw": { - "importance": { - "learning_rate": 0.676, - "only_head": 0.250, - "batch_size": 0.043, - "margin": 0.020, - }, - "correlation": { - "learning_rate": -0.464, - "only_head": 0.406, - "batch_size": 0.170, - "margin": -0.116, - }, - }, - "rof-m": { - "importance": { - "learning_rate": 0.695, - "only_head": 0.229, - "batch_size": 0.047, - "margin": 0.017, - }, - "correlation": { - "learning_rate": -0.491, - "only_head": 0.402, - "batch_size": 0.187, - "margin": -0.014, - }, - }, - "rof-s": { - "importance": { - "learning_rate": 0.662, - "only_head": 0.228, - "batch_size": 0.077, - "margin": 0.029, - }, - "correlation": { - "learning_rate": -0.508, - "only_head": 0.427, - "batch_size": 0.195, - "margin": -0.037, - }, - }, - "val_accuracy": { - "importance": { - "learning_rate": 0.484, - "only_head": 0.127, - "batch_size": 0.130, - "margin": 0.107, - }, - "correlation": { - "learning_rate": -0.533, - "only_head": 0.305, - "batch_size": -0.312, - "margin": -0.319, - }, - }, -} - -experiment_06_parameter_importance = { - "lfw": { - "importance": { - "learning_rate": 0.733, - "only_head": 0.206, - "batch_size": 0.012, - "margin": 0.031, - }, - "correlation": { - "learning_rate": -0.490, - "only_head": 0.424, - "batch_size": 0.062, - "margin": -0.141, - }, - }, - "rof-m": { - "importance": { - "learning_rate": 0.755, - "only_head": 0.175, - "batch_size": 0.015, - "margin": 0.033, - }, - "correlation": { - "learning_rate": -0.481, - "only_head": 0.412, - "batch_size": 0.062, - "margin": -0.153, - }, - }, - "rof-s": { - "importance": { - "learning_rate": 0.676, - "only_head": 0.249, - "batch_size": 0.011, - "margin": 0.040, - }, - "correlation": { - "learning_rate": -0.419, - "only_head": 0.499, - "batch_size": -0.009, - "margin": -0.104, - }, - }, - "val_accuracy": { - "importance": { - "learning_rate": 0.370, - "only_head": 0.435, - "batch_size": 0.019, - "margin": 0.042, - }, - "correlation": { - "learning_rate": -0.290, - "only_head": -0.706, - "batch_size": 0.103, - "margin": -0.160, - }, - }, -} - - -@dataclass -class SweepData: - sweep_id: str - sweep_runs_data: pd.DataFrame - run_histories: list[pd.DataFrame] - - -class WandbClient: - def __init__(self, project_name: str, cache_dir: Path): - self.api = wandb.Api() - self.project_name = project_name - self.cache_dir = cache_dir - self.cache_dir.mkdir(parents=True, exist_ok=True) - - def get_sweep_data(self, sweep_id: str) -> SweepData: - if self._get_cache_filepath(sweep_id).exists(): - print(f"Loading sweep data from cache: {sweep_id}") - return self._load_sweep_data_from_json(sweep_id) - else: - print(f"Fetching sweep data from API: {sweep_id}") - sweep_data = self.fetch_sweep_data_from_api(sweep_id) - self._save_sweep_data_to_json(sweep_data) - return sweep_data - - def fetch_sweep_data_from_api(self, sweep_id: str): - sweep = self._get_sweep_object(sweep_id) - sweep_runs_data = self._extract_sweep_run_data(sweep) - sweep_runs_data = sweep_runs_data[sweep_runs_data["state"] == "finished"] - run_histories = [self._extract_run_history(run) for run in sweep.runs] - return SweepData( - sweep_id=sweep_id, - sweep_runs_data=sweep_runs_data, - run_histories=run_histories, - ) - - def _get_cache_filepath(self, sweep_id: str) -> Path: - return self.cache_dir / f"sweep_{sweep_id}_data.json" - - def _save_sweep_data_to_json(self, sweep_data: SweepData): - with open(self._get_cache_filepath(sweep_data.sweep_id), "w") as f: - json.dump( - { - "sweep_runs_data": sweep_data.sweep_runs_data.to_dict( - orient="records" - ), - "run_histories": [ - rh.to_dict(orient="records") for rh in sweep_data.run_histories - ], - }, - f, - indent=4, - default=str, - ) - - def _load_sweep_data_from_json(self, sweep_id: str) -> SweepData: - with open(self._get_cache_filepath(sweep_id), "r") as f: - data = json.load(f) - sweep_runs_data = pd.DataFrame(data["sweep_runs_data"]) - run_histories = [pd.DataFrame(rh) for rh in data["run_histories"]] - return SweepData( - sweep_id=sweep_id, - sweep_runs_data=sweep_runs_data, - run_histories=run_histories, - ) - - def _extract_sweep_run_data(self, sweep) -> pd.DataFrame: - data = [] - - for run in sweep.runs: - row = { - "run_id": run.id, - "run_name": run.name, - "state": run.state, # finished, failed, running, etc. - "created_at": run.created_at, - "runtime": run.summary["_runtime"], - } - - # Add config parameters (hyperparameters) - for key, value in run.config.items(): - row[f"config_{key}"] = value - - # Add summary metrics (final values) - for key, value in run.summary.items(): - if not key.startswith("_"): # Skip internal wandb fields - row[f"summary_{key}"] = value - - # Add history metrics (you can get specific values) - history = run.history() - if not history.empty: - # Get final values - for col in history.columns: - if not col.startswith("_"): - row[f"final_{col}"] = ( - history[col].iloc[-1] if len(history) > 0 else None - ) - - # Get maximum values for accuracy metrics - accuracy_cols = [ - col for col in history.columns if "accuracy" in col.lower() - ] - for col in accuracy_cols: - row[f"max_{col}"] = history[col].max() - - data.append(row) - - return pd.DataFrame(data) - - def _extract_run_history(self, run) -> pd.DataFrame: - history = run.history() - metrics = [ - "training/train_loss", - "training/val_loss", - "training/train_accuracy", - "training/val_accuracy", - ] - - # Filter to specific metrics (plus _step and _timestamp) - available_metrics = [m for m in metrics if m in history.columns] - cols_to_keep = ["_step", "_timestamp"] + available_metrics - history = history[cols_to_keep] - - history["run_id"] = run.id - history["run_name"] = run.name - - return history - - def _get_sweep_object(self, sweep_id: str): - return self.api.sweep(f"{self.project_name}/{sweep_id}") - - -def make_lr_plot( - run_data: pd.DataFrame, - output_dir: Path, - filename: str, - metric_key: str, - metric_name: str, - show: bool = False, -): - plot_data = run_data.copy() - - plt.figure(figsize=(9, 6)) - - head_only_data = plot_data[plot_data["config_only_head"] == True] - full_model_data = plot_data[plot_data["config_only_head"] == False] - - plt.scatter( - head_only_data["config_learning_rate"], - head_only_data[metric_key], - alpha=0.7, - s=60, - c="steelblue", - edgecolors="black", - linewidth=0.5, - label="only_head=True", - ) - - plt.scatter( - full_model_data["config_learning_rate"], - full_model_data[metric_key], - alpha=0.7, - s=60, - c="red", - edgecolors="black", - linewidth=0.5, - label="only_head=False", - ) - - plt.xscale("log") - plt.xlabel("Współczynnik uczenia (skala log)", fontsize=12) - plt.ylabel(metric_name, fontsize=12) - plt.grid(True, alpha=0.3) - plt.legend(fontsize=11, loc="lower left") - plt.gca().xaxis.set_major_formatter(plt.FuncFormatter(lambda x, p: f"{x:.1e}")) - plt.tight_layout() - - plt.savefig(output_dir / filename, dpi=300, bbox_inches="tight") - if show: - plt.show() - plt.close() - - -def make_margin_plot( - run_data: pd.DataFrame, - output_dir: Path, - filename: str, - metric_key: str, - metric_name: str, - show: bool = False, -): - plot_data = run_data.copy() - - plt.figure(figsize=(9, 6)) - - plt.scatter( - plot_data["config_margin"], - plot_data[metric_key], - alpha=0.7, - s=60, - c="green", - edgecolors="black", - linewidth=0.5, - ) - - plt.xlabel("Margines", fontsize=12) - plt.ylabel(metric_name, fontsize=12) - plt.grid(True, alpha=0.3) - plt.tight_layout() - - plt.savefig(output_dir / filename, dpi=300, bbox_inches="tight") - if show: - plt.show() - plt.close() - - -def make_augmentation_plot( - run_data: pd.DataFrame, - output_dir: Path, - filename: str, - metric_key: str, - metric_name: str, - show: bool = False, -): - plot_data = run_data.copy() - plot_data["config_augmentation"] = plot_data["config_augmentation"].fillna("None") - - plt.figure(figsize=(9, 6)) - - augmentations = ["None", "AddRandomRectangleAverageColor"] - - data_by_augmentation = [ - plot_data[plot_data["config_augmentation"] == aug][metric_key].values - for aug in augmentations - ] - - box_plot = plt.boxplot( - data_by_augmentation, - tick_labels=["", ""], - patch_artist=True, - showmeans=True, - vert=False, - ) - - box_plot["boxes"][0].set_facecolor("lightblue") - box_plot["boxes"][0].set_alpha(0.7) - box_plot["boxes"][1].set_facecolor("lightcoral") - box_plot["boxes"][1].set_alpha(0.7) - - # Get the left edge of the entire plot area with some margin - x_min = plt.xlim()[0] - x_range = plt.xlim()[1] - plt.xlim()[0] - x_margin = x_min + (0.05 * x_range) # 5% margin from left edge - - for i, label in enumerate(augmentations): - # Get the top edge of the box for positioning above - box_top = box_plot["boxes"][i].get_path().vertices[:, 1].max() - - plt.text( - x_margin, - box_top + 0.15, - label, - horizontalalignment="left", # Left aligned to plot area - verticalalignment="bottom", - fontsize=10, - ) - - plt.xlabel(metric_name, fontsize=12) - plt.grid(True, alpha=0.3, axis="x") - plt.tight_layout() - - plt.savefig(output_dir / filename, dpi=300, bbox_inches="tight") - if show: - plt.show() - plt.close() - - -def make_only_head_plot( - run_data: pd.DataFrame, - output_dir: Path, - filename: str, - metric_key: str, - metric_name: str, - show: bool = False, -): - plot_data = run_data.copy() - plot_data["config_only_head"] = plot_data["config_only_head"].map( - {True: "True", False: "False"} - ) - - plt.figure(figsize=(9, 6)) - - data_by_only_head = [ - plot_data[plot_data["config_only_head"] == val][metric_key].values - for val in ["True", "False"] - ] - - labels = ["only_head=True", "only_head=False"] - - box_plot = plt.boxplot( - data_by_only_head, - tick_labels=["", ""], - patch_artist=True, - showmeans=True, - vert=False, - ) - - box_plot["boxes"][0].set_facecolor("lightblue") - box_plot["boxes"][0].set_alpha(0.7) - box_plot["boxes"][1].set_facecolor("lightcoral") - box_plot["boxes"][1].set_alpha(0.7) - - # Get the left edge of the entire plot area with some margin - x_min = plt.xlim()[0] - x_range = plt.xlim()[1] - plt.xlim()[0] - x_margin = x_min + (0.05 * x_range) # 5% margin from left edge - - for i, label in enumerate(labels): - # Get the top edge of the box for positioning above - box_top = box_plot["boxes"][i].get_path().vertices[:, 1].max() - - plt.text( - x_margin, - box_top + 0.15, - label, - horizontalalignment="left", # Left aligned to plot area - verticalalignment="bottom", - fontsize=10, +class Plotting: + experiment_05_id = "bwom9hlj" + experiment_06_id = "d5qbai3t" + experiment_07_id = "afr8ycgh" + experiment_08_id = "lgn9xwtm" + experiment_09_id = "72yrgb2d" + + def __init__(self, wandb_client: WandbClient, output_dir: Path): + self.wandb_client = wandb_client + self.output_dir = output_dir + self.output_dir.mkdir(parents=True, exist_ok=True) + + def get_experiment_dir(self, experiment_name: str) -> Path: + dir = self.output_dir / experiment_name + dir.mkdir(parents=True, exist_ok=True) + return dir + + def experiment_05(self): + print("Generating plots for experiment 05...") + sweep_data = self.wandb_client.get_sweep_data(self.experiment_05_id) + experiment_05_output_dir = self.get_experiment_dir("experiment_05") + plot_groups = PlotGroups(sweep_data, experiment_05_output_dir) + + plot_groups.learning_rate_v_metrics() + plot_groups.margin_v_metrics() + plot_groups.augmentation_v_metrics() + plot_groups.only_head_v_metrics() + plot_groups.batch_size_v_metrics() + plot_groups.training_curves() + plot_groups.parameter_importance_and_correlation( + experiment_05_parameter_importance ) - plt.xlabel(metric_name, fontsize=12) - plt.grid(True, alpha=0.3, axis="x") - plt.tick_params(axis="both", which="major", labelsize=9) - plt.tight_layout() - - plt.savefig(output_dir / filename, dpi=300, bbox_inches="tight") - if show: - plt.show() - plt.close() - - -def make_mini_batch_size_plot( - run_data: pd.DataFrame, - output_dir: Path, - filename: str, - metric_key: str, - metric_name: str, - show: bool = False, -): - plot_data = run_data.copy() - - plt.figure(figsize=(9, 6)) - - mini_batch_sizes = sorted(plot_data["config_batch_size"].unique()) - - data_by_mini_batch_size = [ - plot_data[plot_data["config_batch_size"] == size][metric_key].values - for size in mini_batch_sizes - ] - - box_plot = plt.boxplot( - data_by_mini_batch_size, - tick_labels=[str(size) for size in mini_batch_sizes], - patch_artist=True, - showmeans=True, - vert=False, - ) - - colors = plt.cm.viridis(np.linspace(0, 1, len(mini_batch_sizes))) - for i, box in enumerate(box_plot["boxes"]): - box.set_facecolor(colors[i]) - box.set_alpha(0.7) - - plt.xlabel(metric_name, fontsize=12) - plt.ylabel("Rozmiar mini-pakietu", fontsize=12) - plt.grid(True, alpha=0.3, axis="x") - plt.tight_layout() - - plt.savefig(output_dir / filename, dpi=300, bbox_inches="tight") - if show: - plt.show() - plt.close() - - -def make_loss_or_acc_plot( - run_histories: list[pd.DataFrame], - output_dir: Path, - filename: str, - metric_key: str, - metric_name: str, - show: bool = False, -): - plt.figure(figsize=(10, 6)) - - for history in run_histories: - plt.plot( - history["_step"], - history[metric_key], - alpha=0.3, - linewidth=1, - ) - - plt.xlabel("Epoka", fontsize=12) - plt.ylabel(metric_name, fontsize=12) - plt.grid(True, alpha=0.3) - plt.tight_layout() - - plt.savefig(output_dir / filename, dpi=300, bbox_inches="tight") - if show: - plt.show() - plt.close() - - -def make_aug_comparison_plot( - run_data: pd.DataFrame, - output_dir: Path, - filename: str, - metric_key: str, - metric_name: str, - base: float, - aggregate_best: str = "max", # "max" or "min" - xmin: float = 0.0, - xmax: float = 1.0, - show: bool = False, -): - plot_data = run_data.copy() - plot_data["config_augmentation"] = plot_data["config_augmentation"].fillna("None") - - plt.figure(figsize=(9, 6)) - - aug_names = plot_data["config_augmentation"].unique() - best_result_by_aug = ( - plot_data.groupby("config_augmentation") - .agg({metric_key: aggregate_best}) - .reset_index() - .sort_values(by=metric_key, ascending=True) - ) # Changed to True for horizontal bars - - # Create horizontal bar plot - y_pos = np.arange(len(best_result_by_aug)) - plt.barh( - y_pos, - best_result_by_aug[metric_key], - color=plt.cm.viridis(np.linspace(0, 1, len(aug_names))), - alpha=0.7, - edgecolor="black", - ) - - # Set y-axis labels to augmentation names - plt.yticks(y_pos, best_result_by_aug["config_augmentation"]) - - plt.xlim(xmin, xmax) - plt.xlabel(metric_name, fontsize=12) - plt.grid(True, alpha=0.3, axis="x") # Changed to x-axis grid - - # Add vertical line for base score - plt.axvline(x=base, color="red", linestyle="--", label="Baseline") - - plt.tight_layout() - - plt.savefig(output_dir / filename, dpi=300, bbox_inches="tight") - if show: - plt.show() - plt.close() - -def make_parameter_analysis_plot( - importance_dict: dict, - correlation_dict: dict, - output_dir: Path, - filename: str, - title: str, - show: bool = False -): - parameters = list(importance_dict.keys()) - importances = [importance_dict[param] for param in parameters] - correlations = [correlation_dict[param] for param in parameters] - - # Get absolute values for correlation bar heights - abs_correlations = [abs(corr) for corr in correlations] - - # Create colors based on original correlation sign (positive=green, negative=red) - corr_colors = ['green' if corr > 0 else 'red' for corr in correlations] - - fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(9, 6)) - - # Left subplot: Parameter Importance - bars1 = ax1.bar( - parameters, importances, color="skyblue", edgecolor="black", alpha=0.7 - ) - ax1.set_ylabel("Ważność", fontsize=12) - ax1.set_ylim(0, max(importances) * 1.1) - ax1.grid(axis="y", alpha=0.3) - - # Add value labels on top of bars for importance - for bar in bars1: - height = bar.get_height() - ax1.text( - bar.get_x() + bar.get_width() / 2, - height, - f"{height:.3f}", - ha="center", - va="bottom", - fontsize=10, + def experiment_06(self): + print("Generating plots for experiment 06...") + sweep_data = self.wandb_client.get_sweep_data(self.experiment_06_id) + experiment_06_output_dir = self.get_experiment_dir("experiment_06") + plot_groups = PlotGroups(sweep_data, experiment_06_output_dir) + + plot_groups.learning_rate_v_metrics() + plot_groups.margin_v_metrics() + plot_groups.augmentation_v_metrics() + plot_groups.only_head_v_metrics() + plot_groups.batch_size_v_metrics() + plot_groups.training_curves() + plot_groups.parameter_importance_and_correlation( + experiment_06_parameter_importance ) - # Right subplot: Parameter Correlation (absolute values with color coding) - bars2 = ax2.bar( - parameters, abs_correlations, color=corr_colors, edgecolor="black", alpha=0.7 - ) - ax2.set_ylabel("Korelacja", fontsize=12) - ax2.set_ylim(0, max(abs_correlations) * 1.1) - ax2.grid(axis="y", alpha=0.3) + def experiment_07(self): + print("Generating plots for experiment 07...") + sweep_data = self.wandb_client.get_sweep_data(self.experiment_07_id) + experiment_07_output_dir = self.get_experiment_dir("experiment_07") + plot_groups = PlotGroups(sweep_data, experiment_07_output_dir) - # Add value labels on top of bars for correlation (showing original values) - for bar, original_corr in zip(bars2, correlations): - height = bar.get_height() - ax2.text( - bar.get_x() + bar.get_width() / 2, - height, - f"{original_corr:.3f}", - ha="center", - va="bottom", - fontsize=10, - ) + plot_groups.augmentation_comparison() + plot_groups.training_curves() - # Set overall title - fig.suptitle(title, fontsize=16) - - plt.tight_layout() - plt.savefig(output_dir / filename, dpi=300, bbox_inches="tight") - if show: - plt.show() - plt.close() + def experiment_08(self): + print("Generating plots for experiment 08...") + sweep_data = self.wandb_client.get_sweep_data(self.experiment_08_id) + experiment_08_output_dir = self.get_experiment_dir("experiment_08") + plot_groups = PlotGroups(sweep_data, experiment_08_output_dir) -def make_hyperparam_search_experiment_plots(sweep_data: SweepData, output_dir: Path): - # TODO parameter importance plots - """Generate plots for experiments 05 and 06 (hyperparameter search).""" - plot_args = [ - ("lr-vs-lfw.png", "summary_benchmark/lfw_accuracy", "Dokładność LFW"), - ( - "lr-vs-rof-m.png", - "summary_benchmark/rof_masked_accuracy", - "Dokładność ROF-m", - ), - ( - "lr-vs-rof-s.png", - "summary_benchmark/rof_sunglasses_accuracy", - "Dokładność ROF-s", - ), - ( - "lr-vs-train-accuracy.png", - "max_training/train_accuracy", - "Dokładność na zbiorze treningowym", - ), - ( - "lr-vs-val-accuracy.png", - "max_training/val_accuracy", - "Dokładność na zbiorze walidacyjnym", - ), - ] + plot_groups.augmentation_comparison() + plot_groups.training_curves() - dir = output_dir / "lr" - dir.mkdir(parents=True, exist_ok=True) - for filename, metric_key, metric_name in plot_args: - make_lr_plot( - run_data=sweep_data.sweep_runs_data, - output_dir=dir, - filename=filename, - metric_key=metric_key, - metric_name=metric_name, - show=False, - ) + def experiment_09(self): + print("Generating plots for experiment 09...") + sweep_data = self.wandb_client.get_sweep_data(self.experiment_09_id) + experiment_09_output_dir = self.get_experiment_dir("experiment_09") + plot_groups = PlotGroups(sweep_data, experiment_09_output_dir) - plot_args = [ - ("margin-vs-lfw.png", "summary_benchmark/lfw_accuracy", "Dokładność LFW"), - ( - "margin-vs-rof-m.png", - "summary_benchmark/rof_masked_accuracy", - "Dokładność ROF-m", - ), - ( - "margin-vs-rof-s.png", - "summary_benchmark/rof_sunglasses_accuracy", - "Dokładność ROF-s", - ), - ( - "margin-vs-train-accuracy.png", - "max_training/train_accuracy", - "Dokładność na zbiorze treningowym", - ), - ( - "margin-vs-val-accuracy.png", - "max_training/val_accuracy", - "Dokładność na zbiorze walidacyjnym", - ), - ] - - dir = output_dir / "margin" - dir.mkdir(parents=True, exist_ok=True) - for filename, metric_key, metric_name in plot_args: - make_margin_plot( - run_data=sweep_data.sweep_runs_data, - output_dir=dir, - filename=filename, - metric_key=metric_key, - metric_name=metric_name, - ) - - plot_args = [ - ("augmentation-vs-lfw.png", "summary_benchmark/lfw_accuracy", "Dokładność LFW"), - ( - "augmentation-vs-rof-m.png", - "summary_benchmark/rof_masked_accuracy", - "Dokładność ROF-m", - ), - ( - "augmentation-vs-rof-s.png", - "summary_benchmark/rof_sunglasses_accuracy", - "Dokładność ROF-s", - ), - ( - "augmentation-vs-train-accuracy.png", - "max_training/train_accuracy", - "Dokładność na zbiorze treningowym", - ), - ( - "augmentation-vs-val-accuracy.png", - "max_training/val_accuracy", - "Dokładność na zbiorze walidacyjnym", - ), - ] - - dir = output_dir / "augmentation" - dir.mkdir(parents=True, exist_ok=True) - for filename, metric_key, metric_name in plot_args: - make_augmentation_plot( - run_data=sweep_data.sweep_runs_data, - output_dir=dir, - filename=filename, - metric_key=metric_key, - metric_name=metric_name, - show=False, - ) - - plot_args = [ - ("only-head-vs-lfw.png", "summary_benchmark/lfw_accuracy", "Dokładność LFW"), - ( - "only-head-vs-rof-m.png", - "summary_benchmark/rof_masked_accuracy", - "Dokładność ROF-m", - ), - ( - "only-head-vs-rof-s.png", - "summary_benchmark/rof_sunglasses_accuracy", - "Dokładność ROF-s", - ), - ( - "only-head-vs-train-accuracy.png", - "max_training/train_accuracy", - "Dokładność na zbiorze treningowym", - ), - ( - "only-head-vs-val-accuracy.png", - "max_training/val_accuracy", - "Dokładność na zbiorze walidacyjnym", - ), - ("only-head-vs-runtime.png", "runtime", "Czas obliczeń [s]"), - ] - - dir = output_dir / "only_head" - dir.mkdir(parents=True, exist_ok=True) - for filename, metric_key, metric_name in plot_args: - make_only_head_plot( - run_data=sweep_data.sweep_runs_data, - output_dir=dir, - filename=filename, - metric_key=metric_key, - metric_name=metric_name, - show=False, - ) - - plot_args = [ - ( - "mini-batch-size-vs-lfw.png", - "summary_benchmark/lfw_accuracy", - "Dokładność LFW", - ), - ( - "mini-batch-size-vs-rof-m.png", - "summary_benchmark/rof_masked_accuracy", - "Dokładność ROF-m", - ), - ( - "mini-batch-size-vs-rof-s.png", - "summary_benchmark/rof_sunglasses_accuracy", - "Dokładność ROF-s", - ), - ( - "mini-batch-size-vs-train-accuracy.png", - "max_training/train_accuracy", - "Dokładność na zbiorze treningowym", - ), - ( - "mini-batch-size-vs-val-accuracy.png", - "max_training/val_accuracy", - "Dokładność na zbiorze walidacyjnym", - ), - ] - - dir = output_dir / "mini-batch-size" - dir.mkdir(parents=True, exist_ok=True) - for filename, metric_key, metric_name in plot_args: - make_mini_batch_size_plot( - run_data=sweep_data.sweep_runs_data, - output_dir=dir, - filename=filename, - metric_key=metric_key, - metric_name=metric_name, - show=False, - ) - - plot_args = [ - ( - "train-loss-over-epochs.png", - "training/train_loss", - "Strata na zbiorze treningowym", - ), - ( - "val-loss-over-epochs.png", - "training/val_loss", - "Strata na zbiorze walidacyjnym", - ), - ( - "train-accuracy-over-epochs.png", - "training/train_accuracy", - "Dokładność na zbiorze treningowym", - ), - ( - "val-accuracy-over-epochs.png", - "training/val_accuracy", - "Dokładność na zbiorze walidacyjnym", - ), - ] - - dir = output_dir / "training" - dir.mkdir(parents=True, exist_ok=True) - for filename, metric_key, metric_name in plot_args: - make_loss_or_acc_plot( - run_histories=sweep_data.run_histories, - output_dir=dir, - filename=filename, - metric_key=metric_key, - metric_name=metric_name, - show=False, - ) - - -def make_aug_eval_experiment_plots(sweep_data: SweepData, output_dir: Path): - """Generate plots for experiments 07 and 08 (augmentation evaluation).""" - - plot_args = [ - ( - "aug-comparison-lfw.png", - "summary_benchmark/lfw_accuracy", - "Dokładność LFW", - "max", - 0.90, - 1.0, - 0.971, - ), - ( - "aug-comparison-rof-m.png", - "summary_benchmark/rof_masked_accuracy", - "Dokładność ROF-m", - "max", - 0.70, - 0.90, - 0.859, - ), - ( - "aug-comparison-rof-s.png", - "summary_benchmark/rof_sunglasses_accuracy", - "Dokładność ROF-s", - "max", - 0.70, - 0.90, - 0.872, - ), - ( - "aug-comparison-eer.png", - "summary_rococo-evaluation/eer", - "EER", - "min", - 0.0, - 0.5, - 0.102, - ), - ( - "aug-comparison-frr-at-far-zero.png", - "summary_rococo-evaluation/frr_at_far_zero", - "FRR @ FAR=0", - "min", - 0.0, - 1.0, - 0.763, - ), - ] - - dir = output_dir / "augmentation-comparison" - dir.mkdir(parents=True, exist_ok=True) - for ( - filename, - metric_key, - metric_name, - aggregate_best, - xmin, - xmax, - base, - ) in plot_args: - make_aug_comparison_plot( - run_data=sweep_data.sweep_runs_data, - output_dir=dir, - filename=filename, - metric_key=metric_key, - metric_name=metric_name, - aggregate_best=aggregate_best, - xmin=xmin, - xmax=xmax, - base=base, - show=False, - ) - - plot_args = [ - ( - "train-loss-over-epochs.png", - "training/train_loss", - "Strata na zbiorze treningowym", - ), - ( - "val-loss-over-epochs.png", - "training/val_loss", - "Strata na zbiorze walidacyjnym", - ), - ( - "train-accuracy-over-epochs.png", - "training/train_accuracy", - "Dokładność na zbiorze treningowym", - ), - ( - "val-accuracy-over-epochs.png", - "training/val_accuracy", - "Dokładność na zbiorze walidacyjnym", - ), - ] - - dir = output_dir / "training" - dir.mkdir(parents=True, exist_ok=True) - for filename, metric_key, metric_name in plot_args: - make_loss_or_acc_plot( - run_histories=sweep_data.run_histories, - output_dir=dir, - filename=filename, - metric_key=metric_key, - metric_name=metric_name, - show=False, - ) - - -def make_experiment_09_plots(sweep_data: SweepData, output_dir: Path): - plot_args = [ - ( - "train-loss-over-epochs.png", - "training/train_loss", - "Strata na zbiorze treningowym", - ), - ( - "val-loss-over-epochs.png", - "training/val_loss", - "Strata na zbiorze walidacyjnym", - ), - ( - "train-accuracy-over-epochs.png", - "training/train_accuracy", - "Dokładność na zbiorze treningowym", - ), - ( - "val-accuracy-over-epochs.png", - "training/val_accuracy", - "Dokładność na zbiorze walidacyjnym", - ), - ] - - dir = output_dir / "training" - dir.mkdir(parents=True, exist_ok=True) - for filename, metric_key, metric_name in plot_args: - make_loss_or_acc_plot( - run_histories=sweep_data.run_histories, - output_dir=dir, - filename=filename, - metric_key=metric_key, - metric_name=metric_name, - show=False, - ) + plot_groups.training_curves() def main(): @@ -1044,51 +100,12 @@ def main(): wandb_client = WandbClient(config.project_name, Path(config.cache_dir)) - print("Generating plots for experiment 05...") - experiment_05_sweep_id = "bwom9hlj" - sweep_data = wandb_client.get_sweep_data(experiment_05_sweep_id) - experiment_05_output_dir = output_dir / "experiment_05" - experiment_05_output_dir.mkdir(parents=True, exist_ok=True) - make_hyperparam_search_experiment_plots(sweep_data, experiment_05_output_dir) - - - make_parameter_analysis_plot( - importance_dict=experiment_05_parameter_importance["lfw"]["importance"], - correlation_dict=experiment_05_parameter_importance["lfw"]["correlation"], - output_dir=experiment_05_output_dir, - filename="parameter-analysis-lfw.png", - title="LFW", - ) - - # print("Generating plots for experiment 06...") - # experiment_06_sweep_id = "d5qbai3t" - # sweep_data = wandb_client.get_sweep_data(experiment_06_sweep_id) - # experiment_06_output_dir = output_dir / "experiment_06" - # experiment_06_output_dir.mkdir(parents=True, exist_ok=True) - # # Same as 05 - # make_hyperparam_search_experiment_plots(sweep_data, experiment_06_output_dir) - - # print("Generating plots for experiment 07...") - # experiment_07_sweep_id = "afr8ycgh" - # sweep_data = wandb_client.get_sweep_data(experiment_07_sweep_id) - # experiment_07_output_dir = output_dir / "experiment_07" - # experiment_07_output_dir.mkdir(parents=True, exist_ok=True) - # make_aug_eval_experiment_plots(sweep_data, experiment_07_output_dir) - - # print("Generating plots for experiment 08...") - # experiment_08_sweep_id = "lgn9xwtm" - # sweep_data = wandb_client.get_sweep_data(experiment_08_sweep_id) - # experiment_08_output_dir = output_dir / "experiment_08" - # experiment_08_output_dir.mkdir(parents=True, exist_ok=True) - # # Same as 07 - # make_aug_eval_experiment_plots(sweep_data, experiment_08_output_dir) - - # print("Generating plots for experiment 09...") - # experiment_09_id = "72yrgb2d" - # sweep_data = wandb_client.get_sweep_data(experiment_09_id) - # experiment_09_output_dir = output_dir / "experiment_09" - # experiment_09_output_dir.mkdir(parents=True, exist_ok=True) - # make_experiment_09_plots(sweep_data, experiment_09_output_dir) + plotting = Plotting(wandb_client, output_dir) + plotting.experiment_05() + plotting.experiment_06() + plotting.experiment_07() + plotting.experiment_08() + plotting.experiment_09() if __name__ == "__main__": From 521497fc99b93de1a7cabcc158d385c09b3a1f93 Mon Sep 17 00:00:00 2001 From: mgarbowski Date: Thu, 2 Oct 2025 20:20:01 +0200 Subject: [PATCH 13/14] fmt --- notebooks/11-rococo-train-test-split.ipynb | 52 ++-- notebooks/12-plotting_from_wandb_api.ipynb | 301 +++++++++++++++------ src/dataset/rococo_training.py | 22 +- src/training/data_loaders.py | 5 +- src/training/fine_tuning.py | 9 +- 5 files changed, 263 insertions(+), 126 deletions(-) diff --git a/notebooks/11-rococo-train-test-split.ipynb b/notebooks/11-rococo-train-test-split.ipynb index 780ec25..b9b0763 100644 --- a/notebooks/11-rococo-train-test-split.ipynb +++ b/notebooks/11-rococo-train-test-split.ipynb @@ -34,7 +34,7 @@ "import os\n", "import shutil\n", "import random\n", - "import csv\n" + "import csv" ] }, { @@ -52,9 +52,8 @@ } ], "source": [ - "\n", "os.chdir(\"..\")\n", - "print(os.getcwd())\n" + "print(os.getcwd())" ] }, { @@ -98,7 +97,7 @@ "frame_files = sorted(os.listdir(frames_dir))\n", "\n", "print(f\"Number of face files: {len(face_files)}\")\n", - "print(f\"Number of frame files: {len(frame_files)}\") " + "print(f\"Number of frame files: {len(frame_files)}\")" ] }, { @@ -147,7 +146,7 @@ " train_frames = []\n", " test_frames = []\n", " train_face_set = set(face_id_from_filename(f) for f in train_faces)\n", - " \n", + "\n", " for frame in all_frames:\n", "\n", " face_id = face_id_from_filename(frame)\n", @@ -155,8 +154,8 @@ " train_frames.append(frame)\n", " else:\n", " test_frames.append(frame)\n", - " \n", - " return train_frames, test_frames\n" + "\n", + " return train_frames, test_frames" ] }, { @@ -290,7 +289,7 @@ " src = os.path.join(frames_dir, f)\n", " dst = os.path.join(train_frames_dir, f)\n", " shutil.copyfile(src, dst)\n", - " \n", + "\n", " for f in te_frames:\n", " src = os.path.join(frames_dir, f)\n", " dst = os.path.join(test_frames_dir, f)\n", @@ -328,7 +327,7 @@ "for ratio in split_ratios:\n", " split_dir = os.path.join(splits_root, f\"split_{int(ratio*100)}\")\n", " os.makedirs(split_dir, exist_ok=True)\n", - " create_partitioned_set(ratio, dataset_root, split_dir)\n" + " create_partitioned_set(ratio, dataset_root, split_dir)" ] }, { @@ -353,14 +352,17 @@ " frame_files = sorted(os.listdir(frames_dir))\n", " return face_files, frame_files\n", "\n", + "\n", "def get_matches(face_id, frame_files, n_pairs):\n", " matches = [f for f in frame_files if face_id_from_filename(f) == face_id]\n", " return random.sample(matches, n_pairs)\n", "\n", + "\n", "def get_mismatches(face_id, frame_files, n_pairs):\n", " mismatches = [f for f in frame_files if face_id_from_filename(f) != face_id]\n", " return random.sample(mismatches, n_pairs)\n", "\n", + "\n", "def create_pairs(face_files, frame_files, n_pairs_per_face):\n", " match_pairs = []\n", " mismatch_pairs = []\n", @@ -403,7 +405,7 @@ "print(f\"Number of faces: {len(faces)}\")\n", "print(f\"Number of frames: {len(frames)}\")\n", "\n", - "n_train = int(len(faces) * 2/3)\n", + "n_train = int(len(faces) * 2 / 3)\n", "n_val = len(faces) - n_train\n", "print(f\"Number of training faces: {n_train}\")\n", "print(f\"Number of validation faces: {n_val}\")\n", @@ -411,8 +413,8 @@ "train_faces = faces[:n_train]\n", "val_faces = faces[n_train:]\n", "\n", - "train_frames = frames[:n_train*31]\n", - "val_frames = frames[n_train*31:]\n", + "train_frames = frames[: n_train * 31]\n", + "val_frames = frames[n_train * 31 :]\n", "\n", "print(train_faces[-4:])\n", "print(val_faces[:4])\n", @@ -1488,13 +1490,14 @@ ], "source": [ "def save_csv(pairs, filepath):\n", - " with open(filepath, mode='w', newline='') as file:\n", + " with open(filepath, mode=\"w\", newline=\"\") as file:\n", " writer = csv.writer(file)\n", " writer.writerow([\"face\", \"frame\"])\n", " for face, frame in pairs:\n", " writer.writerow([face, frame])\n", " print(f\"Saved {len(pairs)} pairs to {filepath}\")\n", "\n", + "\n", "save_csv(match_train_pairs, \"data/rococo2v3-dev/train_match_pairs.csv\")\n", "save_csv(mismatch_train_pairs, \"data/rococo2v3-dev/train_mismatch_pairs.csv\")\n", "save_csv(match_val_pairs, \"data/rococo2v3-dev/val_match_pairs.csv\")\n", @@ -1529,14 +1532,23 @@ } ], "source": [ - "used_frames = set([\n", - " *(frame for _, frame in match_train_pairs),\n", - " *(frame for _, frame in mismatch_train_pairs),\n", - " *(frame for _, frame in match_val_pairs),\n", - " *(frame for _, frame in mismatch_val_pairs),\n", - "])\n", + "used_frames = set(\n", + " [\n", + " *(frame for _, frame in match_train_pairs),\n", + " *(frame for _, frame in mismatch_train_pairs),\n", + " *(frame for _, frame in match_val_pairs),\n", + " *(frame for _, frame in mismatch_val_pairs),\n", + " ]\n", + ")\n", "\n", - "len(used_frames), sum((len(match_train_pairs), len(mismatch_train_pairs), len(match_val_pairs), len(mismatch_val_pairs)))" + "len(used_frames), sum(\n", + " (\n", + " len(match_train_pairs),\n", + " len(mismatch_train_pairs),\n", + " len(match_val_pairs),\n", + " len(mismatch_val_pairs),\n", + " )\n", + ")" ] }, { diff --git a/notebooks/12-plotting_from_wandb_api.ipynb b/notebooks/12-plotting_from_wandb_api.ipynb index 4b83b8d..ea989c8 100644 --- a/notebooks/12-plotting_from_wandb_api.ipynb +++ b/notebooks/12-plotting_from_wandb_api.ipynb @@ -25,48 +25,52 @@ "def extract_run_data(runs):\n", " \"\"\"\n", " Extract parameters, metrics, and metadata from runs.\n", - " \n", + "\n", " Args:\n", " runs: List of wandb Run objects\n", - " \n", + "\n", " Returns:\n", " pandas.DataFrame with all run data\n", " \"\"\"\n", " data = []\n", - " \n", + "\n", " for run in runs:\n", " row = {\n", - " 'run_id': run.id,\n", - " 'run_name': run.name,\n", - " 'state': run.state, # finished, failed, running, etc.\n", - " 'created_at': run.created_at,\n", - " 'runtime': run.summary[\"_runtime\"],\n", + " \"run_id\": run.id,\n", + " \"run_name\": run.name,\n", + " \"state\": run.state, # finished, failed, running, etc.\n", + " \"created_at\": run.created_at,\n", + " \"runtime\": run.summary[\"_runtime\"],\n", " }\n", - " \n", + "\n", " # Add config parameters (hyperparameters)\n", " for key, value in run.config.items():\n", - " row[f'config_{key}'] = value\n", - " \n", + " row[f\"config_{key}\"] = value\n", + "\n", " # Add summary metrics (final values)\n", " for key, value in run.summary.items():\n", - " if not key.startswith('_'): # Skip internal wandb fields\n", - " row[f'summary_{key}'] = value\n", - " \n", + " if not key.startswith(\"_\"): # Skip internal wandb fields\n", + " row[f\"summary_{key}\"] = value\n", + "\n", " # Add history metrics (you can get specific values)\n", " history = run.history()\n", " if not history.empty:\n", " # Get final values\n", " for col in history.columns:\n", - " if not col.startswith('_'):\n", - " row[f'final_{col}'] = history[col].iloc[-1] if len(history) > 0 else None\n", - " \n", + " if not col.startswith(\"_\"):\n", + " row[f\"final_{col}\"] = (\n", + " history[col].iloc[-1] if len(history) > 0 else None\n", + " )\n", + "\n", " # Get maximum values for accuracy metrics\n", - " accuracy_cols = [col for col in history.columns if 'accuracy' in col.lower()]\n", + " accuracy_cols = [\n", + " col for col in history.columns if \"accuracy\" in col.lower()\n", + " ]\n", " for col in accuracy_cols:\n", - " row[f'max_{col}'] = history[col].max()\n", - " \n", + " row[f\"max_{col}\"] = history[col].max()\n", + "\n", " data.append(row)\n", - " \n", + "\n", " return pd.DataFrame(data)" ] }, @@ -80,25 +84,25 @@ "def get_run_history(run, metrics=None):\n", " \"\"\"\n", " Get full history for specific metrics from a run.\n", - " \n", + "\n", " Args:\n", " run: wandb Run object\n", " metrics: List of metric names to retrieve (None for all)\n", - " \n", + "\n", " Returns:\n", " pandas.DataFrame with timestamped metrics\n", " \"\"\"\n", " history = run.history()\n", - " \n", + "\n", " if metrics:\n", " # Filter to specific metrics (plus _step and _timestamp)\n", " available_metrics = [m for m in metrics if m in history.columns]\n", - " cols_to_keep = ['_step', '_timestamp'] + available_metrics\n", + " cols_to_keep = [\"_step\", \"_timestamp\"] + available_metrics\n", " history = history[cols_to_keep]\n", - " \n", - " history['run_id'] = run.id\n", - " history['run_name'] = run.name\n", - " \n", + "\n", + " history[\"run_id\"] = run.id\n", + " history[\"run_name\"] = run.name\n", + "\n", " return history" ] }, @@ -137,7 +141,7 @@ "metadata": {}, "outputs": [], "source": [ - "run_data = run_data[run_data['state'] == 'finished']" + "run_data = run_data[run_data[\"state\"] == \"finished\"]" ] }, { @@ -190,7 +194,15 @@ "outputs": [], "source": [ "example_run = runs[5]\n", - "example_run_history = get_run_history(example_run, metrics=[\"training/train_loss\", \"training/val_loss\", \"training/train_accuracy\", \"training/val_accuracy\"])" + "example_run_history = get_run_history(\n", + " example_run,\n", + " metrics=[\n", + " \"training/train_loss\",\n", + " \"training/val_loss\",\n", + " \"training/train_accuracy\",\n", + " \"training/val_accuracy\",\n", + " ],\n", + ")" ] }, { @@ -230,48 +242,55 @@ "metadata": {}, "outputs": [], "source": [ - "def make_lr_plot(run_data: pd.DataFrame, output_dir: Path, filename: str, metric_key: str, metric_name: str, show: bool = False):\n", + "def make_lr_plot(\n", + " run_data: pd.DataFrame,\n", + " output_dir: Path,\n", + " filename: str,\n", + " metric_key: str,\n", + " metric_name: str,\n", + " show: bool = False,\n", + "):\n", " plot_data = run_data.copy()\n", "\n", " plt.figure(figsize=(9, 6))\n", "\n", - " head_only_data = plot_data[plot_data['config_only_head'] == True]\n", - " full_model_data = plot_data[plot_data['config_only_head'] == False]\n", + " head_only_data = plot_data[plot_data[\"config_only_head\"] == True]\n", + " full_model_data = plot_data[plot_data[\"config_only_head\"] == False]\n", "\n", " plt.scatter(\n", - " head_only_data['config_learning_rate'], \n", + " head_only_data[\"config_learning_rate\"],\n", " head_only_data[metric_key],\n", " alpha=0.7,\n", " s=60,\n", - " c='steelblue',\n", - " edgecolors='black',\n", + " c=\"steelblue\",\n", + " edgecolors=\"black\",\n", " linewidth=0.5,\n", - " label='only_head=True'\n", + " label=\"only_head=True\",\n", " )\n", "\n", " plt.scatter(\n", - " full_model_data['config_learning_rate'], \n", + " full_model_data[\"config_learning_rate\"],\n", " full_model_data[metric_key],\n", " alpha=0.7,\n", " s=60,\n", - " c='red',\n", - " edgecolors='black',\n", + " c=\"red\",\n", + " edgecolors=\"black\",\n", " linewidth=0.5,\n", - " label='only_head=False'\n", + " label=\"only_head=False\",\n", " )\n", "\n", - " plt.xscale('log')\n", - " plt.xlabel('Współczynnik uczenia (skala log)', fontsize=12)\n", + " plt.xscale(\"log\")\n", + " plt.xlabel(\"Współczynnik uczenia (skala log)\", fontsize=12)\n", " plt.ylabel(metric_name, fontsize=12)\n", " plt.grid(True, alpha=0.3)\n", - " plt.legend(fontsize=11, loc='lower left')\n", - " plt.gca().xaxis.set_major_formatter(plt.FuncFormatter(lambda x, p: f'{x:.1e}'))\n", + " plt.legend(fontsize=11, loc=\"lower left\")\n", + " plt.gca().xaxis.set_major_formatter(plt.FuncFormatter(lambda x, p: f\"{x:.1e}\"))\n", " plt.tight_layout()\n", "\n", - " plt.savefig(output_dir / filename, dpi=300, bbox_inches='tight')\n", + " plt.savefig(output_dir / filename, dpi=300, bbox_inches=\"tight\")\n", " if show:\n", " plt.show()\n", - " plt.close()\n" + " plt.close()" ] }, { @@ -287,9 +306,21 @@ "plot_args = [\n", " (\"lr-vs-lfw.png\", \"summary_benchmark/lfw_accuracy\", \"Dokładność LFW\"),\n", " (\"lr-vs-rof-m.png\", \"summary_benchmark/rof_masked_accuracy\", \"Dokładność ROF-m\"),\n", - " (\"lr-vs-rof-s.png\", \"summary_benchmark/rof_sunglasses_accuracy\", \"Dokładność ROF-s\"),\n", - " (\"lr-vs-train-accuracy.png\", \"max_training/train_accuracy\", \"Dokładność na zbiorze treningowym\"),\n", - " (\"lr-vs-val-accuracy.png\", \"max_training/val_accuracy\", \"Dokładność na zbiorze walidacyjnym\"),\n", + " (\n", + " \"lr-vs-rof-s.png\",\n", + " \"summary_benchmark/rof_sunglasses_accuracy\",\n", + " \"Dokładność ROF-s\",\n", + " ),\n", + " (\n", + " \"lr-vs-train-accuracy.png\",\n", + " \"max_training/train_accuracy\",\n", + " \"Dokładność na zbiorze treningowym\",\n", + " ),\n", + " (\n", + " \"lr-vs-val-accuracy.png\",\n", + " \"max_training/val_accuracy\",\n", + " \"Dokładność na zbiorze walidacyjnym\",\n", + " ),\n", "]\n", "\n", "for filename, metric_key, metric_name in plot_args:\n", @@ -299,7 +330,7 @@ " filename=filename,\n", " metric_key=metric_key,\n", " metric_name=metric_name,\n", - " show=True\n", + " show=True,\n", " )" ] }, @@ -310,30 +341,37 @@ "metadata": {}, "outputs": [], "source": [ - "def make_margin_plot(run_data: pd.DataFrame, output_dir: Path, filename: str, metric_key: str, metric_name: str, show: bool = False):\n", + "def make_margin_plot(\n", + " run_data: pd.DataFrame,\n", + " output_dir: Path,\n", + " filename: str,\n", + " metric_key: str,\n", + " metric_name: str,\n", + " show: bool = False,\n", + "):\n", " plot_data = run_data.copy()\n", "\n", " plt.figure(figsize=(9, 6))\n", "\n", " plt.scatter(\n", - " plot_data['config_margin'], \n", + " plot_data[\"config_margin\"],\n", " plot_data[metric_key],\n", " alpha=0.7,\n", " s=60,\n", - " c='green',\n", - " edgecolors='black',\n", + " c=\"green\",\n", + " edgecolors=\"black\",\n", " linewidth=0.5,\n", " )\n", "\n", - " plt.xlabel('Margines', fontsize=12)\n", + " plt.xlabel(\"Margines\", fontsize=12)\n", " plt.ylabel(metric_name, fontsize=12)\n", " plt.grid(True, alpha=0.3)\n", " plt.tight_layout()\n", - " \n", - " plt.savefig(output_dir / filename, dpi=300, bbox_inches='tight')\n", + "\n", + " plt.savefig(output_dir / filename, dpi=300, bbox_inches=\"tight\")\n", " if show:\n", " plt.show()\n", - " plt.close()\n" + " plt.close()" ] }, { @@ -345,10 +383,26 @@ "source": [ "plot_args = [\n", " (\"margin-vs-lfw.png\", \"summary_benchmark/lfw_accuracy\", \"Dokładność LFW\"),\n", - " (\"margin-vs-rof-m.png\", \"summary_benchmark/rof_masked_accuracy\", \"Dokładność ROF-m\"),\n", - " (\"margin-vs-rof-s.png\", \"summary_benchmark/rof_sunglasses_accuracy\", \"Dokładność ROF-s\"),\n", - " (\"margin-vs-train-accuracy.png\", \"max_training/train_accuracy\", \"Dokładność na zbiorze treningowym\"),\n", - " (\"margin-vs-val-accuracy.png\", \"max_training/val_accuracy\", \"Dokładność na zbiorze walidacyjnym\"),\n", + " (\n", + " \"margin-vs-rof-m.png\",\n", + " \"summary_benchmark/rof_masked_accuracy\",\n", + " \"Dokładność ROF-m\",\n", + " ),\n", + " (\n", + " \"margin-vs-rof-s.png\",\n", + " \"summary_benchmark/rof_sunglasses_accuracy\",\n", + " \"Dokładność ROF-s\",\n", + " ),\n", + " (\n", + " \"margin-vs-train-accuracy.png\",\n", + " \"max_training/train_accuracy\",\n", + " \"Dokładność na zbiorze treningowym\",\n", + " ),\n", + " (\n", + " \"margin-vs-val-accuracy.png\",\n", + " \"max_training/val_accuracy\",\n", + " \"Dokładność na zbiorze walidacyjnym\",\n", + " ),\n", "]\n", "\n", "for filename, metric_key, metric_name in plot_args:\n", @@ -357,7 +411,7 @@ " output_dir=experiment_05_plots_dir,\n", " filename=filename,\n", " metric_key=metric_key,\n", - " metric_name=metric_name\n", + " metric_name=metric_name,\n", " )" ] }, @@ -368,15 +422,25 @@ "metadata": {}, "outputs": [], "source": [ - "def make_augmentation_plot(run_data: pd.DataFrame, output_dir: Path, filename: str, metric_key: str, metric_name: str, show: bool = False):\n", + "def make_augmentation_plot(\n", + " run_data: pd.DataFrame,\n", + " output_dir: Path,\n", + " filename: str,\n", + " metric_key: str,\n", + " metric_name: str,\n", + " show: bool = False,\n", + "):\n", " plot_data = run_data.copy()\n", - " plot_data['config_augmentation'] = plot_data['config_augmentation'].fillna('None')\n", + " plot_data[\"config_augmentation\"] = plot_data[\"config_augmentation\"].fillna(\"None\")\n", "\n", " plt.figure(figsize=(9, 6))\n", "\n", - " augmentations = ['None', 'AddRandomRectangleAverageColor']\n", + " augmentations = [\"None\", \"AddRandomRectangleAverageColor\"]\n", "\n", - " data_by_augmentation = [plot_data[plot_data['config_augmentation'] == aug][metric_key].values for aug in augmentations]\n", + " data_by_augmentation = [\n", + " plot_data[plot_data[\"config_augmentation\"] == aug][metric_key].values\n", + " for aug in augmentations\n", + " ]\n", "\n", " box_plot = plt.boxplot(\n", " data_by_augmentation,\n", @@ -385,7 +449,7 @@ " showmeans=True,\n", " vert=False,\n", " )\n", - " \n", + "\n", " box_plot[\"boxes\"][0].set_facecolor(\"lightblue\")\n", " box_plot[\"boxes\"][0].set_alpha(0.7)\n", " box_plot[\"boxes\"][1].set_facecolor(\"lightcoral\")\n", @@ -410,13 +474,13 @@ " )\n", "\n", " plt.xlabel(metric_name, fontsize=12)\n", - " plt.grid(True, alpha=0.3, axis='x')\n", + " plt.grid(True, alpha=0.3, axis=\"x\")\n", " plt.tight_layout()\n", "\n", - " plt.savefig(output_dir / filename, dpi=300, bbox_inches='tight')\n", + " plt.savefig(output_dir / filename, dpi=300, bbox_inches=\"tight\")\n", " if show:\n", " plt.show()\n", - " plt.close()\n" + " plt.close()" ] }, { @@ -428,10 +492,26 @@ "source": [ "plot_args = [\n", " (\"augmentation-vs-lfw.png\", \"summary_benchmark/lfw_accuracy\", \"Dokładność LFW\"),\n", - " (\"augmentation-vs-rof-m.png\", \"summary_benchmark/rof_masked_accuracy\", \"Dokładność ROF-m\"),\n", - " (\"augmentation-vs-rof-s.png\", \"summary_benchmark/rof_sunglasses_accuracy\", \"Dokładność ROF-s\"),\n", - " (\"augmentation-vs-train-accuracy.png\", \"max_training/train_accuracy\", \"Dokładność na zbiorze treningowym\"),\n", - " (\"augmentation-vs-val-accuracy.png\", \"max_training/val_accuracy\", \"Dokładność na zbiorze walidacyjnym\"),\n", + " (\n", + " \"augmentation-vs-rof-m.png\",\n", + " \"summary_benchmark/rof_masked_accuracy\",\n", + " \"Dokładność ROF-m\",\n", + " ),\n", + " (\n", + " \"augmentation-vs-rof-s.png\",\n", + " \"summary_benchmark/rof_sunglasses_accuracy\",\n", + " \"Dokładność ROF-s\",\n", + " ),\n", + " (\n", + " \"augmentation-vs-train-accuracy.png\",\n", + " \"max_training/train_accuracy\",\n", + " \"Dokładność na zbiorze treningowym\",\n", + " ),\n", + " (\n", + " \"augmentation-vs-val-accuracy.png\",\n", + " \"max_training/val_accuracy\",\n", + " \"Dokładność na zbiorze walidacyjnym\",\n", + " ),\n", "]\n", "\n", "for filename, metric_key, metric_name in plot_args:\n", @@ -525,10 +605,26 @@ "source": [ "plot_args = [\n", " (\"only-head-vs-lfw.png\", \"summary_benchmark/lfw_accuracy\", \"Dokładność LFW\"),\n", - " (\"only-head-vs-rof-m.png\", \"summary_benchmark/rof_masked_accuracy\", \"Dokładność ROF-m\"),\n", - " (\"only-head-vs-rof-s.png\", \"summary_benchmark/rof_sunglasses_accuracy\", \"Dokładność ROF-s\"),\n", - " (\"only-head-vs-train-accuracy.png\", \"max_training/train_accuracy\", \"Dokładność na zbiorze treningowym\"),\n", - " (\"only-head-vs-val-accuracy.png\", \"max_training/val_accuracy\", \"Dokładność na zbiorze walidacyjnym\"),\n", + " (\n", + " \"only-head-vs-rof-m.png\",\n", + " \"summary_benchmark/rof_masked_accuracy\",\n", + " \"Dokładność ROF-m\",\n", + " ),\n", + " (\n", + " \"only-head-vs-rof-s.png\",\n", + " \"summary_benchmark/rof_sunglasses_accuracy\",\n", + " \"Dokładność ROF-s\",\n", + " ),\n", + " (\n", + " \"only-head-vs-train-accuracy.png\",\n", + " \"max_training/train_accuracy\",\n", + " \"Dokładność na zbiorze treningowym\",\n", + " ),\n", + " (\n", + " \"only-head-vs-val-accuracy.png\",\n", + " \"max_training/val_accuracy\",\n", + " \"Dokładność na zbiorze walidacyjnym\",\n", + " ),\n", " (\"only-head-vs-runtime.png\", \"runtime\", \"Czas obliczeń [s]\"),\n", "]\n", "\n", @@ -582,7 +678,6 @@ " box.set_facecolor(colors[i])\n", " box.set_alpha(0.7)\n", "\n", - "\n", " plt.xlabel(metric_name, fontsize=12)\n", " plt.ylabel(\"Rozmiar mini-pakietu\", fontsize=12)\n", " plt.grid(True, alpha=0.3, axis=\"x\")\n", @@ -603,10 +698,26 @@ "source": [ "plot_args = [\n", " (\"mini-batch-size-vs-lfw.png\", \"summary_benchmark/lfw_accuracy\", \"Dokładność LFW\"),\n", - " (\"mini-batch-size-vs-rof-m.png\", \"summary_benchmark/rof_masked_accuracy\", \"Dokładność ROF-m\"),\n", - " (\"mini-batch-size-vs-rof-s.png\", \"summary_benchmark/rof_sunglasses_accuracy\", \"Dokładność ROF-s\"),\n", - " (\"mini-batch-size-vs-train-accuracy.png\", \"max_training/train_accuracy\", \"Dokładność na zbiorze treningowym\"),\n", - " (\"mini-batch-size-vs-val-accuracy.png\", \"max_training/val_accuracy\", \"Dokładność na zbiorze walidacyjnym\"),\n", + " (\n", + " \"mini-batch-size-vs-rof-m.png\",\n", + " \"summary_benchmark/rof_masked_accuracy\",\n", + " \"Dokładność ROF-m\",\n", + " ),\n", + " (\n", + " \"mini-batch-size-vs-rof-s.png\",\n", + " \"summary_benchmark/rof_sunglasses_accuracy\",\n", + " \"Dokładność ROF-s\",\n", + " ),\n", + " (\n", + " \"mini-batch-size-vs-train-accuracy.png\",\n", + " \"max_training/train_accuracy\",\n", + " \"Dokładność na zbiorze treningowym\",\n", + " ),\n", + " (\n", + " \"mini-batch-size-vs-val-accuracy.png\",\n", + " \"max_training/val_accuracy\",\n", + " \"Dokładność na zbiorze walidacyjnym\",\n", + " ),\n", "]\n", "\n", "for filename, metric_key, metric_name in plot_args:\n", @@ -664,10 +775,22 @@ "outputs": [], "source": [ "plot_args = [\n", - " (\"train-loss-over-epochs.png\", \"training/train_loss\", \"Strata na zbiorze treningowym\"),\n", + " (\n", + " \"train-loss-over-epochs.png\",\n", + " \"training/train_loss\",\n", + " \"Strata na zbiorze treningowym\",\n", + " ),\n", " (\"val-loss-over-epochs.png\", \"training/val_loss\", \"Strata na zbiorze walidacyjnym\"),\n", - " (\"train-accuracy-over-epochs.png\", \"training/train_accuracy\", \"Dokładność na zbiorze treningowym\"),\n", - " (\"val-accuracy-over-epochs.png\", \"training/val_accuracy\", \"Dokładność na zbiorze walidacyjnym\"),\n", + " (\n", + " \"train-accuracy-over-epochs.png\",\n", + " \"training/train_accuracy\",\n", + " \"Dokładność na zbiorze treningowym\",\n", + " ),\n", + " (\n", + " \"val-accuracy-over-epochs.png\",\n", + " \"training/val_accuracy\",\n", + " \"Dokładność na zbiorze walidacyjnym\",\n", + " ),\n", "]\n", "\n", "for filename, metric_key, metric_name in plot_args:\n", diff --git a/src/dataset/rococo_training.py b/src/dataset/rococo_training.py index 3785f2e..d2789b5 100644 --- a/src/dataset/rococo_training.py +++ b/src/dataset/rococo_training.py @@ -1,7 +1,10 @@ from typing import override -from src.dataset.face_pairs import FacePairsDataset + from PIL import Image +from src.dataset.face_pairs import FacePairsDataset + + class RococoTrainingDataset(FacePairsDataset): def __init__( self, root_dir: str, pairs: list[tuple], transform_1=None, transform_2=None @@ -11,7 +14,6 @@ def __init__( self.transform_1 = transform_1 self.transform_2 = transform_2 - @staticmethod def load_image(path: str): return Image.open(path).convert("RGB") @@ -34,25 +36,25 @@ def from_match_and_mismatch_pairs( f.readline() # skip header for line in f: face_path, frame_path = line.strip().split(",") - match_pairs.append(( - f"{faces_dir}/{face_path}", f"{frames_dir}/{frame_path}", 1 - )) # 1 for same person + match_pairs.append( + (f"{faces_dir}/{face_path}", f"{frames_dir}/{frame_path}", 1) + ) # 1 for same person with open(mismatch_pairs_file, "r") as f: f.readline() # skip header for line in f: face_path, frame_path = line.strip().split(",") - mismatch_pairs.append(( - f"{faces_dir}/{face_path}", f"{frames_dir}/{frame_path}", -1 - )) # -1 for different people + mismatch_pairs.append( + (f"{faces_dir}/{face_path}", f"{frames_dir}/{frame_path}", -1) + ) # -1 for different people pairs = match_pairs + mismatch_pairs return cls(root_dir, pairs, transform_1, transform_2) - + @override def __len__(self): return len(self.pairs) - + @override def __getitem__(self, idx: int) -> tuple: face_path, frame_path, label = self.pairs[idx] diff --git a/src/training/data_loaders.py b/src/training/data_loaders.py index 8bdc150..28cbce7 100644 --- a/src/training/data_loaders.py +++ b/src/training/data_loaders.py @@ -3,6 +3,7 @@ from src.dataset.lfw import LFWDataset from src.dataset.rococo_training import RococoTrainingDataset + def get_lfw_loaders( transform, transform_with_augmentation, @@ -48,6 +49,7 @@ def get_lfw_loaders( "val": val_loader, } + def get_rococo_loaders( transform_face, transform_frame, @@ -56,7 +58,6 @@ def get_rococo_loaders( num_workers=4, pin_memory=True, ) -> dict[str, DataLoader]: - train_set = RococoTrainingDataset.from_match_and_mismatch_pairs( root_dir=rococo_root, @@ -97,4 +98,4 @@ def get_rococo_loaders( return { "train": train_loader, "val": val_loader, - } \ No newline at end of file + } diff --git a/src/training/fine_tuning.py b/src/training/fine_tuning.py index f44722b..df4650e 100644 --- a/src/training/fine_tuning.py +++ b/src/training/fine_tuning.py @@ -28,7 +28,8 @@ def get_optimizer(config: FineTuningConfig, model_params) -> Optimizer: ) else: raise ValueError(f"Unknown optimizer type: {config.optimizer_type}") - + + def get_data_loaders(config: FineTuningConfig) -> dict[str, DataLoader]: extra_transform = get_augmentation_transform(config.augmentation) aug_transform = transforms.Compose( @@ -48,14 +49,12 @@ def get_data_loaders(config: FineTuningConfig) -> dict[str, DataLoader]: ) if config.dataset == "lfw": - return get_lfw_loaders( - transform, aug_transform, batch_size=config.batch_size - ) + return get_lfw_loaders(transform, aug_transform, batch_size=config.batch_size) elif config.dataset == "rococo": return get_rococo_loaders( transform_face=aug_transform, transform_frame=transform, - batch_size=config.batch_size + batch_size=config.batch_size, ) else: raise ValueError(f"Unknown dataset: {config.dataset}") From 6fd1c803c3f548d28923e8c2debd2dc29a2014c1 Mon Sep 17 00:00:00 2001 From: mgarbowski Date: Thu, 2 Oct 2025 20:21:08 +0200 Subject: [PATCH 14/14] add comment --- src/plots/plot_groups.py | 1 + 1 file changed, 1 insertion(+) diff --git a/src/plots/plot_groups.py b/src/plots/plot_groups.py index 27a84df..5fc917a 100644 --- a/src/plots/plot_groups.py +++ b/src/plots/plot_groups.py @@ -284,6 +284,7 @@ def training_curves(self) -> None: def augmentation_comparison(self) -> None: """Group of plots comparing different augmentations with respect to a metric for all metrics.""" + # Baseline values hardcoded plot_args = [ ( "aug-comparison-lfw.png",