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 new file mode 100644 index 0000000..ea989c8 --- /dev/null +++ b/notebooks/12-plotting_from_wandb_api.ipynb @@ -0,0 +1,829 @@ +{ + "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}\"] = (\n", + " history[col].iloc[-1] if len(history) > 0 else None\n", + " )\n", + "\n", + " # Get maximum values for accuracy metrics\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", + " 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(\n", + " example_run,\n", + " metrics=[\n", + " \"training/train_loss\",\n", + " \"training/val_loss\",\n", + " \"training/train_accuracy\",\n", + " \"training/val_accuracy\",\n", + " ],\n", + ")" + ] + }, + { + "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 \n", + "\n", + "TODO: parameter importance data has to be scraped manually to recreate plots" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "59c345bf", + "metadata": {}, + "outputs": [], + "source": [ + "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", + "\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()" + ] + }, + { + "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", + " (\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", + " 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(\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[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()" + ] + }, + { + "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", + " (\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", + " 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(\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", + "\n", + " plt.figure(figsize=(9, 6))\n", + "\n", + " augmentations = [\"None\", \"AddRandomRectangleAverageColor\"]\n", + "\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", + " 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()" + ] + }, + { + "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", + " (\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", + " 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", + " (\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", + "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", + " 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", + " (\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", + " 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", + " (\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", + " (\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", + " 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 +} 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/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/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..5fc917a --- /dev/null +++ b/src/plots/plot_groups.py @@ -0,0 +1,356 @@ +"""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.""" + # Baseline values hardcoded + 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 new file mode 100644 index 0000000..0c5fa2a --- /dev/null +++ b/src/scripts/plots_for_publication.py @@ -0,0 +1,112 @@ +from dataclasses import dataclass +from pathlib import Path + +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 +class ScriptConfiguration: + output_dir: str = "results/plots" + project_name: str = "thesis" + cache_dir: str = "results/cache" + + +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 + ) + + 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 + ) + + 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) + + plot_groups.augmentation_comparison() + plot_groups.training_curves() + + 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) + + plot_groups.augmentation_comparison() + plot_groups.training_curves() + + 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_groups.training_curves() + + +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)) + + plotting = Plotting(wandb_client, output_dir) + plotting.experiment_05() + plotting.experiment_06() + plotting.experiment_07() + plotting.experiment_08() + plotting.experiment_09() + + +if __name__ == "__main__": + main() 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}")