airfRANS-model-exploration/notebooks/explore_test_run.ipynb

1153 lines
354 KiB
Text
Raw Permalink Normal View History

2026-07-23 05:47:43 +00:00
{
"cells": [
{
"cell_type": "markdown",
"id": "be7a6c7e",
"metadata": {},
"source": [
"# Explore a completed training run\n",
"\n",
"Use this notebook to sanity-check the artifacts written by `airfrans-frontier train`: final metrics, training curves, split manifest, normalization stats, checkpoint contents, and reloaded-model predictions. It reads existing artifacts only.\n"
]
},
{
"cell_type": "markdown",
"id": "23b531cc",
"metadata": {},
"source": [
"## Setup\n",
"\n",
"Run from the repository root or from `notebooks/`. Select the repository `.venv` kernel. If imports fail, run `uv sync --dev` from the repository root and restart the kernel.\n"
]
},
{
"cell_type": "code",
"execution_count": 1,
"id": "18ae45d0",
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"python: /home/aaron/data/airfrans/.venv/bin/python\n",
"repo: /home/aaron/data/airfrans\n"
]
}
],
"source": [
"from __future__ import annotations\n",
"\n",
"import json\n",
"import os\n",
"import sys\n",
"from pathlib import Path\n",
"\n",
"for path in list(sys.path):\n",
" if \"swactor-mvp/pydeps\" in path:\n",
" sys.path.remove(path)\n",
"if \"PYTHONPATH\" in os.environ:\n",
" os.environ[\"PYTHONPATH\"] = os.pathsep.join(\n",
" entry for entry in os.environ[\"PYTHONPATH\"].split(os.pathsep) if \"swactor-mvp/pydeps\" not in entry\n",
" )\n",
"\n",
"print(f\"python: {sys.executable}\")\n",
"\n",
"try:\n",
" import matplotlib.pyplot as plt\n",
" import numpy as np\n",
" import torch\n",
"except (ImportError, ModuleNotFoundError) as exc:\n",
" raise ImportError(\n",
" f\"Required notebook dependencies are not importable in this kernel: {sys.executable}. \"\n",
" \"From the repository root, run `uv sync --dev`, then restart the notebook kernel.\"\n",
" ) from exc\n",
"\n",
"\n",
"def find_repo_root(start: Path) -> Path:\n",
" for candidate in (start, *start.parents):\n",
" if (candidate / \"pyproject.toml\").exists() and (candidate / \"src\" / \"airfrans_frontier\").exists():\n",
" return candidate\n",
" raise RuntimeError(f\"Could not find repository root from {start}\")\n",
"\n",
"\n",
"REPO_ROOT = find_repo_root(Path.cwd())\n",
"SRC_DIR = REPO_ROOT / \"src\"\n",
"if str(SRC_DIR) not in sys.path:\n",
" sys.path.insert(0, str(SRC_DIR))\n",
"\n",
"from airfrans_frontier.models.mlp import PointwiseMLP\n",
"\n",
"print(f\"repo: {REPO_ROOT}\")\n"
]
},
{
"cell_type": "markdown",
"id": "1388d69e",
"metadata": {},
"source": [
"## Select a run\n",
"\n",
"By default this picks the newest directory under `artifacts/runs`. Override `RUN_DIR` manually if you want an older run.\n"
]
},
{
"cell_type": "code",
"execution_count": 2,
"id": "fd3c05a7",
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"selected run: artifacts/runs/20260721T082353Z_mlp_tiny\n",
"available runs:\n",
" artifacts/runs/20260721T082151Z_mlp_tiny\n",
"* artifacts/runs/20260721T082353Z_mlp_tiny\n"
]
}
],
"source": [
"RUNS_DIR = REPO_ROOT / \"artifacts\" / \"runs\"\n",
"run_dirs = sorted([path for path in RUNS_DIR.glob(\"*\") if path.is_dir()], key=lambda path: path.stat().st_mtime)\n",
"if not run_dirs:\n",
" raise FileNotFoundError(f\"No training runs found under {RUNS_DIR}\")\n",
"\n",
"RUN_DIR = run_dirs[-1]\n",
"print(f\"selected run: {RUN_DIR.relative_to(REPO_ROOT)}\")\n",
"print(\"available runs:\")\n",
"for path in run_dirs:\n",
" marker = \"*\" if path == RUN_DIR else \" \"\n",
" print(f\"{marker} {path.relative_to(REPO_ROOT)}\")\n"
]
},
{
"cell_type": "code",
"execution_count": 3,
"id": "e7e15a27",
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"config.toml 438 bytes\n",
"split_manifest.json 153 bytes\n",
"normalization.json 733 bytes\n",
"metrics.jsonl 725 bytes\n",
"final_metrics.json 1,074 bytes\n",
"checkpoint.pt 623,285 bytes\n"
]
}
],
"source": [
"expected_files = [\n",
" \"config.toml\",\n",
" \"split_manifest.json\",\n",
" \"normalization.json\",\n",
" \"metrics.jsonl\",\n",
" \"final_metrics.json\",\n",
" \"checkpoint.pt\",\n",
"]\n",
"missing = [name for name in expected_files if not (RUN_DIR / name).exists()]\n",
"if missing:\n",
" raise FileNotFoundError(f\"Run is missing expected artifacts: {missing}\")\n",
"\n",
"for name in expected_files:\n",
" path = RUN_DIR / name\n",
" print(f\"{name:20} {path.stat().st_size:>10,} bytes\")\n"
]
},
{
"cell_type": "markdown",
"id": "da64dacb",
"metadata": {},
"source": [
"## Final metrics\n",
"\n",
"This is the end-of-run summary written by the training loop. The key checks are: train loss dropped from the initial loss, validation/test losses are finite, and device/GPU fields match the intended hardware.\n"
]
},
{
"cell_type": "code",
"execution_count": 4,
"id": "cb352094",
"metadata": {},
"outputs": [
{
"data": {
"text/plain": [
"{'device': 'cuda:0',\n",
" 'elapsed_seconds': 1.7905638560005173,\n",
" 'gpu_memory_peak_allocated_mb': 17,\n",
" 'gpu_memory_total_mb': 3717,\n",
" 'gpu_name': 'NVIDIA T550 Laptop GPU',\n",
" 'initial_train_loss': 1.0026862990200476,\n",
" 'parameter_count': 50820,\n",
" 'points_per_case': 128,\n",
" 'steps': 500,\n",
" 'test_cases': 1,\n",
" 'test_loss': 0.00022942002189996227,\n",
" 'test_mse_per_channel': {'pressure': 0.000408856492614153,\n",
" 'turbulent_viscosity': 0.00035512470154015644,\n",
" 'velocity_x': 5.778009158681787e-05,\n",
" 'velocity_y': 9.591880185872174e-05},\n",
" 'train_cases': 4,\n",
" 'train_loss': 0.0001373240501380092,\n",
" 'train_mse_per_channel': {'pressure': 0.00023612300953799542,\n",
" 'turbulent_viscosity': 0.00018737132594797248,\n",
" 'velocity_x': 5.6441453127828415e-05,\n",
" 'velocity_y': 6.936041193824046e-05},\n",
" 'val_cases': 1,\n",
" 'val_loss': 0.0007143189445776568,\n",
" 'val_mse_per_channel': {'pressure': 0.0012492708046041544,\n",
" 'turbulent_viscosity': 0.0010587173853462636,\n",
" 'velocity_x': 0.00028163288487499037,\n",
" 'velocity_y': 0.0002676547034852188}}"
]
},
"execution_count": 4,
"metadata": {},
"output_type": "execute_result"
}
],
"source": [
"final_metrics = json.loads((RUN_DIR / \"final_metrics.json\").read_text())\n",
"final_metrics\n"
]
},
{
"cell_type": "code",
"execution_count": 5,
"id": "42010827",
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"initial_train_loss: 1.00269\n",
"train_loss: 0.000137324\n",
"val_loss: 0.000714319\n",
"test_loss: 0.00022942\n",
"loss improvement: 7301.6x\n",
"device: cuda:0\n",
"gpu: NVIDIA T550 Laptop GPU\n",
"parameters: 50,820\n",
"basic metric checks: ok\n"
]
}
],
"source": [
"initial = final_metrics[\"initial_train_loss\"]\n",
"train = final_metrics[\"train_loss\"]\n",
"val = final_metrics.get(\"val_loss\")\n",
"test = final_metrics.get(\"test_loss\")\n",
"\n",
"print(f\"initial_train_loss: {initial:.6g}\")\n",
"print(f\"train_loss: {train:.6g}\")\n",
"print(f\"val_loss: {val:.6g}\" if val is not None else \"val_loss: None\")\n",
"print(f\"test_loss: {test:.6g}\" if test is not None else \"test_loss: None\")\n",
"print(f\"loss improvement: {initial / train:.1f}x\")\n",
"print(f\"device: {final_metrics.get('device')}\")\n",
"print(f\"gpu: {final_metrics.get('gpu_name')}\")\n",
"print(f\"parameters: {final_metrics.get('parameter_count'):,}\")\n",
"\n",
"assert np.isfinite(train), \"train loss is not finite\"\n",
"assert train < initial, \"train loss did not improve\"\n",
"if val is not None:\n",
" assert np.isfinite(val), \"val loss is not finite\"\n",
"if test is not None:\n",
" assert np.isfinite(test), \"test loss is not finite\"\n",
"print(\"basic metric checks: ok\")\n"
]
},
{
"cell_type": "markdown",
"id": "e4cf0126",
"metadata": {},
"source": [
"## Training curves\n",
"\n",
"`metrics.jsonl` has one JSON row per log interval. A sanity run should show train loss falling; validation loss should be finite and usually trend down on this tiny baseline.\n"
]
},
{
"cell_type": "code",
"execution_count": 6,
"id": "b4d36a9d",
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"{'elapsed_seconds': 0.0, 'step': 0, 'train_loss': 1.0026862990200476, 'val_loss': 0.8999216645838499}\n",
"{'elapsed_seconds': 0.6329330740009027, 'step': 100, 'train_loss': 0.005394324883408844, 'val_loss': 0.008841235591800927}\n",
"{'elapsed_seconds': 0.9306512370003475, 'step': 200, 'train_loss': 0.0008050906900035648, 'val_loss': 0.002610322449175867}\n",
"{'elapsed_seconds': 1.130966939001155, 'step': 300, 'train_loss': 0.00029600711922912065, 'val_loss': 0.0015054571447265593}\n",
"{'elapsed_seconds': 1.4691060960012692, 'step': 400, 'train_loss': 0.00019286412112246845, 'val_loss': 0.0012647405838277085}\n",
"{'elapsed_seconds': 1.7875807079999504, 'step': 500, 'train_loss': 0.0001373240501380092, 'val_loss': 0.0007143189445776568}\n"
]
}
],
"source": [
"metric_rows = [json.loads(line) for line in (RUN_DIR / \"metrics.jsonl\").read_text().splitlines() if line.strip()]\n",
"steps = np.array([row[\"step\"] for row in metric_rows])\n",
"train_loss = np.array([row[\"train_loss\"] for row in metric_rows], dtype=float)\n",
"val_loss = np.array([row[\"val_loss\"] if row[\"val_loss\"] is not None else np.nan for row in metric_rows], dtype=float)\n",
"\n",
"for row in metric_rows:\n",
" print(row)\n"
]
},
{
"cell_type": "code",
"execution_count": 7,
"id": "c678379b",
"metadata": {},
"outputs": [
{
"data": {
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAxYAAAG4CAYAAADYN3EQAAAAOnRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjExLjEsIGh0dHBzOi8vbWF0cGxvdGxpYi5vcmcvctoD+AAAAAlwSFlzAAAPYQAAD2EBqD+naQAAmS1JREFUeJzs3Xd8FHX+x/HXbHrvIQkEQoDQi4pIE6yACtjQ0xO7cij27tnLCSr+rKh3cvZyIuAJqFhOBUEpNqQHCL2F9N525/fHkpCQBJKwyWyy7+fjsQ+SmcnMZ/P9MOyH+RbDNE0TERERERGRY2CzOgAREREREWn9VFiIiIiIiMgxU2EhIiIiIiLHTIWFiIiIiIgcMxUWIiIiIiJyzFRYiIiIiIjIMVNhISIiIiIix0yFhYiIiIiIHDMVFiIiIiIicsxUWIiIiLRyGRkZGIbBCy+84JHXFxH3oMJCRNxaWVkZH3/8Meeccw4JCQmEhoZy/PHHM2PGDMrLy+v8mQ8++IABAwbg7+9PXFwckydPJjs7+5jPm5OTw913303Xrl0JCAigb9++vPrqq1RUVDQphgULFmAYRr2vSy+9tMnx7tq1i+eff54hQ4Zgs9no06dPrWPOOOOMI16/8nXKKac06n0BZGZmcu+995KSkkJgYCBxcXGcddZZLFq06JjaYfny5Vx55ZV06dKFgIAAunTpwt/+9je2bt1a69iuXbvW+X7OOOOMJp+3ob+ztmjXrl0YhsErr7xidSgi4q5MERE39t5775leXl7m3XffbW7bts3Mzc013333XTMgIMA877zzah3/xhtvmID50ksvmYWFhebq1avNnj17mgMHDjTLy8ubfN79+/ebKSkp5umnn27++uuvZmlpqZmWlmbefPPN5pIlS5oUQ33uvvtuEzA//PDDJsc7duxY87bbbjN/+uknMyUlxezdu/dRr/vRRx+ZgPnRRx/Vub+h78tut5v9+/c3IyMjzYULF5qFhYVmWlqaed5555k2m8387rvvmvS+KioqzFNOOcWcNWuWuX37drOgoMBctGiRmZKSYkZFRZl79uypcXyXLl3Mc88996jvu7Hnrc8jjzxiAuaJJ57YoONd6cCBAyZgPv/88812jZ07d5qA+fLLLzfbNUSkdVNhISJubd68eTU+iFZ68MEHTcBcuXJl1bbi4mIzIiLCvPDCC2sc++OPP5qAOXPmzCad1zRN88ILLzS7detmFhcXHzHexsRQl/LycjMuLs6MiooyS0pKmhxvdd27dz/mwqIx7+vXX381AfOBBx6ocWx6eroJmFdffbVL3lf1c9T1gbehhUVjz1uXzz77zDQMw2zXrp25c+fOJl+zqVRYiIg7UFcoEXFr48aN49RTT621vWvXrgA1uqp8//33ZGdnc/7559c4dvjw4cTGxvLJJ5806bw7duxg7ty5TJo0CX9//yPG25gY6vL555+zb98+rrzySvz8/JoUb3NozPsKDQ094rnCwsKqvnbF+/Lx8QEgKCjoqMc2RkPPu379eiZOnIi3tzezZ8+mQ4cOjbrO5s2bMQyDmTNn8v7775OSkkJQUBCjR49m165dALz33nt0794df39/hgwZwtq1axt13vfee49u3brh7+/Pcccdx+eff96oGJcsWUJiYiIAN998c1WXr5tuugmoe4xF9evPmTOHnj174ufnR79+/fjqq6+qjissLCQ8PJyJEyfWum5JSQnR0dFcfPHFjYpXRKyhwkJEWqUFCxYA0L1796ptf/75Z61tlXr06MHq1aubdN7FixdjmiaBgYGMGzeO0NBQAgMDGT58eI0PSK6I4c033wTg+uuvP2qs9cXbHBrzvrp27cqkSZN47bXX+OqrrygqKmLbtm1MmjSJhIQEbr311qNeryHvq7S0lBUrVnDvvfdy3HHH8Ze//KXWMd9++y2hoaH4+/vTs2dPHn30UUpKSo547Yact1Jubi7nnXce+fn5vPTSSwwfPvyo760+X3zxBatWreLHH3/kzz//ZN++fUyYMIGPP/6Y3377jR9++IH169dTUlLCxRdfjGmaDTrvggULWLZsGT/88AObN2/mxBNPZPz48SxcuLDBsQ0fPpydO3cC8PLLL2M6ezw0aLzF119/zeLFi/nmm2/Ytm0bnTt35vzzz2f//v2As3C76qqrmD17NgcOHKjxsx9//DGZmZlce+21DY5VRCxk7QMTEZHGmz9/vgmYY8aMqbG9cmzCpk2bav3M+PHjTV9f3yadd+rUqSZgenl5mffee6+5f/9+c/v27eZFF11kGoZhzp8/3yUx7Nu3z/T29jZHjBhxxDiPFu/hXNEVqrHvq6yszJw0aZIJVL06d+5s/vbbb0eN42jv6/vvv69x3qFDh5q7du2qddxVV11lfvHFF2ZWVpa5a9cu88UXXzQDAwPN4cOH1znWpaHnrWS3282zzz7bBMxJkyYd9X3VZ9OmTSZgDh48uMb2WbNmmYB5yimn1Ng+Z84cEzCXLl1ata2urlCV5x0wYECNn3c4HGavXr3Mfv36NSrOI3WFOtL1Bw0aVOPY7du3m4D53HPPVW3buHGjaRiGOW3atBrHnnTSSWZiYqJpt9sbFauIWENPLESkVVm1ahUTJ06kffv2/Pvf/67zmPpm5TnSbD1HOq/D4QDgxBNPZNq0acTGxtKxY0feeecd2rVrx4MPPuiSGN555x0qKioa9LSiIb+H5tCQ91VRUcHo0aOZN28eCxYsID8/ny1btjB48GBGjBjBkiVL6j1/Q97XKaecgmmaFBQU8MMPP1BYWMhJJ53Etm3bahz31ltvcdZZZxEREUH79u255ZZbmDp1KkuWLOE///lPk89b6cEHH+SLL75g6NChvPzyy/W+p4Y666yzanzfo0cPAE4++eQa23v27AlAWlpag847bty4Gt8bhsH48eP5888/ycjIaGq4DXbOOefU+L5jx46EhITUiD8lJYXTTz+df/7zn1V/337//XeWL1/O1Vdfjc2mjysirYH+popIq7FhwwZGjRpFQEAA3333HQkJCTX2R0VFAdQ5/WlOTg6RkZFNOm90dDQAI0eOrLE9ICCAQYMGsWrVKsrKyo4pBnB2g4qMjGTChAn1HtOQeJtDY97XJ598wvfff8+0adM455xzCA4OJjk5mbfffpvQ0FDuuuuuOq/R2PcVFBTEyJEjmTt3Lrt372bq1KlHfR/nnXcewBGLm4ac95NPPmHq1KkkJCQwZ84cfH19j3rto4mPj6/xfUhIyBG35+TkNOi87dq1q3dbZmZmY8NstMPjB+c4nMPjv/HGG9m6dWtVF61XX30VwzC4+uqrmz1GEXENFRYi0ips3ryZ0047DcMw+O6770hJSal1TN++fQHYuHFjrX0bNmygX79+TTpv//79643LNM0aaxc0JQaApUuXsnHjRi6//PIjDhBvSLzNoTHva8OGDTV+ppKvry8pKSlV+6s7lveVnJxMUFBQg/8H/1jPu3r1aq6++mr8/Pz49NNPiYuLc8n1mvKUqyEqxzLUta2yYGxODY1//PjxJCYm8uqrr5Kbm8uHH37I6aefTlJSUvMGKCIuo8JCRNze9u3bOf3007Hb7Xz33XdVXUEOd9pppxEeHs6nn35aY/uSJUtIT0+v9SSgoecdNGgQSUlJLF68uMb2kpISVq5cyQknnFA1g1BjY6hU2e1n0qRJ9fwWGh5vc2jM++rUqRMAa9asqXFsWVkZqampVfsrHev7WrduHYWFhQ36uc8++wyAYcOGNem8WVlZnHfeeRQWFvLaa68xaNCgRsVqhcqB8JVM02T+/Pn069ev6mlcQ1TOjlVaWurS+Cp5eXkxadIkvvzySx599FGKioo0aFuktbF2iIeIyJHt2bPH7NKlixkbG2uuXbv2qMe//vrrdS7idvzxx5tlZWVNPu+8efNMm81m/v3vfzfT09PNHTt2mJdcconp7e1tfvvtt02KoVJ+fr4ZHBxsDhs2zGW/h+pcMXjbNBv+vvLz882kpCQzPj7e/PLLL82CggIzLS3NvPTSS2st/NeY9/XBBx+Yt956q/nLL7+Yubm5ZmZmprl
"text/plain": [
"<Figure size 800x450 with 1 Axes>"
]
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"fig, ax = plt.subplots(figsize=(8, 4.5))\n",
"ax.plot(steps, train_loss, marker=\"o\", label=\"train\")\n",
"if np.isfinite(val_loss).any():\n",
" ax.plot(steps, val_loss, marker=\"o\", label=\"val\")\n",
"ax.set_yscale(\"log\")\n",
"ax.set_xlabel(\"optimization step\")\n",
"ax.set_ylabel(\"normalized MSE\")\n",
"ax.set_title(RUN_DIR.name)\n",
"ax.grid(True, which=\"both\", alpha=0.25)\n",
"ax.legend()\n",
"fig.tight_layout()\n"
]
},
{
"cell_type": "markdown",
"id": "804c8844",
"metadata": {},
"source": [
"## Split and normalization\n",
"\n",
"The split manifest tells you exactly which `.npz` cases produced train/validation/test metrics. Normalization is computed from the train split and stored so the checkpoint can be used later.\n"
]
},
{
"cell_type": "code",
"execution_count": 8,
"id": "d993f11f",
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"[run]\n",
"name = \"mlp_tiny\"\n",
"seed = 0\n",
"artifact_dir = \"artifacts/runs\"\n",
"\n",
"[data]\n",
"root = \"data/processed/minimal\"\n",
"train_cases = 4\n",
"val_cases = 1\n",
"test_cases = 1\n",
"points_per_case = 128\n",
"batch_size = 128\n",
"\n",
"[model]\n",
"type = \"mlp\"\n",
"hidden_width = 128\n",
"depth = 4\n",
"activation = \"gelu\"\n",
"\n",
"[optim]\n",
"lr = 0.001\n",
"weight_decay = 0.0\n",
"steps = 500\n",
"log_interval = 100\n",
"\n",
"[device]\n",
"type = \"cuda\"\n",
"allow_cpu_fallback = false\n",
"benchmark_kernels = true\n",
"\n",
"[loss]\n",
"type = \"normalized_mse\"\n",
"\n",
"split manifest:\n",
"{\n",
" \"test_ids\": [\n",
" \"case_03\"\n",
" ],\n",
" \"train_ids\": [\n",
" \"case_04\",\n",
" \"case_02\",\n",
" \"case_01\",\n",
" \"case_00\"\n",
" ],\n",
" \"val_ids\": [\n",
" \"case_05\"\n",
" ]\n",
"}\n",
"normalization target names: ['velocity_x', 'velocity_y', 'pressure', 'turbulent_viscosity']\n"
]
}
],
"source": [
"split_manifest = json.loads((RUN_DIR / \"split_manifest.json\").read_text())\n",
"normalization = json.loads((RUN_DIR / \"normalization.json\").read_text())\n",
"config_text = (RUN_DIR / \"config.toml\").read_text()\n",
"\n",
"print(config_text)\n",
"print(\"split manifest:\")\n",
"print(json.dumps(split_manifest, indent=2))\n",
"print(\"normalization target names:\", normalization[\"target_names\"])\n"
]
},
{
"cell_type": "code",
"execution_count": 9,
"id": "b16e05bc",
"metadata": {},
"outputs": [
{
"data": {
"image/png": "iVBORw0KGgoAAAANSUhEUgAABKYAAAGGCAYAAABBiol3AAAAOnRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjExLjEsIGh0dHBzOi8vbWF0cGxvdGxpYi5vcmcvctoD+AAAAAlwSFlzAAAPYQAAD2EBqD+naQAAbuBJREFUeJzt3Xl4TOf///HXZJEgImLfSm1FaRuC2oqIpZaglNLalypK0Wp1+aDaUlWU2qmltVS11E7t+761Su1LLRFrQgjJ3L8/fM2v06BBJmciz8d1fa7rc+5zzj3v6T2Jk9fc5z42Y4wRAAAAAAAAkMQ8rC4AAAAAAAAAKRPBFAAAAAAAACxBMAUAAAAAAABLEEwBAAAAAADAEgRTAAAAAAAAsATBFAAAAAAAACxBMAUAAAAAAABLEEwBAAAAAADAEgRTAAAAAAAAsATBFIB4/vrrL1WvXl3p06eXzWbTDz/8YHVJAAAAgMM777wjX19fq8sAkAgIpgDE07JlS0VHR+vYsWMyxuiNN95I9NeYPHmybDabDhw4kOh9AwAAJAdjxoyRzWbT8ePHrS4lwerUqaNixYq5fZ8Akg+CKQBOYmJitHXrVtWuXVuBgYFWlwMAAAAAeIIRTAFwcuHCBRljlDp1aqtLAQAAAAA84QimADi88847ypUrlySpe/fustls8vPzc+w/d+6cOnbsqFy5cilVqlR66qmn1KtXL928edNxzOrVq2Wz2Rz/S5MmjUqUKKGxY8c6junbt69at24tSSpSpIjj2Llz50qSPv74Y9lsNsXGxjrVd7fvJUuWONruToE/fPiwPvroI+XIkUM2m81R06pVq1StWjWlT59evr6+Cg4O1i+//PKf/y0qVKigChUq6MiRI6pevbrSpk2rAgUKOM49ePCgatasKT8/P+XKlUujR4++Zz8Jef127do5/ht4eHgoY8aMqlu3rnbt2nXPmk6dOqU6deoobdq0ypo1q3r16qW4uDinY7ds2aKXX35ZWbJkUbp06RQcHKwJEybEOw4AAFjjgw8+0FtvvSVJevrppx3XAnevc5YsWRLvmio4OFiTJk1y6mfYsGGy2Ww6deqUevXqpezZs8tmszn2Dx8+XAUKFJCvr69KliypVatW3Xd9pnXr1qlmzZoKCAiQr6+vSpQooR9//NGxv1ixYlq4cKH27dvnqMvLy+uB73PHjh2qU6eOsmbNKj8/P5UoUUJjxoxxXOclpM+RI0eqYMGCTu/hUV4LgJsyAPAPp06dMpLM0KFDndrPnj1rcufObUqUKGE2btxooqKizNq1a02+fPlMzZo179vf+fPnzTfffGM8PT3N999/72ifNGmSkWT2798f75yPPvrISDK3b992al+1apWRZBYvXuxoGz16tJFkXnvtNTNmzBhz4cIFM3nyZHPz5k0zffp04+HhYbp27WqOHz9uLl265Khl6tSpD/zvUL58efPCCy+Y+vXrm507d5rLly+b999/33h5eZm1a9eaGjVqmO3bt5srV6446t28ebNTH4/y+rdu3TL79u0z9evXN5kzZzbh4eFONQUFBZlXXnnFbN682URGRpqJEycam81mRowY4fTf3N/f37Rs2dKcPHnSREdHm127dpkOHTqYbdu2PfB9AwCApHP3OubYsWMPPM5ut5tz586Zr776ynh4eJhZs2Y59g0dOtRIMs2aNTMTJ040Fy9eNGPHjjXGGDNgwADj4eFhvvrqK3PhwgVz6NAh07BhQ1OzZk3j4+Pj9Bo///yz8fT0NB07djRHjx41ly9fNqNGjTLe3t5m/PjxjuNq165tnn322QS9v8uXL5sMGTKYpk2bmhMnTpgbN26YPXv2mE6dOpn169cnqM9BgwYZDw8PM3DgQBMREWEOHjxowsLC4r2HhL4WAPdDMAXAyf2Cqfbt25s0adKYv//+26l9+fLlRpL57bffHthv/fr1TYUKFRzbiR1Mde/e3enY6OhokzFjRlOrVq14/bdq1cpkz57dxMXF3bfe8uXLG5vNZv78809HW0xMjAkICDDp0qUzv//+u6P91q1bJjAw0LRv3z7RXj8yMtJ4eHiYMWPGxKvpjz/+cDq2SpUqpnjx4o7tJUuWGElmy5Yt9+0fAABYL6HB1D/VrFnThIaGOrbvBlO9e/d2Oi4qKsr4+fmZpk2bxmvPkCGDU6gTExNjsmbNaqpWrRrv9Tp27GgyZcrkuC57mGBq9erVRpJZs2bNA4+7X5/Xrl0zfn5+pkmTJk7tV65cMenTp3d6Dwl9LQDuh1v5ACTI/PnzVa5cOeXMmdOpvVKlSvL29taaNWskScYYjRgxQiVLlpSfn5/TbXqHDx92WX1hYWFO2xs3btTFixfVuHHjeMeGhobq7NmzOnTo0AP7zJs3r4oUKeLYTpUqlfLly6eAgACnJ8d4e3urQIECOnr06CO9/t1bJPPmzatUqVLJZrPJ399fdrs93n+zvHnz6tlnn3VqK1asmNNrFy5cWKlSpVLXrl01b948RUVFPfB9AgAA92O32zVkyBCVKFFCadOmdbrV717XVP++Ftq2bZuuXbumunXrOrX7+fmpSpUq8Y4NDw/Xq6++Gq/f0NBQXbhwQfv27Xvo91CoUCGlTp1a3bt315w5c3T16tWHOn/r1q26du1avPeWPn36eO/hcV8LgHUIpgAkSHh4uFasWCEvLy95enrK09NTHh4e8vb21u3bt3Xx4kVJUr9+/dS9e3e1a9dOhw4dUmxsrIwxeuONN3T79u3HqsEYc999/w7Mzp07J0lq06aNo2YPDw95eHjojTfekCRHzfeTPXv2eG3p0qW7b/uVK1ce+vVjY2NVpUoVrVq1SlOmTNGFCxdkt9sVGxsrT0/PeP/N7vXa/v7+un79umP9hDx58mjhwoXy9vbWK6+8ooCAAJUpU0bjxo2T3W5/4HsGAADu4cMPP9T777+vTp066ciRI45rqkaNGt3zmurf10J3r3OyZMkS79h/t929bunUqVO8a71GjRo59fcwsmfPrkWLFsnPz0+NGzdWYGCgSpUqpVGjRiVo3cu7r5k1a9Z4+/7d9rivBcA6BFMAEiRjxoyqV6+eYmNjFRcXp7i4ONntdpk7twRr1KhRkqSpU6eqVq1aeuutt5Q9e3Z5enpKko4dO5bg10qfPr0kxZvpc/r06fue4+3t7bSdKVMmSdKsWbMcNdvtdqeay5Ur98A6/rlwaELaH+X1t27dqgMHDqh///6qVKmS/P39HQuY3usiKiGvLd35dnPdunW6fPmyFi9erPz58+vNN9/U8OHDE3Q+AACw1tSpU9WgQQO1a9dO2bJl+89rqn9fC2XMmFGSdP78+XjH/rvt7nXL1KlT73utFxIS8kjvo3LlylqzZo0uX76spUuXqkiRIurcubO++uqr/zz37nsIDw+Pt+9ebY/zWgCsQzAFIEHq1q2r1atX3/Pi5t98fHyctg8fPqzNmzc7taVNm1aSFBMTE+/8/PnzS5L++OMPp/YFCxYkuN4KFSooICDA6UkySelhX//f/82mTp2aKHWkS5dO1atX1/Tp05UpUyatXbs2UfoFAACP70HXQ1L864N9+/Zp586dCeq7VKlS8vPz08KFC53ar1+/rtWrVzu1lSlTRpkzZ9asWbMSVPP96n0QPz8/hYaGaurUqcqVK5fTNcn9+rz7HubPn+/UHhkZGe89JPS1ALgfgikACfL5558rffr0evnll7Vy5UpFRkbq/PnzWrVqlZo2bar169dLurO+wfz58zVv3jxdv35dmzdv1htvvKGqVas69Xd3naRFixbp1q1bTvtq166tHDly6P3339exY8cUERGhzz//3PFNYUKkTZtWo0aN0s8//6zOnTvrr7/+0o0bN3TkyBFNmTJFtWvXfsz/Ionz+kFBQcqdO7c+++wzHTx4UJcvX9a4ceO0d+/eeBejCTV9+nR17NhRmzZt0tWrVxUZGalJkybp4sWL8dZjAAAA1rm7ZuXChQvjXQ+FhYXp559/1qJFi3T9+nVt2LBBbdu2TfC/5X5+fvroo480c+ZMDRkyRJcuXdKRI0fUunVrlS5d2ulYX19
"text/plain": [
"<Figure size 1200x400 with 2 Axes>"
]
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"feature_mean = np.array(normalization[\"feature_mean\"], dtype=np.float32)\n",
"feature_std = np.array(normalization[\"feature_std\"], dtype=np.float32)\n",
"target_mean = np.array(normalization[\"target_mean\"], dtype=np.float32)\n",
"target_std = np.array(normalization[\"target_std\"], dtype=np.float32)\n",
"\n",
"fig, axes = plt.subplots(1, 2, figsize=(12, 4))\n",
"axes[0].bar(normalization[\"feature_names\"], feature_mean)\n",
"axes[0].set_title(\"feature means\")\n",
"axes[0].tick_params(axis=\"x\", rotation=45)\n",
"axes[0].grid(axis=\"y\", alpha=0.25)\n",
"\n",
"axes[1].bar(normalization[\"target_names\"], target_std)\n",
"axes[1].set_title(\"target stds\")\n",
"axes[1].tick_params(axis=\"x\", rotation=45)\n",
"axes[1].grid(axis=\"y\", alpha=0.25)\n",
"fig.tight_layout()\n"
]
},
{
"cell_type": "markdown",
"id": "32f03cf3",
"metadata": {},
"source": [
"## What the input features mean\n",
"\n",
"Each row is one sampled point from one case. The MLP sees five scalar inputs for that point:\n",
"\n",
"| feature | meaning | varies per |\n",
"|---|---|---|\n",
"| `re_norm` | normalized Reynolds-number-like case condition; constant across all points in one case | case |\n",
"| `aoa_norm` | normalized angle-of-attack-like case condition; constant across all points in one case | case |\n",
"| `x` | point x-coordinate in the local 2D domain | point |\n",
"| `y` | point y-coordinate in the local 2D domain | point |\n",
"| `sdf` | signed-distance-like geometry feature; here computed as radial distance from a radius-0.5 body proxy, so negative means inside/near the proxy and positive means farther away | point |\n",
"\n",
"For the current `data/processed/minimal` sanity dataset, these are deliberately compact synthetic features from the training-framework spec, not the full AirfRANS graph feature set. The target formulas used by the synthetic sanity data are simple functions of these columns, which is why the tiny MLP should learn them quickly."
]
},
{
"cell_type": "code",
"execution_count": 10,
"id": "c9b3b2fe",
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"re_norm - case-level normalized Reynolds-number-like condition\n",
"aoa_norm - case-level normalized angle-of-attack-like condition\n",
"x - sampled point x-coordinate\n",
"y - sampled point y-coordinate\n",
"sdf - signed-distance-like geometry signal; negative near/inside radius-0.5 proxy\n"
]
}
],
"source": [
"feature_descriptions = {\n",
" \"re_norm\": \"case-level normalized Reynolds-number-like condition\",\n",
" \"aoa_norm\": \"case-level normalized angle-of-attack-like condition\",\n",
" \"x\": \"sampled point x-coordinate\",\n",
" \"y\": \"sampled point y-coordinate\",\n",
" \"sdf\": \"signed-distance-like geometry signal; negative near/inside radius-0.5 proxy\",\n",
"}\n",
"\n",
"for name in normalization[\"feature_names\"]:\n",
" print(f\"{name:8} - {feature_descriptions.get(name, 'no description recorded')}\")\n"
]
},
{
"cell_type": "code",
"execution_count": 11,
"id": "2a31f98e",
"metadata": {},
"outputs": [
{
"data": {
"text/plain": [
"[{'split': 'test',\n",
" 'case_id': 'case_03',\n",
" 'points': 256,\n",
" 're_norm_min': 0.10000000149011612,\n",
" 're_norm_mean': 0.10000001639127731,\n",
" 're_norm_max': 0.10000000149011612,\n",
" 'aoa_norm_min': 0.10000000149011612,\n",
" 'aoa_norm_mean': 0.10000001639127731,\n",
" 'aoa_norm_max': 0.10000000149011612,\n",
" 'x_min': -0.992418110370636,\n",
" 'x_mean': -0.028246253728866577,\n",
" 'x_max': 0.9915499091148376,\n",
" 'y_min': -0.999594509601593,\n",
" 'y_mean': 0.00812380388379097,\n",
" 'y_max': 0.9907281994819641,\n",
" 'sdf_min': -0.483142614364624,\n",
" 'sdf_mean': 0.24713799357414246,\n",
" 'sdf_max': 0.855535626411438},\n",
" {'split': 'train',\n",
" 'case_id': 'case_04',\n",
" 'points': 256,\n",
" 're_norm_min': 0.30000001192092896,\n",
" 're_norm_mean': 0.30000001192092896,\n",
" 're_norm_max': 0.30000001192092896,\n",
" 'aoa_norm_min': 0.20000000298023224,\n",
" 'aoa_norm_mean': 0.20000003278255463,\n",
" 'aoa_norm_max': 0.20000000298023224,\n",
" 'x_min': -0.9973810315132141,\n",
" 'x_mean': 0.016297580674290657,\n",
" 'x_max': 0.9969028830528259,\n",
" 'y_min': -0.9971365332603455,\n",
" 'y_mean': -0.03903472423553467,\n",
" 'y_max': 0.988767683506012,\n",
" 'sdf_min': -0.4517866373062134,\n",
" 'sdf_mean': 0.29825884103775024,\n",
" 'sdf_max': 0.8830513954162598},\n",
" {'split': 'train',\n",
" 'case_id': 'case_02',\n",
" 'points': 256,\n",
" 're_norm_min': -0.10000000149011612,\n",
" 're_norm_mean': -0.10000001639127731,\n",
" 're_norm_max': -0.10000000149011612,\n",
" 'aoa_norm_min': 0.0,\n",
" 'aoa_norm_mean': 0.0,\n",
" 'aoa_norm_max': 0.0,\n",
" 'x_min': -0.9979289770126343,\n",
" 'x_mean': 0.03891447186470032,\n",
" 'x_max': 0.9994633793830872,\n",
" 'y_min': -0.9955872893333435,\n",
" 'y_mean': -0.018903596326708794,\n",
" 'y_max': 0.9769426584243774,\n",
" 'sdf_min': -0.4469470679759979,\n",
" 'sdf_mean': 0.2548476755619049,\n",
" 'sdf_max': 0.8520128726959229},\n",
" {'split': 'train',\n",
" 'case_id': 'case_01',\n",
" 'points': 256,\n",
" 're_norm_min': -0.30000001192092896,\n",
" 're_norm_mean': -0.30000001192092896,\n",
" 're_norm_max': -0.30000001192092896,\n",
" 'aoa_norm_min': -0.10000000149011612,\n",
" 'aoa_norm_mean': -0.10000001639127731,\n",
" 'aoa_norm_max': -0.10000000149011612,\n",
" 'x_min': -0.9986386895179749,\n",
" 'x_mean': -0.039472322911024094,\n",
" 'x_max': 0.991798460483551,\n",
" 'y_min': -0.9973899126052856,\n",
" 'y_mean': 0.029526807367801666,\n",
" 'y_max': 0.998464822769165,\n",
" 'sdf_min': -0.4342630207538605,\n",
" 'sdf_mean': 0.2738206684589386,\n",
" 'sdf_max': 0.9004930257797241},\n",
" {'split': 'train',\n",
" 'case_id': 'case_00',\n",
" 'points': 256,\n",
" 're_norm_min': -0.5,\n",
" 're_norm_mean': -0.5,\n",
" 're_norm_max': -0.5,\n",
" 'aoa_norm_min': -0.20000000298023224,\n",
" 'aoa_norm_mean': -0.20000003278255463,\n",
" 'aoa_norm_max': -0.20000000298023224,\n",
" 'x_min': -0.9898697733879089,\n",
" 'x_mean': -0.004540023393929005,\n",
" 'x_max': 0.9964226484298706,\n",
" 'y_min': -0.9919538497924805,\n",
" 'y_mean': -0.0564127042889595,\n",
" 'y_max': 0.9963796734809875,\n",
" 'sdf_min': -0.40864098072052,\n",
" 'sdf_mean': 0.25872567296028137,\n",
" 'sdf_max': 0.8634771108627319},\n",
" {'split': 'val',\n",
" 'case_id': 'case_05',\n",
" 'points': 256,\n",
" 're_norm_min': 0.5,\n",
" 're_norm_mean': 0.5,\n",
" 're_norm_max': 0.5,\n",
" 'aoa_norm_min': 0.30000001192092896,\n",
" 'aoa_norm_mean': 0.30000001192092896,\n",
" 'aoa_norm_max': 0.30000001192092896,\n",
" 'x_min': -0.986579179763794,\n",
" 'x_mean': -0.012095458805561066,\n",
" 'x_max': 0.9900404810905457,\n",
" 'y_min': -0.9959282875061035,\n",
" 'y_mean': 0.022574514150619507,\n",
" 'y_max': 0.9964964985847473,\n",
" 'sdf_min': -0.4551549255847931,\n",
" 'sdf_mean': 0.25610190629959106,\n",
" 'sdf_max': 0.8619978427886963}]"
]
},
"execution_count": 11,
"metadata": {},
"output_type": "execute_result"
}
],
"source": [
"import tomllib\n",
"\n",
"config = tomllib.loads(config_text)\n",
"DATA_ROOT = REPO_ROOT / config[\"data\"][\"root\"]\n",
"\n",
"\n",
"def load_case(case_id: str) -> tuple[np.ndarray, np.ndarray]:\n",
" data = np.load(DATA_ROOT / f\"{case_id}.npz\")\n",
" return data[\"features\"].astype(np.float32), data[\"targets\"].astype(np.float32)\n",
"\n",
"\n",
"feature_rows = []\n",
"for split_name, ids in split_manifest.items():\n",
" for case_id in ids:\n",
" features, _ = load_case(case_id)\n",
" row = {\"split\": split_name.replace(\"_ids\", \"\"), \"case_id\": case_id, \"points\": len(features)}\n",
" for index, name in enumerate(normalization[\"feature_names\"]):\n",
" values = features[:, index]\n",
" row[f\"{name}_min\"] = float(values.min())\n",
" row[f\"{name}_mean\"] = float(values.mean())\n",
" row[f\"{name}_max\"] = float(values.max())\n",
" feature_rows.append(row)\n",
"\n",
"feature_rows\n"
]
},
{
"cell_type": "code",
"execution_count": 12,
"id": "ccab3feb",
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"case-level conditions for case_05:\n",
"re_norm = 0.500\n",
"aoa_norm = 0.300\n"
]
},
{
"data": {
"image/png": "iVBORw0KGgoAAAANSUhEUgAABEMAAAGkCAYAAADaJet/AAAAOnRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjExLjEsIGh0dHBzOi8vbWF0cGxvdGxpYi5vcmcvctoD+AAAAAlwSFlzAAAPYQAAD2EBqD+naQABAABJREFUeJzs3Xd8G0X6P/DP7K6qq9zjFpfEcXpIJ5AeSIFQQ8nlgECoRwnlOODo5S7HHfDjvrSDC4TejxZa6OkJJKT3YsdOcS9yUd2d3x+K5ciSbEmWLcl+3q+XX4l2d2ZnR5KtfTTzDOOccxBCCCGEEEIIIYT0EkKoG0AIIYQQQgghhBDSnSgYQgghhBBCCCGEkF6FgiGEEEIIIYQQQgjpVSgYQgghhBBCCCGEkF6FgiGEEEIIIYQQQgjpVSgYQgghhBBCCCGEkF6FgiGEEEIIIYQQQgjpVSgYQgghhBBCCCGEkF6FgiGEEEIIIYQQQgjpVSgYQkgnffnll2CMYcOGDUGrc+vWrWCM4f333w9anb7qiusJJ+F6fX/+858hSVKHx82aNQsjRozo+gZ1grdr+fjjjzFkyBBoNBowxtDY2BiC1hFCSHBce+21SEpKctv+xhtvoLCwEGq12qff677UO2vWLIwePTrgtvrL0/nGjx+PKVOmdFsbvLWjO3X2uQyVUPdbZ3h7XxHSFSgYQgAAdrsdTzzxBPLy8qDRaJCfn4+//e1vsNvtLsdVVVWBMebx54EHHghR60k4+fjjj8EYw6ZNm0LdFBJGSkpKsGDBAsybNw8NDQ3gnCM6OjrUzSKEkKDavXs3rrnmGtxwww1oampy+xzV3aZMmYLx48eHtA0dCdc2httzGU7C9TkjxF+RE+IkXerGG2/Exx9/jPfffx9TpkzBypUrcdlll6G0tBT/+c9/3I5fsmQJ7r333hC0lHS1c889F5zzUDeD9DDr16+H1WrFZZddBrVaHermEEJIl1i9ejUURcHll18OlUoVtHq//fbboNUVjufzJpTt6KrnkhASPmhkCMG2bdvw6quv4oEHHsCsWbOg1Woxc+ZMPPDAA3jllVewc+fOUDeREBLhqqqqAAA6nS7ELSGEkK5Dv+t6DnouCen5KBgSRFVVVVi8eLFzqklBQQEefvhhNDU1AQA2bdrkMq1Ep9Nh2LBh+L//+z+XepqamvDnP/8ZeXl50Ol0yMvLww033IBjx465ne/WW29F3759oVarkZmZidtvv915Pl998sknAIALL7zQZftFF10Ezjk+/vhjf7vCKS0tDYWFhR0e58s1+9p/b7/9Nhhj2L59Ox566CH06dMH8fHxWLRoESwWCxRFwYMPPoj09HTo9XpcfPHFqK2t9VrHgw8+iLS0NOj1epx99tk+B4d8fX62bt2K6dOnQ6/Xo0+fPnj44YehKIpP5/C3nQcOHMBll12GpKQkaDQaFBYW4sknn4Qsy85jPOXUaDnPjh07sGTJEmRkZECn02HatGnYu3ev87innnoKl1xyCQBgzJgxzufq7bffBuD7a9tbf7b3/vL1+rzxpeyp/f3YY48hKysLgiA4PzCtXbsWs2fPhsFggFarxYgRI/Duu++6nes///kP+vfvD61Wi9NOOw0//vhjh+1rq6SkBHPmzEF0dDRSU1Nx++23w2QyOfffcccd0Gg0qKiocCt7zz33QJIknDhxwmv9W7Zswdy5c5Gamoro6GiMHDkSL774otswYV+uJScnB7fccgsAIDc3F4wx/PGPf/T7mgkhJNjWr1+PmTNnIjk5GTExMRgzZgyWLVvm9nf46aefdv7tGjduHNatW+dWV3x8vHPKsMFgAGPM+bvPG1/qBTzngOio7Tk5OVi5ciU2btzo/HscHx/vLN+SA+TQoUOYPXs2YmJicPnll3s9X4t9+/ZhxowZ0Ov1SE9Px3333QebzeZyTGZmpsff8wsXLkRaWprzcUdt9NaON954AyNHjoROp0NcXBxmzpyJjRs3uhzTcn1HjhzBnDlzEBUVhbS0NNx3330dfs7q6Ln05/ye+tcTX1+LP/zwA2bMmIG4uDjodDqMGTMGn3/+ebvX42/ZH3/8ETNnzoTBYEBsbCwmT56MH374AUDHz5k/5/H19e9Ne+0EHK+3ljYKgoCkpCScd9552L59u0s93dH3JExxEhQVFRU8NzeX9+/fn69YsYLX19fzAwcO8Mcff5y/9957HstUVlby//znP1ytVvOXXnrJuX3hwoU8IyODr127lptMJl5SUsKXLl3K77//fucxVVVVPD8/nw8ZMoSvWrWKNzQ08PXr1/MBAwbwSZMmcVmWfW77BRdcwFUqlVsZRVG4RqPhF110kUubAXCDwcA1Gg2PioriY8eO5a+++qrHulNTU/mAAQM6bIMv19yWt/576623OAB++eWX81deeYXX1tby1atXc4PBwBcvXszvuece/tJLL/Gamhq+du1anpiYyBcuXOhSd0sd8+bN4//3f//Hq6ur+a5du/iZZ57JDQYDP3LkiPPY5cuXcwB8/fr1zm2+Pj8HDx7ksbGxfMqUKXzPnj28qqqKP/PMM/ziiy/mALy+dgJpZ1FREU9ISOBjx47lW7du5bW1tfy1117jWq2WX3HFFe1eT8t5rrjiCv7888/z6upqvnPnTj5w4EA+ZMgQriiK89iPPvqIA+C//fZbUJ5nzn17f3Xm+nwt29IPl1xyCX/22Wd5ZWUl//DDD3l1dTX/7LPPuCRJ/LrrruOHDh3itbW1/D//+Q9XqVQur89nnnmGM8b43/72N15ZWckPHDjAL7jgAj579mwuimK7/cA55zNnzuSFhYV8zpw5fM2aNby+vp5/8sknPC4ujp933nnO4/bv3+88z6nMZjNPSkri5557rtdz1NfX88TERH7ppZfy4uJibjKZ+Pbt2/ktt9zCf/nll4Cu5bnnnuMAeFFRUYfXSAgh3eHEiRM8JiaGX3PNNfzo0aO8ubmZ//7773zRokX8999/dx730EMPcVEU+bPPPsurq6v57t27+Zw5c/iMGTN4YmKiS52PP/44B8Bra2s7PL8/9c6cOZOPGjXK77ZPnjyZjxs3zuP5x40bx0eNGsVnzZrFN2zYwCsqKvi7777r8Xxtj//11195XV0df+edd3hUVBS/8sorXY7NyMjgCxYscDvnVVddxVNTU122tddGT+3429/+xhlj/PHHH+cVFRX80KFD/MILL+RqtZqvWrXKpb2jR4/mF154If/11195fX09f+WVVzgAl7/L3nh7Lv05v7f+bcvX5/PNN9/kjDF+xx138CNHjvCamhr+zDPPcFEUXer21G++ln3ttdc4Y4zfcMMNfN++fdxoNPJVq1bxuXPnOo9p7znz9Tz+vP498aWdp7JarXzXrl38vPPO46mpqbyysrJL+p5EFgqGBMktt9zCJUni+/bt87vsH//4Rz5ixAjn45ycHH711Ve3W+aOO+7garWaHzp0yGX7hg0bOAD+6aef+nz+iRMnev2lk5KSwidNmuR8XFVVxRcsWMDXr1/PGxoa+MGDB/nixYs5AH777bf7fM62fLlmb9r2X8sNa9v23HnnnVyn0/HFixe7bL/77ru5SqXiZrPZrY4bbrjB5dgTJ05wrVbLr7/+euc2TzfXvj4/V199NY+KiuJVVVUuxy1cuNCvYIgv7bzmmmu4Wq3mpaWlLsc+/PDDLsGL9oIht912m0vZDz74gAPga9eudW5rLxgS6PPsy/urM9fna9mWfrjuuutcjrNarTw9PZ1PnjzZY9sTEhK4xWLhzc3NPC4ujl988cUuxxiNRm4wGHwOhgDgq1evdtn+/PPPcwAuH8TOOuss3rdvX5dAZ8s1tPc7Ys2aNRwA//HHH70e4++1UDCEEBJuvvzySw6Ab9682esx1dXVXKPR8EWLFrlsr6io4Hq9PuBgiL/1tr259aXtnHccDAHAt27d6rbPWzAEAN++fbvL9ieeeMJte1cFQ6qrq7lWq+WXXnqpy3Fms5lnZGS41DNu3DguCALfs2ePy7ETJ07kp512msfzncrTc+nv+b31b1u+PJ+NjY3cYDDw888/323
"text/plain": [
"<Figure size 1150x430 with 3 Axes>"
]
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"# Visualize one case's point-level features. re_norm and aoa_norm are constant for the case;\n",
"# x, y, and sdf vary over sampled points.\n",
"case_id = (split_manifest.get(\"val_ids\") or split_manifest.get(\"test_ids\") or split_manifest.get(\"train_ids\"))[0]\n",
"features, _ = load_case(case_id)\n",
"feature_index = {name: idx for idx, name in enumerate(normalization[\"feature_names\"])}\n",
"\n",
"x_values = features[:, feature_index[\"x\"]]\n",
"y_values = features[:, feature_index[\"y\"]]\n",
"sdf_values = features[:, feature_index[\"sdf\"]]\n",
"\n",
"fig, axes = plt.subplots(1, 2, figsize=(11.5, 4.3))\n",
"scatter = axes[0].scatter(x_values, y_values, c=sdf_values, s=18, cmap=\"coolwarm\")\n",
"axes[0].set_title(f\"{case_id}: sampled points colored by sdf\")\n",
"axes[0].set_xlabel(\"x\")\n",
"axes[0].set_ylabel(\"y\")\n",
"axes[0].set_aspect(\"equal\", adjustable=\"box\")\n",
"axes[0].grid(alpha=0.25)\n",
"fig.colorbar(scatter, ax=axes[0], label=\"sdf\")\n",
"\n",
"axes[1].hist(sdf_values, bins=30, edgecolor=\"white\")\n",
"axes[1].set_title(\"sdf distribution for selected case\")\n",
"axes[1].set_xlabel(\"sdf\")\n",
"axes[1].set_ylabel(\"sampled points\")\n",
"axes[1].grid(alpha=0.25)\n",
"fig.tight_layout()\n",
"\n",
"print(f\"case-level conditions for {case_id}:\")\n",
"print(f\"re_norm = {features[0, feature_index['re_norm']]:.3f}\")\n",
"print(f\"aoa_norm = {features[0, feature_index['aoa_norm']]:.3f}\")\n"
]
},
{
"cell_type": "markdown",
"id": "3c22f84b",
"metadata": {},
"source": [
"## Checkpoint contents\n",
"\n",
"The checkpoint is saved with CPU-compatible tensors. This section verifies it can be loaded, shows metadata, and reconstructs the model.\n"
]
},
{
"cell_type": "code",
"execution_count": 13,
"id": "ed0720ec",
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"checkpoint keys:\n",
" step: int\n",
" model_type: str\n",
" input_dim: int\n",
" output_dim: int\n",
" model_config: dict\n",
" model_state_dict: 10 tensors\n",
" optimizer_state_dict: 2 tensors\n",
" normalization: dict\n",
" target_names: tuple\n",
" feature_names: tuple\n",
" final_metrics: dict\n",
"\n",
"step: 500\n",
"model_type: mlp\n",
"input_dim: 5\n",
"output_dim: 4\n",
"feature_names: ('re_norm', 'aoa_norm', 'x', 'y', 'sdf')\n",
"target_names: ('velocity_x', 'velocity_y', 'pressure', 'turbulent_viscosity')\n",
"model_config: {'type': 'mlp', 'hidden_width': 128, 'depth': 4, 'activation': 'gelu'}\n",
"\n",
"first model tensors:\n",
"network.0.weight (128, 5) torch.float32\n",
"network.0.bias (128,) torch.float32\n",
"network.2.weight (128, 128) torch.float32\n",
"network.2.bias (128,) torch.float32\n",
"network.4.weight (128, 128) torch.float32\n",
"network.4.bias (128,) torch.float32\n",
"network.6.weight (128, 128) torch.float32\n",
"network.6.bias (128,) torch.float32\n"
]
}
],
"source": [
"checkpoint = torch.load(RUN_DIR / \"checkpoint.pt\", map_location=\"cpu\")\n",
"\n",
"print(\"checkpoint keys:\")\n",
"for key, value in checkpoint.items():\n",
" if key.endswith(\"state_dict\"):\n",
" print(f\" {key}: {len(value)} tensors\")\n",
" else:\n",
" print(f\" {key}: {type(value).__name__}\")\n",
"\n",
"print()\n",
"print(\"step:\", checkpoint[\"step\"])\n",
"print(\"model_type:\", checkpoint[\"model_type\"])\n",
"print(\"input_dim:\", checkpoint[\"input_dim\"])\n",
"print(\"output_dim:\", checkpoint[\"output_dim\"])\n",
"print(\"feature_names:\", checkpoint[\"feature_names\"])\n",
"print(\"target_names:\", checkpoint[\"target_names\"])\n",
"print(\"model_config:\", checkpoint[\"model_config\"])\n",
"\n",
"print()\n",
"print(\"first model tensors:\")\n",
"for name, tensor in list(checkpoint[\"model_state_dict\"].items())[:8]:\n",
" print(name, tuple(tensor.shape), tensor.dtype)\n"
]
},
{
"cell_type": "markdown",
"id": "94b4e5d8",
"metadata": {},
"source": [
"## MLP architecture\n",
"\n",
"The training loop derives `input_dim` from normalized feature columns and `output_dim` from normalized target columns. This cell expands that into the concrete network shape, layer table, and parameter count so the abstraction is visible at a glance."
]
},
{
"cell_type": "code",
"execution_count": 14,
"id": "e48801e5",
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"MLP shape at a glance\n",
"----------------------\n",
"inputs (5): ['re_norm', 'aoa_norm', 'x', 'y', 'sdf']\n",
"outputs (4): ['velocity_x', 'velocity_y', 'pressure', 'turbulent_viscosity']\n",
"architecture: 5 -> 128 -> 128 -> 128 -> 128 -> 4\n",
"activation after hidden layers: gelu\n",
"hidden layers: 4\n",
"hidden width: 128\n",
"\n",
"layer table:\n",
" 0: Linear 5 -> 128 params=768\n",
" 1: GELU -> params=0\n",
" 2: Linear 128 -> 128 params=16512\n",
" 3: GELU -> params=0\n",
" 4: Linear 128 -> 128 params=16512\n",
" 5: GELU -> params=0\n",
" 6: Linear 128 -> 128 params=16512\n",
" 7: GELU -> params=0\n",
" 8: Linear 128 -> 4 params=516\n",
"\n",
"total parameters: 50,820\n"
]
}
],
"source": [
"input_names = list(checkpoint[\"feature_names\"])\n",
"target_names = list(checkpoint[\"target_names\"])\n",
"model_config = checkpoint[\"model_config\"]\n",
"\n",
"print(\"MLP shape at a glance\")\n",
"print(\"----------------------\")\n",
"print(f\"inputs ({len(input_names)}): {input_names}\")\n",
"print(f\"outputs ({len(target_names)}): {target_names}\")\n",
"print(\n",
" \"architecture: \"\n",
" f\"{checkpoint['input_dim']} -> \"\n",
" + \" -> \".join([str(model_config[\"hidden_width\"])] * model_config[\"depth\"])\n",
" + f\" -> {checkpoint['output_dim']}\"\n",
")\n",
"print(f\"activation after hidden layers: {model_config['activation']}\")\n",
"print(f\"hidden layers: {model_config['depth']}\")\n",
"print(f\"hidden width: {model_config['hidden_width']}\")\n",
"print()\n",
"\n",
"model = PointwiseMLP(\n",
" input_dim=checkpoint[\"input_dim\"],\n",
" output_dim=checkpoint[\"output_dim\"],\n",
" hidden_width=model_config[\"hidden_width\"],\n",
" depth=model_config[\"depth\"],\n",
" activation=model_config[\"activation\"],\n",
")\n",
"model.load_state_dict(checkpoint[\"model_state_dict\"])\n",
"model.eval()\n",
"\n",
"layer_rows = []\n",
"for index, layer in enumerate(model.network):\n",
" row = {\"index\": index, \"type\": layer.__class__.__name__}\n",
" if isinstance(layer, torch.nn.Linear):\n",
" row[\"in_features\"] = layer.in_features\n",
" row[\"out_features\"] = layer.out_features\n",
" row[\"parameters\"] = layer.weight.numel() + layer.bias.numel()\n",
" else:\n",
" row[\"in_features\"] = \"\"\n",
" row[\"out_features\"] = \"\"\n",
" row[\"parameters\"] = 0\n",
" layer_rows.append(row)\n",
"\n",
"print(\"layer table:\")\n",
"for row in layer_rows:\n",
" print(\n",
" f\"{row['index']:>2}: {row['type']:<8} \"\n",
" f\"{str(row['in_features']):>4} -> {str(row['out_features']):<4} \"\n",
" f\"params={row['parameters']}\"\n",
" )\n",
"\n",
"parameter_count = sum(parameters.numel() for parameters in model.parameters())\n",
"print()\n",
"print(f\"total parameters: {parameter_count:,}\")\n",
"assert parameter_count == final_metrics[\"parameter_count\"]\n"
]
},
{
"cell_type": "markdown",
"id": "6f638448",
"metadata": {},
"source": [
"## Reloaded-model predictions\n",
"\n",
"This section runs the checkpoint on real `.npz` cases from the split. Predictions are de-normalized back to the target scale before comparison.\n"
]
},
{
"cell_type": "code",
"execution_count": 15,
"id": "09e67ecb",
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"/home/aaron/data/airfrans/data/processed/minimal\n",
"['case_00.npz', 'case_01.npz', 'case_02.npz', 'case_03.npz', 'case_04.npz', 'case_05.npz']\n"
]
}
],
"source": [
"def config_value(config: str, section: str, key: str) -> str:\n",
" current_section = None\n",
" for raw_line in config.splitlines():\n",
" line = raw_line.strip()\n",
" if not line or line.startswith(\"#\"):\n",
" continue\n",
" if line.startswith(\"[\") and line.endswith(\"]\"):\n",
" current_section = line[1:-1]\n",
" continue\n",
" if current_section == section and line.startswith(f\"{key} =\"):\n",
" return line.split(\"=\", 1)[1].strip().strip('\\\"')\n",
" raise KeyError(f\"{section}.{key}\")\n",
"\n",
"\n",
"DATA_ROOT = REPO_ROOT / config_value(config_text, \"data\", \"root\")\n",
"print(DATA_ROOT)\n",
"print(sorted(path.name for path in DATA_ROOT.glob(\"*.npz\")))\n"
]
},
{
"cell_type": "code",
"execution_count": 16,
"id": "924e2cef",
"metadata": {},
"outputs": [
{
"data": {
"text/plain": [
"[{'split': 'test',\n",
" 'case_id': 'case_03',\n",
" 'rows': 256,\n",
" 'normalized_mse': 0.00025704235304147005,\n",
" 'mse_velocity_x': 5.847604370501358e-06,\n",
" 'mse_velocity_y': 6.842037691967562e-06,\n",
" 'mse_pressure': 4.433627691469155e-05,\n",
" 'mse_turbulent_viscosity': 1.3116708032612223e-05},\n",
" {'split': 'train',\n",
" 'case_id': 'case_04',\n",
" 'rows': 256,\n",
" 'normalized_mse': 0.00018982301116921008,\n",
" 'mse_velocity_x': 6.853222657809965e-06,\n",
" 'mse_velocity_y': 6.716706138831796e-06,\n",
" 'mse_pressure': 2.8528093025670387e-05,\n",
" 'mse_turbulent_viscosity': 9.420269634574652e-06},\n",
" {'split': 'train',\n",
" 'case_id': 'case_02',\n",
" 'rows': 256,\n",
" 'normalized_mse': 0.0001731793163344264,\n",
" 'mse_velocity_x': 5.711051926482469e-06,\n",
" 'mse_velocity_y': 4.464081939659081e-06,\n",
" 'mse_pressure': 2.741067874012515e-05,\n",
" 'mse_turbulent_viscosity': 9.072812645172235e-06},\n",
" {'split': 'train',\n",
" 'case_id': 'case_01',\n",
" 'rows': 256,\n",
" 'normalized_mse': 0.00011508660099934787,\n",
" 'mse_velocity_x': 3.957362423534505e-06,\n",
" 'mse_velocity_y': 2.683263573999284e-06,\n",
" 'mse_pressure': 2.2211312170838937e-05,\n",
" 'mse_turbulent_viscosity': 5.00482428833493e-06},\n",
" {'split': 'train',\n",
" 'case_id': 'case_00',\n",
" 'rows': 256,\n",
" 'normalized_mse': 0.00016892868734430522,\n",
" 'mse_velocity_x': 8.186969353118911e-06,\n",
" 'mse_velocity_y': 7.417508186335908e-06,\n",
" 'mse_pressure': 3.0236758902901784e-05,\n",
" 'mse_turbulent_viscosity': 5.846786280017113e-06},\n",
" {'split': 'val',\n",
" 'case_id': 'case_05',\n",
" 'rows': 256,\n",
" 'normalized_mse': 0.0007373533444479108,\n",
" 'mse_velocity_x': 2.7578798835747875e-05,\n",
" 'mse_velocity_y': 1.9889514078386128e-05,\n",
" 'mse_pressure': 0.00013202131958678365,\n",
" 'mse_turbulent_viscosity': 3.3118620194727555e-05}]"
]
},
"execution_count": 16,
"metadata": {},
"output_type": "execute_result"
}
],
"source": [
"def load_case(case_id: str) -> tuple[np.ndarray, np.ndarray]:\n",
" data = np.load(DATA_ROOT / f\"{case_id}.npz\")\n",
" return data[\"features\"].astype(np.float32), data[\"targets\"].astype(np.float32)\n",
"\n",
"\n",
"def predict_targets(features: np.ndarray, batch_size: int = 4096) -> np.ndarray:\n",
" features_norm = (features - feature_mean) / feature_std\n",
" outputs = []\n",
" with torch.no_grad():\n",
" for start in range(0, len(features_norm), batch_size):\n",
" batch = torch.from_numpy(features_norm[start : start + batch_size])\n",
" pred_norm = model(batch).numpy()\n",
" outputs.append(pred_norm)\n",
" pred_norm_all = np.concatenate(outputs, axis=0)\n",
" return pred_norm_all * target_std + target_mean\n",
"\n",
"\n",
"def normalized_mse(predictions: np.ndarray, targets: np.ndarray) -> float:\n",
" pred_norm = (predictions - target_mean) / target_std\n",
" target_norm = (targets - target_mean) / target_std\n",
" return float(np.mean((pred_norm - target_norm) ** 2))\n",
"\n",
"\n",
"case_rows = []\n",
"for split_name, ids in split_manifest.items():\n",
" for case_id in ids:\n",
" features, targets = load_case(case_id)\n",
" predictions = predict_targets(features)\n",
" per_channel_mse = np.mean((predictions - targets) ** 2, axis=0)\n",
" case_rows.append(\n",
" {\n",
" \"split\": split_name.replace(\"_ids\", \"\"),\n",
" \"case_id\": case_id,\n",
" \"rows\": len(features),\n",
" \"normalized_mse\": normalized_mse(predictions, targets),\n",
" **{f\"mse_{name}\": float(value) for name, value in zip(checkpoint[\"target_names\"], per_channel_mse)},\n",
" }\n",
" )\n",
"\n",
"case_rows\n"
]
},
{
"cell_type": "code",
"execution_count": 17,
"id": "4fe67c98",
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"case: case_05\n",
"target names: ('velocity_x', 'velocity_y', 'pressure', 'turbulent_viscosity')\n",
"first 5 predictions:\n",
"[[ 0.44094524 -0.02950683 0.19028291 0.0982319 ]\n",
" [ 0.02009592 0.14751981 -0.21558036 0.03777217]\n",
" [ 0.45918038 -0.00707006 0.2556273 0.12591958]\n",
" [ 0.31692797 0.19948506 -0.60716456 0.3552236 ]\n",
" [ 0.23262833 0.05736347 -0.33526948 0.07037625]]\n",
"first 5 actual targets:\n",
"[[ 0.44482228 -0.03109705 0.18664566 0.09674599]\n",
" [ 0.01824543 0.1526474 -0.20674932 0.03895786]\n",
" [ 0.4640298 -0.00911091 0.25503194 0.12455965]\n",
" [ 0.32461593 0.1904991 -0.6266293 0.35012734]\n",
" [ 0.23149356 0.05860274 -0.33879325 0.06851949]]\n"
]
},
{
"data": {
"image/png": "iVBORw0KGgoAAAANSUhEUgAABjYAAAGMCAYAAAB51ps5AAAAOnRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjExLjEsIGh0dHBzOi8vbWF0cGxvdGxpYi5vcmcvctoD+AAAAAlwSFlzAAAPYQAAD2EBqD+naQABAABJREFUeJzs3Xd0VNXax/HvTHonjQQIvVdBei/SBKUrVb1SpIjlioogVlTsqCggTRQVQZAiKCLSexOkE3pNT0hCymQy5/2DN3ONCZAgkAz8PmvdtW727LPPs3fiHM55zt7bZBiGgYiIiIiIiIiIiIiIiAMwF3QAIiIiIiIiIiIiIiIieaXEhoiIiIiIiIiIiIiIOAwlNkRERERERERERERExGEosSEiIiIiIiIiIiIiIg5DiQ0REREREREREREREXEYSmyIiIiIiIiIiIiIiIjDUGJDREREREREREREREQchhIbIiIiIiIiIiIiIiLiMJTYEBEREXFwaWlpTJ06lX379hXI+WfMmMHmzZsLZZuJiYlMnTqVw4cP34SoCv95oeD/HgrK8ePHmTp1KrGxsfayvXv3MnXqVCwWy007z61o81b466+/+Pbbb5k6dSoXLlwo6HBERERERG4qJTZEREREbqPz588zdepU+/+mTZvGggULOH369A23mZyczPDhw/njjz9uYqR5N3LkSObPn18o24yKimL48OFs3LjxJkRV+M8L//7v4fLly0ydOpUDBw7c5MhurR07djB8+HDOnj1rL/v9998ZPnw4KSkp+Wpr9+7dTJ06FavVmuOzG23zdnr11Vdp3LgxK1asYM+ePSQnJxd0SIVacnIyS5cuZfbs2WzcuBGbzZajTkRERLbv7r//Lzw8vACiFhEREbm7ORd0ACIiIiJ3k0OHDjF8+HCaN29OtWrVsFqtHDhwgG3btvH4448zbdo0nJycCjpMcWAeHh4MHTqUWrVq3dDx8fHxDB8+nEmTJlG9evWbHN3tVbt2bYYOHYqbm1u+jvvll1945ZVXGDBgAN7e3jelzdslMTGRCRMm8OabbzJmzJiCDqfQ+/333+nbty/lypWjatWqvPzyy4SFhbFs2TKCg4Pt9Y4dO8bw4cPp0KEDZcqUydbGvffee5ujFhERERElNkREREQKQL9+/Rg2bJj951dffZXx48fTqFEjhgwZUoCRiaPz8vJi6tSpBR1GodC2bVvatm1b6Nu8mU6dOoXVaqVcuXIFHUqhFxERQc+ePenSpQtz5szBZDIRFxdHnTp1eOyxx/jll19yHDNs2DC6det2+4MVERERkWyU2BAREREpBIYMGcL48eNZuXJljsSG1Wpl8+bNnDhxAk9PT1q0aEFoaGie2k1OTmb9+vVEREQQEhJCixYt8PHxyVZn2rRp9qVXXF1dCQsLo3nz5nh4eORoLyoqitWrVwPQpk0bihYtetVz5zXu/LSZm8zMTLZu3cqxY8cIDAykcePGBAYG5lp33bp1nDhxgsqVK9OkSZN/FXd+zpvl7NmzLF++nPLly9OuXTsSExP5/vvvadWqFZUqVWLNmjWcPXuWqlWr0rBhw1zbuN7vNC0tjdmzZ9O0aVNq1qwJkO08VapUueo4RERE8O233wKwceNGnJ2v3C7UqFGDZs2a5RpPfvrw97qVK1dm48aNHD16lHbt2lGqVCkAYmJi2LhxI3FxcZQqVYoWLVrg6uqa47w7d+5k//79FC9enPvuuy/X2Pbu3cuWLVsYOHBgjjbi4uLYtGkTsbGxVKpUiUaNGmE2m1m3bh07duwAYObMmfaZGT179iQ4OPiabZ44cYJt27ZhsVioWrUqDRo0yPb5jh072LVrF0OHDiUpKYnffvsNi8VCy5YtCQsLyxF/fHw8mzdvJiYmhjJlytCgQYNc/7vM8vvvv7Nq1SoAVq9eTXx8PP7+/vTu3fuGYkxMTOSPP/4gKSmJxx577KrnvdZ4AoSHh9uXRjOZTHh5eVGzZk3uueeeG+5zXv9OrmXatGkkJSXx+uuvYzKZAAgICOCpp57ihRde4MCBAw4/a0lERETkTqU9NkREREQKgawHcv9c03/Hjh1UqVKF/v37s2LFCqZNm0bZsmX59NNPr9vm0qVLKV26NP/9739ZtWoVzz//PKVLl2bRokXZ6u3du5c9e/awZ88e/vjjD4YOHUrp0qXZtm1btnrff/89ZcuW5Z133mHp0qXcd999/P7777meO69x56fN3Gzfvp3KlSvTs2dPli5dyvTp06lfvz7ff/99tnppaWl069aNDz74gCVLltCyZUv69+9/w3Hn9bz/bLtBgwZ899131KlTB/jfXhy//fYbHTt2ZOLEiSxYsICWLVvSrVs30tPTs7WRl99pbntsZJ1n1apV1xyHy5cv2/fWOH/+vP3v4vz581ftV3768Pe6nTp1YsKECcycOZO//voLgHfeeYeSJUvy3nvvsXr1aoYPH061atWy7fdhsVh4+OGHadKkCfPmzeOLL76gbdu2XLp0KUdsV9sP4+2336ZEiRK8+uqrrFq1itGjR1O/fn3OnTvH2bNnuXjxIgD79u2zj0FqaupV28zMzGTo0KFUrlyZWbNmsXTpUtq3b0+TJk3sbQEsWbKE4cOHs23bNlq3bs2iRYv44IMPqFChAsuWLcsW49y5cylZsiQTJkxg7dq1jB8/nnvuuYd169Zd9Xdx8uRJ+34PJ0+eZM+ePRw6dOiGYly/fj0tWrRg3rx5fPDBB1c95/XGE64kK7LGcffu3cybN49GjRrRrl27bJuw57XPefk7yYs1a9YQEhJChQoVspW3bNkSINd9anbt2sWsWbOYO3cuR48ezdf5REREROQmMkRERETktvn9998NwJgyZUq28unTpxuA8d5779nLIiIijICAAKNDhw7G5cuX7eWzZ882AGP16tWGYRhGdHS0ARgTJ0601wkPDzfc3d2Nnj17GhaLxTAMw8jIyDB69+5tuLm5GYcOHbpqjFar1ejSpYtRsWJFe9mRI0cMV1dX4z//+Y+RmZlpGIZhpKSkGH369DGcnZ2NZ555Jt9x56fN3Jw/f94oUqSI0aZNGyMhIcFenpSUZD9HeHi4ARgVK1Y0du7caa8zbdo0AzDWrl2b77jzc97p06cbhmEYP/74o+Hh4WEMGDDASEtLsx+TVa9s2bLGtm3b7OUbNmwwnJycjNGjR2erm5ff6dX+HvI6DmfPnjUAY9KkSdcc/xvtQ1bdrVu3GoZhGDabzTh//rw9lhkzZtjrp6WlGR06dDAqVKhgZGRkGIZhGK+99pphMpmMVatW2evt3r3bqFChggEYf/75p738gw8+MAAjPj7eXjZ58uRc+7d3717j1KlThmEYxvjx4w3ASEpKytHf3NrMqv/TTz/Zy44dO2YULVrUaNGihb3s5ZdfNgBjwIABRnJysmEYV36HLVq0MMqXL29YrVb7mPj7+xtPPvlktnNfuHDBPm5Xs2XLFgMwfvzxx2zl+Y3x4Ycftvf/7NmzVz1fXsYzN8ePHzd8fX2Nt99+O199zuvfSV6UKlXKqFevXo7yixcvGkC276ANGzYYZrPZaNSokTFgwACjadOmBmA8+OCD2f4WREREROT2UGJDRERE5DbKSmz069fPmDJlijFp0iRj6NChhouLi9GpU6dsD9TffPNNAzCOHDmSo53KlSsbDz/8sGEYuT/IHj16tAEYJ06cyHbc2bNnDZPJZDz77LPZyo8dO2bMnz/f+PLLL40pU6YYjz76qAEYFy9etLdnNpuNiIiIbMdt3LgxxwPAvMadnzZz8+qrrxrANZM0WQ/S//Of/2QrT09PN8xms/HKK6/kO+78nHf69OnGhAkTDLPZbLz++utXrdevX78cn/Xq1cvw8/OzJ33y+ju9VmIjL+Nwo4mNvPQhq27v3r1z1C1XrpzRuHHjHOVZD+p/+eUXwzAMIyQkxOjUqVOOeoMHD85TYqNUqVJGw4YNr9mn/CY2QkNDsyUHsrz99tsGYOzZs8cwjP8lDf6eSDIMw5gxY4YBGMePHzcM48qDeicnpxwP+fPiaomN/Mb466+/5ul8eRlPw7iSMN2
"text/plain": [
"<Figure size 1600x380 with 4 Axes>"
]
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"# Pick a validation case if present, otherwise test, otherwise train.\n",
"case_id = (split_manifest.get(\"val_ids\") or split_manifest.get(\"test_ids\") or split_manifest.get(\"train_ids\"))[0]\n",
"features, targets = load_case(case_id)\n",
"predictions = predict_targets(features)\n",
"\n",
"print(f\"case: {case_id}\")\n",
"print(\"target names:\", checkpoint[\"target_names\"])\n",
"print(\"first 5 predictions:\")\n",
"print(predictions[:5])\n",
"print(\"first 5 actual targets:\")\n",
"print(targets[:5])\n",
"\n",
"fig, axes = plt.subplots(1, len(checkpoint[\"target_names\"]), figsize=(4 * len(checkpoint[\"target_names\"]), 3.8))\n",
"if len(checkpoint[\"target_names\"]) == 1:\n",
" axes = [axes]\n",
"for index, (ax, name) in enumerate(zip(axes, checkpoint[\"target_names\"])):\n",
" ax.scatter(targets[:, index], predictions[:, index], s=12, alpha=0.7)\n",
" low = float(min(targets[:, index].min(), predictions[:, index].min()))\n",
" high = float(max(targets[:, index].max(), predictions[:, index].max()))\n",
" ax.plot([low, high], [low, high], color=\"black\", linewidth=1)\n",
" ax.set_title(name)\n",
" ax.set_xlabel(\"actual\")\n",
" ax.set_ylabel(\"predicted\")\n",
" ax.grid(alpha=0.25)\n",
"fig.suptitle(f\"Reloaded checkpoint predictions for {case_id}\", y=1.03)\n",
"fig.tight_layout()\n"
]
},
{
"cell_type": "markdown",
"id": "5901d5d7",
"metadata": {},
"source": [
"## What good looks like\n",
"\n",
"For this tiny baseline sanity run:\n",
"\n",
"- `train_loss` is finite and much lower than `initial_train_loss`.\n",
"- `val_loss` and `test_loss` are finite.\n",
"- The loss curve falls on a log-scale plot rather than staying flat or exploding.\n",
"- `checkpoint.pt` loads on CPU and reconstructs `PointwiseMLP` without missing/unexpected keys.\n",
"- Reloaded predictions have the expected target channels and roughly follow the identity line in predicted-vs-actual plots.\n",
"\n",
"This does not prove the model is scientifically useful. It proves the training framework can run end-to-end, write readable artifacts, reload the model, and produce numerically sensible outputs on the processed mini dataset.\n"
]
}
],
"metadata": {
"kernelspec": {
"display_name": ".venv",
"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.12"
}
},
"nbformat": 4,
"nbformat_minor": 5
}