airfRANS-model-exploration/notebooks/airfrans_aggressive_mechinterp.ipynb

1220 lines
1 MiB
Text
Raw Normal View History

{
"cells": [
{
"cell_type": "markdown",
"id": "ca507338",
"metadata": {},
"source": [
"# Aggressive AirfRANS model mechinterp\n",
"\n",
"Notebook for poking at `artifacts/remote_runs/airfrans-aggressive-smoke-20260723T083947Z`.\n",
"\n",
"It loads the large `film_fourier_mlp` checkpoint, runs small deterministic probes, and keeps defaults intentionally light so a full top-to-bottom run works on CPU. Increase the sample-count constants after the first pass.\n"
]
},
{
"cell_type": "markdown",
"id": "cf0fd290",
"metadata": {},
"source": [
"## 1. Setup and run selection\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": "48d1e21a",
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"python: /home/aaron/data/airfrans/.venv/bin/python\n",
"repo: /home/aaron/data/airfrans\n",
"run: artifacts/remote_runs/airfrans-aggressive-smoke-20260723T083947Z\n"
]
}
],
"source": [
"from __future__ import annotations\n",
"\n",
"import hashlib\n",
"import json\n",
"import math\n",
"import os\n",
"import sys\n",
"import tomllib\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",
"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.film import FourierFiLMMLP\n",
"\n",
"RUN_DIR = REPO_ROOT / \"artifacts\" / \"remote_runs\" / \"airfrans-aggressive-smoke-20260723T083947Z\"\n",
"CHECKPOINT_NAME = \"checkpoint_best.pt\" # best val checkpoint; switch to checkpoint_final.pt if desired.\n",
"USE_CUDA = False # CPU default is slower but avoids small-GPU OOM for the 253M-param model.\n",
"\n",
"CASE_SAMPLE_POINTS = 768\n",
"FAST_SAMPLE_POINTS = 192\n",
"GRAD_SAMPLE_POINTS = 32\n",
"INFER_BATCH_SIZE = 128\n",
"RNG_SEED = 20260724\n",
"\n",
"print(f\"python: {sys.executable}\")\n",
"print(f\"repo: {REPO_ROOT}\")\n",
"print(f\"run: {RUN_DIR.relative_to(REPO_ROOT)}\")\n",
"assert RUN_DIR.is_dir(), RUN_DIR\n"
]
},
{
"cell_type": "markdown",
"id": "b5665781",
"metadata": {},
"source": [
"## 2. Load artifacts and checkpoint\n",
"\n",
"The checkpoint is memory-mapped on CPU. That avoids eagerly copying the full optimizer state while still letting us reconstruct the model weights.\n"
]
},
{
"cell_type": "code",
"execution_count": 2,
"id": "7ee0592e",
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"checkpoint: checkpoint_best.pt\n",
"step: 3500\n",
"model: film_fourier_mlp\n",
"params: 253,681,156\n",
"best val: 0.0760959\n",
"final val: 0.0898282\n",
"test loss: 0.220556\n",
"features: ['x', 'y', 'sdf', 'u_inf', 'log_re', 'aoa_deg', 'aoa_sin', 'aoa_cos', 'naca_param_0', 'naca_param_1', 'naca_param_2', 'naca_param_3', 'naca_param_0_mask', 'naca_param_1_mask', 'naca_param_2_mask', 'naca_param_3_mask']\n",
"targets: ['velocity_x', 'velocity_y', 'pressure', 'turbulent_viscosity']\n"
]
}
],
"source": [
"final_metrics = json.loads((RUN_DIR / \"final_metrics.json\").read_text())\n",
"split_manifest = json.loads((RUN_DIR / \"split_manifest.json\").read_text())\n",
"normalization = json.loads((RUN_DIR / \"normalization.json\").read_text())\n",
"checkpoint_path = RUN_DIR / CHECKPOINT_NAME\n",
"checkpoint = torch.load(checkpoint_path, map_location=\"cpu\", weights_only=False, mmap=True)\n",
"\n",
"print(f\"checkpoint: {checkpoint_path.name}\")\n",
"print(f\"step: {checkpoint['step']}\")\n",
"print(f\"model: {checkpoint['model_type']}\")\n",
"print(f\"params: {final_metrics['parameter_count']:,}\")\n",
"print(f\"best val: {final_metrics['best_val_loss']:.6g}\")\n",
"print(f\"final val: {final_metrics['val_loss']:.6g}\")\n",
"print(f\"test loss: {final_metrics['test_loss']:.6g}\")\n",
"print(f\"features: {list(checkpoint['feature_names'])}\")\n",
"print(f\"targets: {list(checkpoint['target_names'])}\")\n",
"\n",
"assert checkpoint[\"model_type\"] == \"film_fourier_mlp\"\n",
"assert tuple(checkpoint[\"feature_names\"]) == tuple(normalization[\"feature_names\"])\n",
"assert tuple(checkpoint[\"target_names\"]) == tuple(normalization[\"target_names\"])\n"
]
},
{
"cell_type": "markdown",
"id": "5b818291",
"metadata": {},
"source": [
"## 3. Rebuild the FiLM/Fourier MLP\n",
"\n",
"The model receives normalized rows. Its first three features are coordinates (`x`, `y`, `sdf`); the remaining features are case conditions used by the FiLM path.\n"
]
},
{
"cell_type": "code",
"execution_count": 3,
"id": "6dbc747b",
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"<All keys matched successfully>\n",
"device: cpu\n",
"architecture: 16 normalized inputs -> 12 x 4096 trunk -> 4 targets\n",
"condition names: ['u_inf', 'log_re', 'aoa_deg', 'aoa_sin', 'aoa_cos', 'naca_param_0', 'naca_param_1', 'naca_param_2', 'naca_param_3', 'naca_param_0_mask', 'naca_param_1_mask', 'naca_param_2_mask', 'naca_param_3_mask']\n",
"parameter count: 253,681,156\n"
]
}
],
"source": [
"model_config = checkpoint[\"model_config\"]\n",
"model = FourierFiLMMLP(\n",
" feature_names=checkpoint[\"feature_names\"],\n",
" output_dim=checkpoint[\"output_dim\"],\n",
" coordinate_features=model_config[\"coordinate_features\"],\n",
" fourier_scales=model_config[\"fourier_scales\"],\n",
" trunk_width=model_config[\"hidden_width\"],\n",
" trunk_depth=model_config[\"depth\"],\n",
" condition_width=model_config[\"condition_width\"],\n",
" condition_depth=model_config[\"condition_depth\"],\n",
" condition_dim=model_config[\"condition_dim\"],\n",
" activation=model_config[\"activation\"],\n",
")\n",
"state_result = model.load_state_dict(checkpoint[\"model_state_dict\"], strict=True)\n",
"model.eval()\n",
"for parameter in model.parameters():\n",
" parameter.requires_grad_(False)\n",
"\n",
"if USE_CUDA and torch.cuda.is_available():\n",
" device = torch.device(\"cuda:0\")\n",
" model.to(device)\n",
"else:\n",
" device = torch.device(\"cpu\")\n",
"\n",
"parameter_count = sum(parameter.numel() for parameter in model.parameters())\n",
"print(state_result)\n",
"print(f\"device: {device}\")\n",
"print(f\"architecture: {checkpoint['input_dim']} normalized inputs -> \"\n",
" f\"{model_config['depth']} x {model_config['hidden_width']} trunk -> {checkpoint['output_dim']} targets\")\n",
"print(f\"condition names: {list(model.condition_names)}\")\n",
"print(f\"parameter count: {parameter_count:,}\")\n",
"assert parameter_count == final_metrics[\"parameter_count\"]\n"
]
},
{
"cell_type": "markdown",
"id": "66cff427",
"metadata": {},
"source": [
"## 4. Shared helpers\n",
"\n",
"These helpers keep every probe deterministic and cheap. Model outputs are de-normalized only for plots/tables; model losses stay in normalized target space.\n"
]
},
{
"cell_type": "code",
"execution_count": 4,
"id": "da5cd3d9",
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"data root: data/processed/full\n",
"example: airFoil2D_SST_83.526_0.275_3.114_4.789_1.0_12.835\n",
"{'case_id': 'airFoil2D_SST_83.526_0.275_3.114_4.789_1.0_12.835', 'split': 'val', 'u_inf': 83.5260009765625, 'log_re': 15.49339771270752, 'aoa_deg': 0.2750000059604645, 'aoa_sin': 0.004799637012183666, 'aoa_cos': 0.9999884963035583, 'naca_param_3': 12.835000038146973}\n"
]
}
],
"source": [
"config = tomllib.loads(checkpoint[\"config\"])\n",
"DATA_ROOT = REPO_ROOT / config[\"data\"][\"root\"]\n",
"FEATURE_NAMES = tuple(checkpoint[\"feature_names\"])\n",
"TARGET_NAMES = tuple(checkpoint[\"target_names\"])\n",
"FEATURE_INDEX = {name: index for index, name in enumerate(FEATURE_NAMES)}\n",
"TARGET_INDEX = {name: index for index, name in enumerate(TARGET_NAMES)}\n",
"COORDINATE_NAMES = tuple(model.coordinate_names)\n",
"CONDITION_NAMES = tuple(model.condition_names)\n",
"CONDITION_INDICES = [FEATURE_INDEX[name] for name in CONDITION_NAMES]\n",
"COORDINATE_INDICES = [FEATURE_INDEX[name] for name in COORDINATE_NAMES]\n",
"\n",
"feature_mean = np.asarray(normalization[\"feature_mean\"], dtype=np.float32)\n",
"feature_std = np.asarray(normalization[\"feature_std\"], dtype=np.float32)\n",
"target_mean = np.asarray(normalization[\"target_mean\"], dtype=np.float32)\n",
"target_std = np.asarray(normalization[\"target_std\"], dtype=np.float32)\n",
"\n",
"\n",
"def load_case(case_id: str, *, max_points: int | None = None, seed: int = RNG_SEED) -> dict[str, np.ndarray | str]:\n",
" path = DATA_ROOT / f\"{case_id}.npz\"\n",
" with np.load(path, allow_pickle=False) as npz:\n",
" features = np.asarray(npz[\"features\"], dtype=np.float32)\n",
" targets = np.asarray(npz[\"targets\"], dtype=np.float32)\n",
" if max_points is not None and max_points < len(features):\n",
" digest = hashlib.blake2b(f\"{case_id}:{seed}\".encode(), digest_size=8).digest()\n",
" stable_seed = int.from_bytes(digest, \"little\") % (2**32)\n",
" rng = np.random.default_rng(stable_seed)\n",
" indices = np.sort(rng.choice(len(features), size=max_points, replace=False))\n",
" features = np.ascontiguousarray(features[indices], dtype=np.float32)\n",
" targets = np.ascontiguousarray(targets[indices], dtype=np.float32)\n",
" return {\"case_id\": case_id, \"features\": features, \"targets\": targets}\n",
"\n",
"\n",
"def normalize_features_raw(features: np.ndarray) -> np.ndarray:\n",
" return np.ascontiguousarray((features - feature_mean) / feature_std, dtype=np.float32)\n",
"\n",
"\n",
"def normalize_targets_raw(targets: np.ndarray) -> np.ndarray:\n",
" return np.ascontiguousarray((targets - target_mean) / target_std, dtype=np.float32)\n",
"\n",
"\n",
"def denormalize_targets(targets_norm: np.ndarray) -> np.ndarray:\n",
" return np.ascontiguousarray(targets_norm * target_std + target_mean, dtype=np.float32)\n",
"\n",
"\n",
"def predict_norm(features: np.ndarray, *, batch_size: int = INFER_BATCH_SIZE) -> np.ndarray:\n",
" features_norm = normalize_features_raw(features)\n",
" outputs: list[np.ndarray] = []\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]).to(device)\n",
" outputs.append(model(batch).detach().cpu().numpy())\n",
" return np.concatenate(outputs, axis=0)\n",
"\n",
"\n",
"def predict_targets(features: np.ndarray, *, batch_size: int = INFER_BATCH_SIZE) -> np.ndarray:\n",
" return denormalize_targets(predict_norm(features, batch_size=batch_size))\n",
"\n",
"\n",
"def normalized_mse(predictions: np.ndarray, targets: np.ndarray) -> tuple[float, dict[str, float]]:\n",
" errors = (predictions - target_mean) / target_std - normalize_targets_raw(targets)\n",
" per_channel = np.mean(errors**2, axis=0)\n",
" return float(np.mean(per_channel)), {name: float(value) for name, value in zip(TARGET_NAMES, per_channel)}\n",
"\n",
"\n",
"def split_name_for_case(case_id: str) -> str:\n",
" for split_key, ids in split_manifest.items():\n",
" if case_id in ids:\n",
" return split_key.removesuffix(\"_ids\")\n",
" return \"unknown\"\n",
"\n",
"\n",
"def case_condition_frame(case_ids: list[str]) -> list[dict[str, float | str]]:\n",
" rows = []\n",
" for case_id in case_ids:\n",
" features = load_case(case_id, max_points=1)[\"features\"]\n",
" row = {\"case_id\": case_id, \"split\": split_name_for_case(case_id)}\n",
" for name in [\"u_inf\", \"log_re\", \"aoa_deg\", \"aoa_sin\", \"aoa_cos\", \"naca_param_3\"]:\n",
" row[name] = float(features[0, FEATURE_INDEX[name]])\n",
" rows.append(row)\n",
" return rows\n",
"\n",
"\n",
"example_case = split_manifest[\"val_ids\"][0]\n",
"print(f\"data root: {DATA_ROOT.relative_to(REPO_ROOT)}\")\n",
"print(f\"example: {example_case}\")\n",
"print(case_condition_frame([example_case])[0])\n"
]
},
{
"cell_type": "markdown",
"id": "5953b04d",
"metadata": {},
"source": [
"## 5. Prediction and residual maps\n",
"\n",
"A fast qualitative check: truth, prediction, and normalized residual for one validation case.\n"
]
},
{
"cell_type": "code",
"execution_count": 5,
"id": "adf0e581",
"metadata": {},
"outputs": [
{
"data": {
"image/png": "iVBORw0KGgoAAAANSUhEUgAABsMAAAQfCAYAAABGeZeuAAAAOnRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjExLjEsIGh0dHBzOi8vbWF0cGxvdGxpYi5vcmcvctoD+AAAAAlwSFlzAAAPYQAAD2EBqD+naQABAABJREFUeJzs3XdYFMf/B/D3HV2aimLHgl0Rey8o2Auxx14QNWoSS4xR0zWaaIz5Ro0FC3bFbuy9Yq+gWFCsiGClCAJ38/uD32047jhuj6bk/XoenofbnZ2dmd393N3O7YxCCCFARERERERERERERERElAcpc7sARERERERERERERERERNmFnWFERERERERERERERESUZ7EzjIiIiIiIiIiIiIiIiPIsdoYRERERERERERERERFRnsXOMCIiIiIiIiIiIiIiIsqz2BlGREREREREREREREREeRY7w4iIiIiIiIiIiIiIiCjPYmcYERERERERERERERER5VnsDCMiIiIiIiIiIiIiIqI8i51hRERERERERERERERElGexM4yIiIjoAxIeHg4vLy8cOHAgt4vywcjtNvn666/xxRdfZLiMiIiIiIiIiD5M5rldACIiIiL617t373D48GH0798/t4vywZDbJtHR0ejWrVu663ft2gVra2uj93/58mXExsZmuCwxMRE7duzA8ePH8ejRI5QoUQKenp7o3r07FAqFVtqff/4ZJ06cAAAoFArky5cPhQoVQvXq1dG5c2eUL1/e6PKldf/+fSxfvhyhoaFISkqCq6srOnXqhObNm8tOu3fvXsyZM8eo/W7atAkFChQwKu2DBw+wdOlS3LlzB8nJyShXrhyGDBmCatWq6aQNDAzEtm3bcPfuXdjb28PNzQ0+Pj5wcnLKcD/+/v5Ys2aN3nVbt26Fg4OD9FrO8dO4f/8+VqxYgaCgIDg6OqJTp07o2bOnUW2Q1vv377F48WIcP34cKpUK9evXx5gxY7TKmB45Ze/cuTPi4+P15qNQKLB//34olSm/mZTTfsbKzPEEgJcvX2Lr1q3YtWsX4uLisHTpUpQpUybTaVObM2cO9u7dC3d3d6PP/6zMT84xAuRdT/fu3cPixYsREhICa2trtGjRAr6+vrCysjKpbqa2sUZmzvvMlMeU6z2t7Ihj2XHNEREREVEKdoYRERER0QetRIkSOHjwoN4bjPokJibi8OHD6NSpE7788kud9RYWFrL2P3v2bCQnJxtMo1arUbJkSajVanz++edo06YNbty4gWHDhmH+/PnYu3cvbGxspPTXr1/H4cOHcfDgQQBAQkICnjx5gv3792PixIno3bs3Fi1aBHt7e1llXbx4McaMGYPWrVujT58+cHBwwNGjR9GmTRt4e3tj48aNstK6u7vjm2++0dpH69at0bhxY/z0009ay21tbY0q465du9CtWzc0bdoUw4cPh7m5OdavXw83NzcsWbIEw4YNk9KOHj0aoaGhaN++PZo3b47w8HD89ddf+OWXX3Do0CHUq1fP4L5CQ0Nx+PBh7NmzR+e4pz4eco8fAKxbtw6+vr7o378/Bg4cCCEEtmzZgitXrmDGjBlGtYXGu3fv0LJlS7x48QLff/89bGxs8Ntvv2HFihUIDAxE4cKF091WbtknTJigcz6Hh4dj0KBBqFOnjlYni7HtZ6zMHs8VK1ZgypQp6NSpE8zNzXH48GGdTmlT0qZ25swZTJ48GUKIDK97Y5iSn5xjJOd6OnToELy9vdG4cWMMGzYMb9++xfTp07Fq1SocPXrU6GtYw9Q21sjMeZ+Z8phyvaeVXXEsq685IiIiIkpFEBEREdEH4+7duwKAWLFiRW4X5aMVFRUlAIgRI0Zk2z48PT1FgwYNpNdJSUmiadOmIiIiQivd7t27BQAxa9YsreXdu3cX6X0U37dvn7CwsBDt2rWTVaa3b98Ka2tr0apVK511V69eFX379jUpbVoAhLe3t6yypVa7dm1RuHBhER8fLy1Tq9XCzc1NFClSRCvty5cvdbaPjIwUFhYWomPHjhnua+rUqQKA1r70kXv8rly5IiwsLMTChQt18nrz5k2G5Urrxx9/FEqlUty8eVNaFhUVJRwcHISPj0+Wll2fX375RQAQixcv1lpubPsZK7PH8/nz5yIpKUkIIcQPP/wgAIigoKBMp9WIjY0V5cuXF59//rlwcnISLVq0yLBMOZVfesfI2OspOTlZlCpVSlStWlUkJydLy2/duiXMzc3F5MmTZZfJlDZOLTPnfWbKkxXXTHbFsay+5oiIiIjoX5wzjIiIiMhEL168QOvWreHv76+zTqVS4ZNPPsEvv/wCIOXJHy8vL+mvY8eO+Oyzz3D8+PFMleH06dP44osv0KlTJ/j4+GDbtm3SOjn7PH/+PD7//HN4e3tj8ODBWLJkCRITE3XSnTx5EqNHj0bHjh3Rp08frFq1CiqVSivNhAkT4OXllWHZjS2fvjnD7ty5Ay8vL5w8eRI3btzAqFGj0K5dOwQHB2e439Tu3LmDiRMnonPnzujRowfmzJmDmJgYrTTGzA9mZmaGAwcOoEiRIlrLPT09AQDnzp0zukxt27bFmDFjsG/fPhw+fNjo7cLCwpCQkIAmTZrorHN3d8fSpUtNSpvV4uLiUKxYMa2hKhUKBcqUKYO4uDittAULFtTZvnDhwnBycsLr16+zrExyj9/06dNRrFgxjBgxQicvR0dH2ftfsWIFGjVqhCpVqkjLChUqBG9vb6xfvx4JCQlZVvb09m9nZ4e+ffvKLrscmT2ezs7OMDc3bnATOWk1JkyYgKSkJNlP9uVEfukdI2Ovp6CgIDx+/Bg9evSAmZmZtLxSpUqoVasWli1bBiGErDKZ0sapZea8z0x5suKa+RDjGBEREREZxs4wIiIiIhMVKlQI7969w6+//qqzbv/+/dixYweqVq0KIGVovm+++Ub6GzJkCJRKJVq1aoUVK1aYtP9JkyahWbNmePfuHQYMGIBmzZph7dq1UgecsfvcsWMHGjVqBIVCgSFDhqBNmza4evUqevXqpbW/KVOmoGXLllAqlfDx8UHjxo3x1VdfoUePHlCr1VK6S5cuGdWJY2z5NHOGhYeHS8uio6Nx+PBh7N27F19++SUaNGiApk2bIjo62uj227NnD9zd3XHt2jX07dsXXl5emDdvHmrVqoVnz55J6S5fvozz588bzEuhUOgdwiooKAgAZA/3pZl3ateuXUZvU7ZsWZibm+PIkSN6OzJTl09O2qw2aNAg3LhxA8eOHZOWBQcH49ixYxg4cGCG22/atAkRERGy5tUbMmQI2rdvjwEDBmDFihU6dZZz/FQqFfbv348WLVpg3759GDx4MDp06ICRI0ciMDDQ6DJpREVF4eHDh6hVq5bOulq1auHdu3cGO3kze+4dP34coaGh6NOnD+zs7PSmyaj9MsOU45kd9u3bhyVLlmDx4sXptkNu5WfoGBl7PWliY/78+XXyz58/PyIjIxEWFpapcsqR2fM+M7IiXmd3HMvOa46IiIjov4pzhhERERFlgo+PD3x8fHD69Gmtp2yWL18OZ2dndOrUCUDKL9HTPi3Vo0cPKJVKqSNIjt27d2PWrFn4/fffMWHCBGn54MGD8fbtW1n7XLFiBerVq4e//vpLSte3b1+8efNGen3gwAHMnDkTf/31Fz7//HNpeePGjVGvXj2sXbsWAwYMAADMmTPHqF+7Z0WbbN26FdeuXYOVlRUAICkpSar/rl27dPLv378/Bg8ejISEBAwZMgQ1a9bEvn37pDl4OnbsiIoVK2L8+PFYv359hvs3RK1WY+LEiQBSbpzKUbFiRQDA/fv3jd7GwcEBv/32G7766itUqFAB3bp1Q506ddCkSROULVvW5LRZbfLkyXB2dsYnn3widcrdvHkT3377LSZNmqR3m169euHFixd48uQJXr16hfXr1+PTTz/NcF8KhQLdunVDq1at4OTkhPPnz2P06NGYN28eDh8+jAIFCqS7bXrHLyIiArGxsTh48CB2796N7777DqVKlcKGDRvQtGlTzJ8/H6NGjTK6PTSdvPpuwDs7OwOAVuesMeS
"text/plain": [
"<Figure size 1050x1040 with 24 Axes>"
]
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"def plot_case_prediction(case_id: str, *, max_points: int = CASE_SAMPLE_POINTS) -> dict[str, object]:\n",
" case = load_case(case_id, max_points=max_points)\n",
" features = case[\"features\"]\n",
" targets = case[\"targets\"]\n",
" predictions = predict_targets(features)\n",
" loss, per_channel = normalized_mse(predictions, targets)\n",
" x = features[:, FEATURE_INDEX[\"x\"]]\n",
" y = features[:, FEATURE_INDEX[\"y\"]]\n",
" target_norm = normalize_targets_raw(targets)\n",
" pred_norm = (predictions - target_mean) / target_std\n",
" residual_norm = pred_norm - target_norm\n",
"\n",
" fig, axes = plt.subplots(len(TARGET_NAMES), 3, figsize=(10.5, 2.6 * len(TARGET_NAMES)), sharex=True, sharey=True)\n",
" for row, name in enumerate(TARGET_NAMES):\n",
" index = TARGET_INDEX[name]\n",
" values = [targets[:, index], predictions[:, index], residual_norm[:, index]]\n",
" titles = [f\"truth {name}\", f\"prediction {name}\", f\"normalized residual {name}\"]\n",
" for col, (ax, values_col, title) in enumerate(zip(axes[row], values, titles)):\n",
" cmap = \"coolwarm\" if col == 2 else \"viridis\"\n",
" scatter = ax.scatter(x, y, c=values_col, s=4, cmap=cmap)\n",
" ax.set_title(title, fontsize=9)\n",
" ax.set_aspect(\"equal\", adjustable=\"box\")\n",
" ax.set_xticks([])\n",
" ax.set_yticks([])\n",
" fig.colorbar(scatter, ax=ax, fraction=0.045, pad=0.02)\n",
" fig.suptitle(f\"{split_name_for_case(case_id)} case: {case_id}\\nnormalized MSE={loss:.4g}; per-channel={per_channel}\", y=1.01)\n",
" fig.tight_layout()\n",
" return {\"case\": case, \"predictions\": predictions, \"loss\": loss, \"per_channel\": per_channel, \"residual_norm\": residual_norm}\n",
"\n",
"\n",
"PREDICTION_RESULT = plot_case_prediction(example_case)\n"
]
},
{
"cell_type": "markdown",
"id": "98f0cac6",
"metadata": {},
"source": [
"## 6. Where failure lives spatially\n",
"\n",
"Aggregate pointwise normalized residual by spatial coordinates and by signed-distance-like `sdf`. Quantile bins avoid assuming a particular distance scale.\n"
]
},
{
"cell_type": "code",
"execution_count": 6,
"id": "3eec699a",
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"near-wall-ish sdf <= 0.001306: mean RMSE 0.3099\n",
"farfield-ish sdf >= 18.48: mean RMSE 0.2385\n",
"worst 5 sampled points: [1.6673461 1.7071781 1.7142189 1.7244339 1.7760984]\n"
]
},
{
"data": {
"image/png": "iVBORw0KGgoAAAANSUhEUgAABKQAAAFICAYAAABunSaOAAAAOnRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjExLjEsIGh0dHBzOi8vbWF0cGxvdGxpYi5vcmcvctoD+AAAAAlwSFlzAAAPYQAAD2EBqD+naQAA6eJJREFUeJzs3XdcE/cbB/BPCHtvQUBAnIh74UDce+DWonVVEXdb909FrBW31r1HxbqLW+rAhRNFcOAGFWUP2SMk9/uDkhKSQBICSeB5v16+JJfvffPc5ZLLPfcdLIZhGBBCCCGEEEIIIYQQUknUFB0AIYQQQgghhBBCCKleKCFFCCGEEEIIIYQQQioVJaQIIYQQQgghhBBCSKWihBQhhBBCCCGEEEIIqVSUkCKEEEIIIYQQQgghlYoSUoQQQgghhBBCCCGkUlFCihBCCCGEEEIIIYRUKkpIEUIIIYQQQgghhJBKRQkpQkiV8OrVK+zatQs5OTkVvu6rV69w9OhR7Nq1C9HR0eV6nfLEXZFOnjyJK1euSFT23LlzOH/+fAVHJBqXy8WNGzewf/9+HDx4UCExlKTI/fH+/Xvs2rUL379/L3UZIYQQQgghiqau6AAIIarn69eveP78ORITE1GrVi00a9YMJiYmFf664eHhePDgAcaPHw9tbW2B527evImZM2fCw8MDOjo6UtUrzbq///47fv/9dwwePBgGBgbo1KlTuV6nPHFXpBUrVsDBwQF9+vQps+y6deugrq6OgQMHVkJk/8nPz0eHDh2QkpKCLl26wNDQsFJfXxxp98eVK1fw+fNnkc+5ubmhUaNGEr/2o0eP4O3tjc6dO8PY2FjsMgBgGAbPnz/H69evoaWlBRcXF9StW1eozi9fvuDy5cv8x5qamjAyMkKdOnXg7OwMDQ0NieMTJS4uDmFhYYiPj4eNjQ2aNWsGc3NzqcsyDIPdu3dL9Jqurq5o1qyZVDE+ffoUSUlJsLKyQtu2bQX2ZXFRUVEIDw9HdnY27O3t4erqCjabXeZrpKen46+//hL5XLt27dC0aVOBZZK+f8UlJiYiODgY2dnZaNmyJRo0aFBmXOLk5ubi5s2biIuLg729PTp16gR1dcl+Vkoae3BwMF6+fCm2ns6dO/O3Qdr9JylZ388iBQUFCAoKQmRkJJo1awZXV1e5lC2SkZGBv/76CwzDwNPTEwYGBhLHJo/6pHmPikjzeYqJieEfsw0bNkSbNm3AYrGk3q4isuzj4spz3JcnHlk+7yXJ+3usoj5zhJDqhRJShBCJZWdnY/r06fD390e7du3g6OiId+/eISwsDMOHD8e+ffugqalZYa9/7do1zJs3D8OGDRNKSLm4uMDLywu6uroV9vq5ubn47bffsGDBAvj6+kq9fmXEqAgeHh5QU6v8BrdnzpzBkydPEBERgYYNG1b664sj7f7Yvn07AgMD8dNPPwk9J23CoF69evDy8iozQXzixAksXboUSUlJ6NmzJ3JycvDPP/+gX79+2Ldvn8D6z58/5ye06tevDx6Ph+TkZDx58gSZmZn46aefsGzZMujp6UkVa35+Pn7++Wfs27cPLVq0QL169RAVFYVHjx7Bw8MD+/bt418MS1JWT08PYWFhAq9x8+ZNvHv3DhMmTBD4bnJ0dJQ4znnz5mHLli1o3bo1nJyc8Pz5c3z8+BG//fYbZs+ezS/38eNHTJkyBd++fUOTJk0AALdu3YKBgQH27t2Lrl27lvo6CQkJ8Pb2hpubG5ydnQWeq1OnjsBjad4/AODxeFi0aBG2b98ONzc3WFtbY/369WjQoAGOHTsm8b4oEhISgsGDB8PExAQtW7bE3bt3oa2tjUuXLsHBwaHUdaWJ/du3b0LvKVCYxC1KlBZ9RqTZf5Io7/uZnZ2N2bNn4+LFi3B0dMSDBw/w66+/ikw4SFO2pJ9//hn79+8HAPTu3bvcCSlp65PmPQIk/zwBgI+PD1atWoX27dvD2toa8+bNQ7169XDu3DmxSWtxyrOPi5TnuC9PPNJ+3kWpiO8xeX/mCCHVFEMIIRKaNWsWo6amxly+fFlg+dWrVxlLS0smNTW1Ql9/3bp1DAAmMTFRrvVu3bqVAcDExsaWWu7t27cMAObQoUOV/tqVrVGjRky/fv0UHUapfH19GQBMbm6uokMpl379+jFaWloVVv+RI0cYAMzr16/5yzw9PZlJkyYx6enp/GVhYWGMrq4uM3jwYIH1L1y4wABgDh48KLCcx+Mxp06dYgwNDZlWrVoxmZmZUsW1bNkyBgBz/PhxgeXBwcGMnZ0dExUVJVPZ4saNG8cAkPm76cyZMwwAxtfXV6heFovFhIWF8Ze9efOGefz4sUC59PR0plGjRoyJiQmTnZ1d6mu9f/+eAcDs3LmzzLikef8YhmFmz57NWFpaMq9evRJYfu7cuTJfq6SMjAzG2tqa6dmzJ8PhcBiGYZisrCymSZMmTIsWLRgulyvX2EvKzs5mjIyMGBsbG6agoIC/XJr9J4nyvp/p6enMnj17mKSkJCYqKooBwPz666/lLlvc+fPnGQ0NDaZHjx4MALGfA0nJqz5x75E0nyd/f38GALN+/Xr+spiYGKZmzZpMr169pI5J1n1cpLzHfXniKe9npqK+x+T9mSOEVE/UQooQIrGLFy+ibt26Qt24evTogbt370JLS4u/LDQ0FI8fP8ZPP/2EzMxMXL16FRwOB507d4aNjY3A+pGRkbh69SoAgMViQU9PD40aNULz5s35ZYKDg/Hw4UMAwOHDh/mtMTw8PGBlZYVXr17h7t27GDduHL/rmyT1SiooKAhBQUEAgNu3byMnJwcGBgbw9PSU+HVExSjK3bt38eHDB0yYMEHkfho6dCgsLCwACO7nnJwcXL9+HampqZg4cSJ/vTdv3iA0NBQcDgfNmjWTuhl9REQEnjx5AjMzM3Tr1k2oddq5c+fAYrEEuqidPHkSBgYG6NOnD96/f48HDx7A2NgYPXr0ELnt7969w/Pnz5Gfnw9nZ+cyu1Pt2bMH9+/fBwDs378fampqaN68Odq2bQugsHvDw4cP8fbtW2hra6Ndu3awt7cXqKNkjI8ePYKdnR3c3d1Fvua1a9fw8eNHAACbzYapqSn/zr20+6Os1xLn27dvuHfvHrKzs1GnTh20b99eoDXW+/fvcePGDYwaNUpsVwwA+PXXX4WOz6ZNm6Jnz544f/488vPzy2ztyGKxMGzYMOTl5WHMmDHYsGEDli1bJvG2XLx4ETVq1MDIkSMFlnfo0AHBwcECLTOkKStPL168AAAMGjRIYLmHhwcOHz6Mly9f8j9P9evXF1rfwMAAgwYNwqpVqxAZGSlV98vSSPP+RUREYMuWLdizZ49QKwZZutn+9ddfiI2NxcmTJ/ldlXR1dTFv3jyMHTsWN2/eRLdu3eQSuyhnzpxBWloaZs+eLVXXOWmV9/00MDDA5MmTARR2gyuNNGWLJCYmYvLkyViwYAFycnJw7do1idarjPrEvUfSfJ727dsHY2NjgdY71tbW8Pb2xtKlSxEWFiZVt1tZ9nFx5T3uyxNPeT8zyvo9RgghAA1qTgiRgrGxMVJTU5Gfny/0XL169QQSDZcvX4a3tzdu3bqFTp06ISAgAJs2bULt2rWFxnlJS0tDWFgYwsLCEBoailOnTqFjx45wd3fnD/b97ds3fPv2DUDhj6ui8llZWQAKu+Z4e3sjLS1Nqnol9fnzZ7x9+5b/d1hYGCIiIqR6HVExinL06FH8/PPPQstDQ0Ph7e0tMN5Q0X6+c+cO3NzccPz4caxevRoA8P37dwwaNAgtWrTA0aNHcfHiRbi7u2PIkCESb7+vry8mTpyIy5cvY/LkyWjYsCHevHkjUGbdunXYuHGjwLIVK1Zg+/bt2LBhA3788UcEBgZi8uTJcHFxQXx8PL8cwzCYNGkSmjdvjiNHjvDLderUCampqWLjCg8P59cTHh6OsLAwxMbGAigc86hVq1bo378/Ll26hD179qBOnTqYM2cOGIYRinHlypUYM2YM/v77bxw/flzsa0ZFRfHf5wcPHmDDhg2ws7PDunXrJN4fkr6WKIsWLYKDgwO
"text/plain": [
"<Figure size 1200x330 with 3 Axes>"
]
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"def binned_mean(x_values: np.ndarray, y_values: np.ndarray, *, bins: int = 12) -> tuple[np.ndarray, np.ndarray]:\n",
" edges = np.quantile(x_values, np.linspace(0, 1, bins + 1))\n",
" edges = np.unique(edges)\n",
" centers, means = [], []\n",
" for low, high in zip(edges[:-1], edges[1:]):\n",
" mask = (x_values >= low) & (x_values <= high if high == edges[-1] else x_values < high)\n",
" if mask.any():\n",
" centers.append(float(0.5 * (low + high)))\n",
" means.append(float(np.mean(y_values[mask])))\n",
" return np.asarray(centers), np.asarray(means)\n",
"\n",
"\n",
"case = PREDICTION_RESULT[\"case\"]\n",
"features = case[\"features\"]\n",
"residual_norm = PREDICTION_RESULT[\"residual_norm\"]\n",
"point_rmse = np.sqrt(np.mean(residual_norm**2, axis=1))\n",
"sdf = features[:, FEATURE_INDEX[\"sdf\"]]\n",
"x_coord = features[:, FEATURE_INDEX[\"x\"]]\n",
"y_coord = features[:, FEATURE_INDEX[\"y\"]]\n",
"\n",
"fig, axes = plt.subplots(1, 3, figsize=(12, 3.3))\n",
"for ax, values, label in [(axes[0], sdf, \"sdf\"), (axes[1], x_coord, \"x\"), (axes[2], y_coord, \"y\")]:\n",
" centers, means = binned_mean(values, point_rmse)\n",
" ax.plot(centers, means, marker=\"o\")\n",
" ax.set_xlabel(label)\n",
" ax.set_ylabel(\"mean normalized point RMSE\")\n",
" ax.grid(alpha=0.25)\n",
"fig.suptitle(f\"Spatial failure bins for {case['case_id']}\")\n",
"fig.tight_layout()\n",
"\n",
"near_cut = np.quantile(sdf, 0.20)\n",
"far_cut = np.quantile(sdf, 0.80)\n",
"print(f\"near-wall-ish sdf <= {near_cut:.4g}: mean RMSE {point_rmse[sdf <= near_cut].mean():.4g}\")\n",
"print(f\"farfield-ish sdf >= {far_cut:.4g}: mean RMSE {point_rmse[sdf >= far_cut].mean():.4g}\")\n",
"print(f\"worst 5 sampled points: {np.sort(point_rmse)[-5:]}\")\n"
]
},
{
"cell_type": "markdown",
"id": "715e7829",
"metadata": {},
"source": [
"## 7. Condition sensitivity sweeps\n",
"\n",
"Hold sampled coordinates/geometry fixed and perturb case-level conditions. This is a local probe, not a physically valid new CFD case.\n"
]
},
{
"cell_type": "code",
"execution_count": 7,
"id": "864622fd",
"metadata": {},
"outputs": [
{
"data": {
"image/png": "iVBORw0KGgoAAAANSUhEUgAABKYAAAGGCAYAAABBiol3AAAAOnRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjExLjEsIGh0dHBzOi8vbWF0cGxvdGxpYi5vcmcvctoD+AAAAAlwSFlzAAAPYQAAD2EBqD+naQABAABJREFUeJzs3Xd8U9X7wPFPku7dQhdtocjeu2xkW/YQlaUMlSGCgigbRPEHOL8Kioiyp7Jk77333qO0QAsFuqAzyf39URobuqFtOp7395XXl56ce+5zcnPNzZNzzlUpiqIghBBCCCGEEEIIIUQuU5s6ACGEEEIIIYQQQghROEliSgghhBBCCCGEEEKYhCSmhBBCCCGEEEIIIYRJSGJKCCGEEEIIIYQQQpiEJKaEEEIIIYQQQgghhElIYkoIIYQQQgghhBBCmIQkpoQQQgghhBBCCCGESUhiSgghhBBCCCGEEEKYhCSmhBBCCCGEEEIIIYRJSGJKCFFo2NnZ8fHHHxv+Dg8PR6VS8f3332fbPnKiTSGEEEKIVzFo0CCcnJyMytq3b0/16tWzdT850WZBlh2v1+LFi6lQoQIWFhaoVKrsCUyIXCaJKSGEyKKQkBBUKhX/+9//TB2KyURFRWFnZ4dKpWLz5s3Z1u6///6LSqXCwcGBZ8+eZVu7QgghhMgeLVu2pHbt2qYOI9dVrlyZ8uXLp/pcbGwsKpWK7t2752pMV69epW/fvvTv359nz56hKEqu7l+I7GJm6gCEEMJUnJycsv0DPCfazIuWLl1KdHQ0xYoVY86cObRp0yZb2v3zzz/x8vLi3r17rFixgv79+2dLu0IIIYQwtmHDhnzRZkH2qq/XgQMH0Ol0dO/eHXNz82yKSojcJyOmhBBCZNmcOXNo3bo1I0eOZP369Tx48OCV27x37x6bN29m9OjRNG/enDlz5mRDpEIIIYQQBdOjR48AsLa2NnEkQrwaSUwJkcc0atSIRo0acfPmTVq3bo2trS2lS5dm9erVAFy7dg1/f3/s7Ozw9vZm1qxZqbaze/duWrVqhaOjI1ZWVtSuXdvQRpIPPvgAlUqFSqVCrVZTpEgROnTowOnTp1ONKSgoiPbt22Nra4u7uztffPEFOp0u0326fv06rVq1wtbWlmLFijFq1Cji4+NTrXvr1i3atWuHvb093bp1A0Cr1TJ9+nQqVaqElZUVLi4uvP3229y6dcuojdDQUN59912cnZ1xcnKiV69ehIeHp4grrfWgtFot3333HVWrVsXa2hovLy/69evH3bt3OXLkCJ6engAMHz7c8PoNGjQo3TYjIyP59NNPKV68OBYWFnh7ezNkyBCePHliqBMQEIBKpeL3339nzZo1hn5WrlyZTZs2pfsaR0dH4+TkRO/evVM8p9Vq8fDw4M033zSUzZw5k6pVq2JnZ4enpyedOnXi5MmT6e4jyZkzZzh58iQfffQR/fr1w8LCgvnz56da98KFC3Tp0gUXFxdDX3755ZdUR5XNnTsXa2tr3nvvPT766COOHDnChQsXMhWTEEKI/O/3339HpVJx48YNvvjiC9zc3HB2dmbIkCFotVq0Wi2ff/45Hh4e2Nra0qNHD6KiolK0ExISwqBBg/D29sbCwoLixYvzxRdfEBsba6izZ88ew2e4SqXCxsaGmjVrMnv27FRjunXrFpMmTcLT0xMbGxveeOONFNcfGfXp888/x83NDVtbW9q1a8e1a9fSrDtu3DiKFSuGSqUyxJ2ZazuA//3vf5QqVQpra2vq1KnDgQMHUo0trfWNdu/ejb+/Py4uLtjb29OkSRO2bdsGQOnSpdm5cycnT540vHZ2dnYZtrlo0SJq1aqFtbU1Dg4OtG7dmsOHDxvVeZXrzX79+mFnZ8fTp09TPDds2DCsrKwICwsD4OjRo7Rp0wY3Nzfs7e2pXbs2f/75Z6auabNbaq9XZl+HokWLMnr0aABcXV2NrkeFyHcUIUSe0rBhQ6V69epK586dlVOnTilhYWHKqFGjFDMzM2Xfvn3KG2+8oZw4cUIJDw9Xxo0bpwDKkSNHjNpYunSpolarlWHDhikBAQHKkydPlJ9//lnRaDTKwoULU91vfHy8cvHiRaVz586Kq6ur8uDBA6OYatSooXTt2lU5cuSIEhkZqfz111+KSqVSZsyYkek++fv7K0ePHlUiIiKU5cuXK3Z2dkrPnj1TrdumTRvl0KFDSmhoqLJ06VJFr9crHTt2VFxcXJSlS5cqYWFhyo0bN5R27dopHh4eSnBwsKIoihIbG6tUrVpVKVGihLJ3714lIiJCWbdunfLOO+8otra2ypAhQwz7CgsLUwDlu+++M5TpdDqlbdu2iqOjo/LXX38pDx48UIKDg5X58+cr48ePVxRFUYKDgxVA+emnn1L0NbU24+LilNq1ayseHh7K5s2blYiICGXXrl2Kj4+PUrlyZeXZs2eKoijK7du3FUB56623lCFDhiiBgYFKcHCw0qVLF8XKykq5f/9+uq/zoEGDFCsrKyUsLMyofM2aNQqgbNy4UVEURfnzzz8VMzMzZcmSJUpkZKTy6NEjZcOGDUr37t0zOJKJPvroI6V48eKKVqtVFEVR+vfvr5QuXVrR6/VG9S5evKjY2dkpr7/+unLp0iXl8ePHysyZMxVzc3Nl6NChRnX1er3i6+urDBw4UFEURUlISFC8vLyUTz75JFMxCSGEyP9mzZqlAEqvXr2UBQsWKOHh4cqOHTsUOzs7Zdy4ccqwYcOUefPmKeHh4cru3bsVBwcH5eOPPzZqIzg4WPHx8VFq1qypHDp0SImKilL27dunvPbaa4q/v3+a+3748KHhWmnRokUpYurbt68yZ84c5cmTJ8qZM2eU1157Talbt26m+/TOO+8of/zxhxIWFqacO3dOqVOnjuLm5ma4fklet3v37srvv/+uPHr0SJk/f74SGxub6Wu7r776StFoNMoPP/ygPHr0SLl8+bLSrl07pXXr1oqjo6NRbO3atVOqVatmVLZw4UJFpVIpH3zwgXL16lUlMjJS2b9/v9KuXTtDnRYtWii1atVKtb+ptTl9+nRFpVIpX375pfLgwQPl5s2bSrdu3RRzc3Nl9+7dhnqvcr154MABBVD+/PNPo/LY2FilSJEiSo8ePRRFSTzODg4OSp8+fZTAwEAlOjpaOX36tDJgwADl+PHj6e6jUqVKSrly5VJ9LiYmxnCcsyK11ysrr8PUqVMVQAkNDc3SfoXIayQxJUQe07BhQ0WlUimXLl0ylMXFxSlOTk6Kvb29cv78eUN5fHy84uLionz44YeGsujoaKVIkSJK27ZtU7Tdt29fxdPTU9HpdGnuPzIyUlGr1crvv/+eIqYLFy4Y1W3WrJlSpUqVTPUJUE6fPm1UPm3aNAVQTp06laLuiRMnjOquWrVKAZRly5aliNfFxUX57LPPFEVRlL/++ksBlO3btxvV+/PPPxUgw8TU8uXLFSDNBJ6iZD0xNXfuXAVQVq9ebVR3+/btRu0kJaZq1qxpVO/u3bsKoEyfPj3NmBRFUU6cOKEAyq+//mpU3r59e8XLy8uQSOrdu7dSpkyZdNtKS3R0tOLo6Kh88803hrKTJ08qgLJr1y6jul27dlXs7OyUx48fG5V/8sknikqlUq5du2Yo27p1qwIoZ8+eNZR9+eWXiouLixIbG/tSsQohhMhfkhIzST8EJRkwYIBiY2OjjBkzxqh88ODBio2NjdEPIx9++KFiY2Oj3L1716jujh07Ur0+eFHnzp2VRo0apYhp7NixRvWSrjeSf26l16ek65Qkt2/fVszMzJQRI0akqDt8+HCjupm9tgsLC1Osra2VPn36GNV59OiRYmtrm2Fi6unTp4qTk5PSqlWrdPuUlcRUWFiYYmNjo3Tt2tWoXlxcnOLj42PUzqteb5YvX16pX7++UVnSdd2OHTsURVGULVu2KIBy9OjRDNt7UW4mpjL7OkhiShQUMpVPiDzI19eXChUqGP62sLDgtddew8nJicqVKxvKzc3NKV26tNFQ8kOHDvH48WPefvvtFO22bNmS4OBgrl+/Dvw31N3X19dwi1kHBwf0ej03btx
"text/plain": [
"<Figure size 1200x400 with 2 Axes>"
]
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"NU_DATASET = 1.56e-5\n",
"\n",
"\n",
"def with_condition_updates(features: np.ndarray, *, aoa_deg: float | None = None, u_inf: float | None = None) -> np.ndarray:\n",
" updated = np.array(features, copy=True)\n",
" if aoa_deg is not None:\n",
" radians = math.radians(float(aoa_deg))\n",
" updated[:, FEATURE_INDEX[\"aoa_deg\"]] = float(aoa_deg)\n",
" updated[:, FEATURE_INDEX[\"aoa_sin\"]] = math.sin(radians)\n",
" updated[:, FEATURE_INDEX[\"aoa_cos\"]] = math.cos(radians)\n",
" if u_inf is not None:\n",
" updated[:, FEATURE_INDEX[\"u_inf\"]] = float(u_inf)\n",
" updated[:, FEATURE_INDEX[\"log_re\"]] = math.log(float(u_inf) / NU_DATASET)\n",
" return np.ascontiguousarray(updated, dtype=np.float32)\n",
"\n",
"\n",
"def sweep_condition(features: np.ndarray, *, values: list[float], field: str) -> np.ndarray:\n",
" rows = []\n",
" for value in values:\n",
" counterfactual = with_condition_updates(features, aoa_deg=value) if field == \"aoa_deg\" else with_condition_updates(features, u_inf=value)\n",
" pred = predict_targets(counterfactual, batch_size=INFER_BATCH_SIZE)\n",
" rows.append([value, *pred.mean(axis=0), *pred.std(axis=0)])\n",
" return np.asarray(rows, dtype=np.float64)\n",
"\n",
"\n",
"base_features = PREDICTION_RESULT[\"case\"][\"features\"][:FAST_SAMPLE_POINTS]\n",
"base_aoa = float(base_features[0, FEATURE_INDEX[\"aoa_deg\"]])\n",
"base_u = float(base_features[0, FEATURE_INDEX[\"u_inf\"]])\n",
"aoa_values = [base_aoa + delta for delta in [-8, -4, 0, 4, 8]]\n",
"u_values = [base_u * factor for factor in [0.75, 0.9, 1.0, 1.1, 1.25]]\n",
"\n",
"aoa_rows = sweep_condition(base_features, values=aoa_values, field=\"aoa_deg\")\n",
"u_rows = sweep_condition(base_features, values=u_values, field=\"u_inf\")\n",
"\n",
"fig, axes = plt.subplots(1, 2, figsize=(12, 4))\n",
"for target_offset, name in enumerate(TARGET_NAMES, start=1):\n",
" axes[0].plot(aoa_rows[:, 0], aoa_rows[:, target_offset], marker=\"o\", label=name)\n",
" axes[1].plot(u_rows[:, 0], u_rows[:, target_offset], marker=\"o\", label=name)\n",
"axes[0].axvline(base_aoa, color=\"black\", linewidth=1, alpha=0.5)\n",
"axes[0].set_xlabel(\"counterfactual AoA [deg]\")\n",
"axes[0].set_title(\"mean prediction vs AoA\")\n",
"axes[1].axvline(base_u, color=\"black\", linewidth=1, alpha=0.5)\n",
"axes[1].set_xlabel(\"counterfactual U_inf [m/s]\")\n",
"axes[1].set_title(\"mean prediction vs U_inf\")\n",
"for ax in axes:\n",
" ax.grid(alpha=0.25)\n",
" ax.legend(fontsize=8)\n",
"fig.tight_layout()\n"
]
},
{
"cell_type": "markdown",
"id": "9fae716e",
"metadata": {},
"source": [
"## 8. FiLM mechanism inspection\n",
"\n",
"Extract condition embeddings and per-block FiLM modulation. If the condition path is doing useful work, embeddings/modulation should vary smoothly with case variables such as AoA and freestream speed.\n"
]
},
{
"cell_type": "code",
"execution_count": 8,
"id": "622651d1",
"metadata": {},
"outputs": [
{
"data": {
"image/png": "iVBORw0KGgoAAAANSUhEUgAABQoAAAFyCAYAAAC9Y+fnAAAAOnRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjExLjEsIGh0dHBzOi8vbWF0cGxvdGxpYi5vcmcvctoD+AAAAAlwSFlzAAAPYQAAD2EBqD+naQABAABJREFUeJzs3Xd8U1UbB/DfvUmb7r0HHYxCW0Zp2bKnsgREAYEXBERAVEBRcSCCiAMHDkBEZA/ZIEtAQJFd9oZOuvdOmnHeP0oCadI2aZM2TZ/v51NKzj333nNy0zuenMExxhgIIYQQQgghhBBCCCENGl/XBSCEEEIIIYQQQgghhNQ9ChQSQgghhBBCCCGEEEIoUEgIIYQQQgghhBBCCKFAISGEEEIIIYQQQgghBBQoJIQQQgghhBBCCCGEgAKFhBBCCCGEEEIIIYQQUKCQEEIIIYQQQgghhBACChQSQgghhBBCCCGEEEJAgUJCCCGEEEIIIYQQQggoUEgIIYQQQgghhBBCCAEFCuuVv/76C6dPn1ZL27t3L65cuaLzNvTNX19t374dN27cqJV9aTsu+uRtKMeEEEIIIYQQQgghpo1jjLG6LgTRTXh4OAIDA7F//35VmpOTE8aOHYsff/xRlbZr1y4EBwejdevWGtvQlt8ccRyHd999F0uWLDH6vrQdF33y1vUxOXjwIIqKigCUvW/W1tZo1qwZmjRpUuE62dnZuHbtGgoKCuDt7Y3GjRvD2dm5wvwnT55ERkYGunbtCk9PT4PXgRBCCCGEEEIIITUnrOsCkJoZOnQoIiIi1NJGjx6N1157Dd99951O+UndqutjMm3aNOTn56NXr14AgNTUVJw9exadOnXC5s2b4efnp8p7584dzJ49G0ePHkWbNm3g5eWFR48e4fr163j22WexZMkShIaGqm0/Pz8fzz33HIqLi/H222/jq6++qtX6EUIIIYQQQgghRDcUKKzn1q5da9T8xPhM4ZgEBwdj+/btqtenTp1Cr169MG7cOPz9998AgOjoaPTo0QMRERG4f/8+AgICVPlv3bqFSZMm4dSpUxqBwk2bNkEsFuPZZ5/FunXrsHjxYlhYWNROxQghhBBCCCGEEKIzChTqICEhAbdu3YK1tTUiIiLg4OCgtpwxhmvXriE+Ph7Ozs6IioqCtbW1Wp7t27ejefPmCA8PR2xsLG7evAlfX98KW5KVlpbi7NmzKC4uRrt27eDq6qo13969e9GoUSO0adMGEokE+/btg0KhwIMHD1SBHy8vLzzzzDMa+Y1dh4rExsbi1q1bsLCwQFRUFFxcXCrcz4MHD3Dv3j2EhoYiMDBQlefBgwe4c+cOGjdujBYtWlS6v/v37+Pu3bsICAhAy5Ytq10uQPfjok9ebcfEEJ+Xv/76CzY2NujSpUuFZaxIt27d0K1bN5w4cQL5+fmws7PDyy+/DEdHR+zbt0/jbyA0NBTHjx/H3bt3Nba1atUqDBgwAJ9//jlat26NPXv24IUXXtC7TIQQQgghhBBCCDEyRiqUnJzMBgwYwEQiEevQoQPr0aMH8/PzYytWrFDluXnzJgsPD2dOTk6sT58+LCgoiLm4uLCNGzeqbQsAe/fdd9l7773HIiIiWI8ePZiFhQUbNGgQk8vlankvXbrEGjVqxNzc3Fjfvn1Zs2bN2I4dO1hYWBgbOHCgWl5HR0c2Y8YMxhhjOTk5bMSIEYznedakSRM2YsQINmLECLZo0SKt+Y1ZB21SUlJY//79mY2NDevevTvr2LEjs7a2ZosXL9a6nzfeeIO1a9eOdenShQmFQvbDDz8whULBZsyYwaKiotgzzzzDOI5j8+bN09iXchszZsxgERERrFu3bkwkErEBAwaw/Pz8apVLn+NS3WNYnff6woULzM/Pj7m7u7N+/fqxkJCQCvelTUBAAIuMjNRIHzJkCAPAUlJS2JEjRxgAtnDhwiq397To6GgGgO3fv58xxliXLl1Yv3799NoGIYQQQgghhBBCagcFCitQWlrKWrZsyYKDg9mdO3dU6QUFBWzLli2MMcaKioqYv78/a9myJUtPT2eMMSaXy9nUqVMZz/Ps33//Va0HgDVr1oytXr1albZ7924GQC0gV1hYyHx9fVnHjh1Zbm4uY4wxiUTCxo4dyzw9PXUKMolEIvbmm29qrVf5/MaogzZSqZS1adOGNWnShMXGxqrS9+7dywCwzZs3q+2nSZMmbP369aq0Tz75hFlYWLD333+f/fbbb6r0hQsXMp7n2b1799T2B4A1btyYrVy5UpV26dIlZmdnx/73v//pXS59joshjqGu73VBQQHz9vZmnTt3Znl5eYyxss/uxIkTte5LG22Bwry8PObu7s58fX2ZQqFgn3zyCQPAjh49WuX2njZt2jQWGBioCm5u3LiR8TzP4uLi9NoOIYQQQgghhBBCjI8ChRXYunUrA8D++OOPCvOsXbuWAWD79u1TS8/Pz2cODg5sxIgRqjQArF27dhrb8PPzY2PHjtXY5t9//62WLzY2lgEweKDQGHXQZvv27QwA27Fjh8ayvn37so4dO6rtp3379mp5EhISGADWtm1btfTk5GQGgP3www9q6QBY69atNfb11ltvMYFAwDIyMvQqlz7HxRDHUNf3+vfff2cA2IkTJ9TyJSQkMI7jdA4UBgcHsz/++IP98ccf7Mcff2StW7dmlpaWqvdl5syZDAC7fPlyldtTKioqYo6Ojuzzzz9XpUkkEubh4cE++ugjnbdDCCGEEEIIIYSQ2kFjFFbg3LlzAICuXbtWmOfy5csAgA4dOqil29vbIywsDNHR0WrpkZGRGtvw9/dHYmKi6rVynaioKLV8gYGB8PDw0KMGujFGHbQ5c+YMACA3Nxe7d+8GKwtSAwAsLCxw7do1tfxt27ZVe+3j4wOO4zTSvby8IBAI8OjRI419ln8PAaB9+/aQy+W4du0aevXqpXO59DkuhjqGurzXyuNXfl/+/v567SsnJwdbtmwBx3GwsrLC8OHDMWbMGDRp0gQAYGNjAwAoKSnReZvbtm1DXl4eHB0d1SZKad26NX777TfMnz8fAoFA5+0RQgghhBBCCCHEuChQWAGxWAwAsLW1rTBPaWlphXlsbW0hkUjU0hwdHTXyWVpaqval3CbP87CystLIqwzWGJIx6qCNMsD0559/guM4tWXW1tZ47rnnKt2PQCAAz/Ma6RzHQSgUat2/tvdLmaasl67l0ue4GOoY6vN5EYlENdpX+VmPywsLCwNQNrtxp06ddNrmr7/+itatW+PYsWNq6Q4ODsjOzsahQ4cwcOBAnctICCGEEEIIIYQQ46JAYQWaNm0KALh9+zbatWunNY9yFt779++jdevWasvu37+PoKAgvfcbFBQEhUKB2NhYVRmAssBlUlKSKmBTmfIBr8oYow7aNG7cGACwcOFChIaGGmSbVXn48KFG2oMHDwBAVS9dy6XPcTHEMdRVYGAgFAoF4uLiVK3/gLIAYlJSksHe64EDB8LOzg4bNmzApEmTKsxXUFAAe3t73L59G6dPn8bff/+NHj16aOQbNGgQVq1aRYFCQgghhBBCCCHEhPB1XQBT9dJLL8HGxgaLFi2CQqFQW5aWlgYAGD58OAQCAZYtW6a2fNeuXYiPj8eoUaP03u+wYcMgEAjw008/qaWvXLlS526arq6uyMvL0ymvMeqgzejRo2Fra4vFixdrXV5V1+Xq+Oeff1SBQaCsFeGqVavQqlUrNG/eXK9y6XNcDHEMdaXc1/Lly9XSV69eDZ433J+3i4sLFi1ahBMnTmDp0qUayxlj+OKLL7Bp0yYAwKpVq2Bvb48uXbpo3d6zzz6LP//8EykpKQYrIyGEEEIIIYQQQmqGWhRWwNvbG5s2bcLo0aPRtWtXvPjii7C2tsapU6fAGMPGjRvRuHFjfPfdd3jjjTdQUFCAAQMG4MGDB/j2228xePBgTJs2Te/9NmnSBIsXL8a7776L3NxcdO/eHVevXkVxcbGq9VtV+vXrhz179uCnn36Cp6cnvLy88Mwzz2jNa4w6aOPt7Y1t27Zh9OjReOaZZzBixAi4uroiJiYG+/fvR+/
"text/plain": [
"<Figure size 1300x380 with 4 Axes>"
]
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"def condition_row_norm(case_id: str) -> np.ndarray:\n",
" features = load_case(case_id, max_points=1)[\"features\"]\n",
" features_norm = normalize_features_raw(features)\n",
" return np.ascontiguousarray(features_norm[:, CONDITION_INDICES], dtype=np.float32)\n",
"\n",
"\n",
"def pca2(values: np.ndarray) -> tuple[np.ndarray, np.ndarray, np.ndarray]:\n",
" centered = values - values.mean(axis=0, keepdims=True)\n",
" u, s, vt = np.linalg.svd(centered, full_matrices=False)\n",
" scores = centered @ vt[:2].T\n",
" explained = (s[:2] ** 2) / np.maximum(np.sum(s**2), 1e-12)\n",
" return scores, vt[:2], explained\n",
"\n",
"\n",
"all_case_ids = split_manifest[\"train_ids\"] + split_manifest[\"val_ids\"] + split_manifest[\"test_ids\"]\n",
"condition_tensor = torch.from_numpy(np.concatenate([condition_row_norm(case_id) for case_id in all_case_ids], axis=0)).to(device)\n",
"with torch.no_grad():\n",
" condition_embedding = model.condition_encoder(condition_tensor).detach().cpu().numpy()\n",
" gamma_norms = []\n",
" beta_norms = []\n",
" for block in model.blocks:\n",
" gamma_beta = block.film(torch.from_numpy(condition_embedding).to(device)).detach().cpu()\n",
" gamma, beta = gamma_beta.chunk(2, dim=1)\n",
" gamma_norms.append(gamma.norm(dim=1).numpy())\n",
" beta_norms.append(beta.norm(dim=1).numpy())\n",
"gamma_norms = np.stack(gamma_norms, axis=1)\n",
"beta_norms = np.stack(beta_norms, axis=1)\n",
"condition_meta = case_condition_frame(all_case_ids)\n",
"aoa = np.asarray([row[\"aoa_deg\"] for row in condition_meta], dtype=float)\n",
"u_inf = np.asarray([row[\"u_inf\"] for row in condition_meta], dtype=float)\n",
"splits = [row[\"split\"] for row in condition_meta]\n",
"\n",
"scores, _, explained = pca2(condition_embedding)\n",
"fig, axes = plt.subplots(1, 3, figsize=(13, 3.8))\n",
"scatter = axes[0].scatter(scores[:, 0], scores[:, 1], c=aoa, cmap=\"coolwarm\", s=42)\n",
"axes[0].set_title(f\"condition embedding PCA\\nexplained={explained.sum():.2%}\")\n",
"fig.colorbar(scatter, ax=axes[0], label=\"AoA deg\")\n",
"for split, marker in [(\"train\", \"o\"), (\"val\", \"s\"), (\"test\", \"^\")]:\n",
" mask = np.asarray([value == split for value in splits])\n",
" axes[1].scatter(u_inf[mask], aoa[mask], marker=marker, label=split, s=42)\n",
"axes[1].set_xlabel(\"U_inf\")\n",
"axes[1].set_ylabel(\"AoA deg\")\n",
"axes[1].set_title(\"split coverage\")\n",
"axes[1].legend()\n",
"axes[2].plot(gamma_norms.mean(axis=0), marker=\"o\", label=\"gamma norm\")\n",
"axes[2].plot(beta_norms.mean(axis=0), marker=\"o\", label=\"beta norm\")\n",
"axes[2].set_xlabel(\"FiLM block\")\n",
"axes[2].set_title(\"mean modulation norm\")\n",
"axes[2].legend()\n",
"for ax in axes:\n",
" ax.grid(alpha=0.25)\n",
"fig.tight_layout()\n"
]
},
{
"cell_type": "markdown",
"id": "541f871e",
"metadata": {},
"source": [
"## 9. Fourier/spatial-frequency probe\n",
"\n",
"Use existing sampled points in a narrow horizontal strip, sorted by `x`, to look for rough prediction artifacts along a quasi-1D cut.\n"
]
},
{
"cell_type": "code",
"execution_count": 9,
"id": "c15b98aa",
"metadata": {},
"outputs": [
{
"data": {
"text/plain": [
"[{'target': 'velocity_x',\n",
" 'truth_mean_abs_diff': 10.329564094543457,\n",
" 'pred_mean_abs_diff': 7.7607245445251465},\n",
" {'target': 'velocity_y',\n",
" 'truth_mean_abs_diff': 11.62238883972168,\n",
" 'pred_mean_abs_diff': 10.390179634094238},\n",
" {'target': 'pressure',\n",
" 'truth_mean_abs_diff': 221.22994995117188,\n",
" 'pred_mean_abs_diff': 277.8408203125},\n",
" {'target': 'turbulent_viscosity',\n",
" 'truth_mean_abs_diff': 0.0007083089440129697,\n",
" 'pred_mean_abs_diff': 0.0005382539238780737}]"
]
},
"execution_count": 9,
"metadata": {},
"output_type": "execute_result"
},
{
"data": {
"image/png": "iVBORw0KGgoAAAANSUhEUgAAA+cAAAMECAYAAADdLlSQAAAAOnRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjExLjEsIGh0dHBzOi8vbWF0cGxvdGxpYi5vcmcvctoD+AAAAAlwSFlzAAAPYQAAD2EBqD+naQABAABJREFUeJzs3Xd4FOXax/Hf7CbZJKSQUAKB0It0FBARC1iwo3hUPCoqR7FXxIKKiqAc29Hz2sCGBVHEhu0odkABQSkKIogIoZcEEgipM+8fyS7ZZJPspk2S+X6uiyvZZ5955p65d5fc+0wxLMuyBAAAAAAAbOOyOwAAAAAAAJyO4hwAAAAAAJtRnAMAAAAAYDOKcwAAAAAAbEZxDgAAAACAzSjOAQAAAACwGcU5AAAAAAA2ozgHAAAAAMBmFOdADRo3bpxmzpxZo8sG6leV9SI0y5Yt0zXXXKNNmzbZHUq51qxZo2uuuUZr1661Zf333XefXn755QrbAAAAnMqwLMuyOwigqrZt26YPP/xQf//9tyIiItS1a1edeeaZaty4cY2v+5ZbbtGgQYM0cuTIUs81bdpU5513nqZOnRryuMEuG6hfVdZbF5W3j+327rvv6vzzz9eSJUvUv39/u8Mp01dffaWTTz5Z3377rYYMGVJh/02bNunhhx8O+FxMTIwef/zxkNbfqVMn9e/fX2+//Xa5bZK0YMECLVy4UNu3b1ebNm102mmnqUuXLqXGnDJlijZu3ChJcrlcatSokZo3b64+ffpoyJAhioiICCnG4nbv3q33339fGzZskGEY6tKli84880w1bdo05L4fffSRPvvsswrXGR4erqeffjroGNPS0vTOO+/or7/+UlhYmLp06aILLrhA0dHRpfquW7dOX3zxhf766y8lJiZqwIABGjZsmAzDqHA906dP1+LFi0u1R0ZG6qmnnirVHmz+vHbu3Kn3339fa9euVVJSkkaMGFFu//IUFBTovffe008//SS3260hQ4botNNOC3r5YGI/cOCAbrvttjLHSEhI0JQpU3yPQ91/wahKPiUpNzdXc+fO1ZdffqmcnBz997//lcfjqXLf4qZOnarly5dr0KBBuuyyy0LavqqOF2qOpNDeTzt27NCsWbO0fv16xcfH69RTT9XRRx9duQ1T5fexV1Vf91WJJ9T3e0k18TlWE+85oLYwc45678knn1SHDh00a9YsxcXFKSYmRv/3f/+nVq1a6cUXX6zx9b/00kuaP39+wOeeeOIJXXzxxTUeQ11Zb00pbx8jON26ddPzzz8f9B9NO3fu1LRp07Rjxw717dvX71+vXr1CXv+kSZN05ZVXVrjOww47TKeeeqrWr1+vVq1a6YcfflDPnj117733lur/3nvvafbs2erbt6969+6tli1bKjU1Vddee61atmwZUqFb3EsvvaR27drplVdeUXR0tBISEvTyyy+rdevWeuKJJ0Lu27JlS7/9l5KSomnTpmndunV+7X369Ak6xo8//lht2rTRK6+8oiZNmig6OloPP/yw2rVrp59//tmv71VXXaWRI0dq/fr1SklJ0fbt23X++edrwIAB2rNnT4Xr+vrrr/XWW29V+DoINX+SNHv2bHXp0kVffPGFUlJSlJubq/PPP1/vvPNO0PvC68CBAzr++ON12223KTExUR6PRxdffLEuuOACFRQUlLtsKLG73e5S+6Jv376KiYnRtGnT9Msvv1Rq/wWrqvl8/vnnlZKSoqlTp2rBggWaNm2a8vLyqty3uO+//1433HCDpk2bpm+//TbkbazqeKHmKJT309dff60uXbrozTffVHJystLS0jRkyBDdfPPNldq2yu5jr6q87qsST2Xe7yXV1OdYdb/ngFplAfXY6tWrLcMwrFGjRvm1m6Zp/ec//7HGjh1b4zE0atTIuv7666t93CZNmlhXX311tfWrz2pqH1eH2bNnW5KsJUuW2B1KtVqyZIklyZoyZUqNraNjx47WyJEjfY/XrVtnHXnkkdbff//t1++BBx6wJFlz5871a+/Xr5/Vtm3bUuMWFBRYd999tyXJevDBB0OKadOmTVZ4eLg1fPjwUs+98MIL1pVXXlmpvsVt2LDBkmTdfPPNIcVWXOvWra2OHTtaOTk5vrbdu3dbcXFx1pAhQ/z6rlq1qtTy33zzjSXJuv322ytc18UXX2wlJSVV2C/U/C1atMhyu93WCy+84Neel5dnbdiwocL1lTR27FjL4/FYf/75p69twYIFliTrueeeq9bYA7ntttssSdbs2bP92oPdf8Gqaj5XrFhhpaenW5ZlWZdddpklycrMzKxyX699+/ZZ7dq1s6677jpLknXZZZdVGFNtjVdWjoJ9Px04cMBq1qyZNWDAACsvL8/X/uabb1qSrHfffTfkmCqzj4uryuu+KvFUx3umpj7Hqvs9B9SmsNr8IgCobosXL5ZlWRo+fLhfu2EYuvXWW7Vz506/9uuvv14nnHCCTjvtNL3xxhv6/fff1aZNG1166aWlDledNGmStmzZIkkKCwtTUlKSTjnlFB155JGSDh02l5OTo++++07XXHONpMIZSu836OPGjdMRRxyhiy66KOhxq0Og9Xq3/cwzz9Sbb76plStXqlWrVrr88svVrFmzUmPs27dPs2fP1urVq+V2uzV48GANHz5cLldwB9xkZGTovffe06pVqxQZGanjjjtOw4YN8z13xx136IILLtAJJ5zgt9ykSZPUpEkTXXfddUHt40BCybO37xlnnKG33npLK1as0JlnnqmTTjpJkjRv3jx9+eWXysjIULt27XT++eerdevWAdf7119/acaMGdq7d6+OOuoonXfeeQH31++//645c+Zo69atatq0qUaMGFHqG/177rlHkvTQQw+Vt5u1detWPfjgg5IKX/eRkZHq1KmT/vGPf6hFixa+fmvWrNFTTz2lsWPH+mbPly1bpmnTpunuu++WZVmaOXOmUlNTyzycvSzbtm3zHZYYGxuroUOH6sQTT/Trc99996lt27a64ooryhynRYsW+uabb9SoUSO/9gsuuEAPPPCA79D8irhcLj300EP64YcfNHnyZP3rX/9Sq1atgtqWn3/+WXl5eTrrrLNKPTdmzBidffbZlepbnSzL0o4dO3TGGWf4HbrfpEkTtW/fXtu2bfPr371791JjDB48WIZhVOu1EkLN3/jx49WrVy+NGTPGr39YWJjatWsX0rrz8vL0yiuv6IwzzlDHjh197YMHD9YRRxyh5557Ttdee221xR5o/W+88YaaN29eY3n3qmo+e/fuHfS6QunrdfPNNyssLEwPPPCAnnvuuZCXr6nxyspRKO+n77//Xrt27dKkSZMUFnboT+h//vOfuu666/T000/rH//4R0hxVWYfF9+mqrzuqxJPVd8zdfVzDLAbh7WjXktJSZEkLV++PODzzZs393vsPSTurLPO0rp169S0aVO99NJL6tGjh1avXu3Xt2vXrr5DoTp27Ki1a9fqmGOO8RUu3sPm3G63mjVr5tfX69VXX9W8efNCGrc6BFrvtGnTNG/ePJ133nlatWqVmjdvrqlTp+rwww/X3r17/fr+8MMP6tSpk5599lklJCQoMjJS1113nU466SRlZ2dXuP558+apY8eOeuyxx+TxeBQdHa0nn3xS//rXvyRJWVlZmjZtmlauXFlq2VmzZvnOzw1mHwcSSp6nTZum7777TmeddZZ+/fVXRUVF6ffff5dlWRo1apROPPFE7dixQ8nJyfrwww/VuXNnvf/++6XWuXjxYl166aVyu92yLEujR4/WGWecofz8fL9+999/v3r27Klly5apTZs22rRpk4444ohS58G9+eabevPNNyvc11FRUX6HRTdt2lRvvPGGunTpohUrVvj6bd68WdOmTdPWrVt9bevXr9e0adP0ySef6OKLL1Z+fr5SU1O1f//+Ctfr9b///U+dO3fW22+/rRYtWig9PV2nn366/vGPf/ht+8yZM/Xll1+WO1ZMTEypP/Qk+f7wCnQOYnkuvfRS5ebm6tNPPw16mVA+U0L9/KkuhmHo9NNP148//qhdu3b52tesWaM1a9b
"text/plain": [
"<Figure size 1000x760 with 4 Axes>"
]
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"def horizontal_strip(features: np.ndarray, targets: np.ndarray, *, max_points: int = 256) -> tuple[np.ndarray, np.ndarray]:\n",
" y = features[:, FEATURE_INDEX[\"y\"]]\n",
" center = np.median(y)\n",
" width = np.quantile(np.abs(y - center), 0.20)\n",
" mask = np.abs(y - center) <= max(width, 1e-6)\n",
" selected_features = features[mask]\n",
" selected_targets = targets[mask]\n",
" if len(selected_features) > max_points:\n",
" order = np.argsort(selected_features[:, FEATURE_INDEX[\"x\"]])\n",
" keep = order[np.linspace(0, len(order) - 1, max_points).astype(int)]\n",
" selected_features = selected_features[keep]\n",
" selected_targets = selected_targets[keep]\n",
" order = np.argsort(selected_features[:, FEATURE_INDEX[\"x\"]])\n",
" return selected_features[order], selected_targets[order]\n",
"\n",
"\n",
"full_case = load_case(example_case, max_points=4096)\n",
"strip_features, strip_targets = horizontal_strip(full_case[\"features\"], full_case[\"targets\"])\n",
"strip_predictions = predict_targets(strip_features)\n",
"strip_x = strip_features[:, FEATURE_INDEX[\"x\"]]\n",
"\n",
"fig, axes = plt.subplots(len(TARGET_NAMES), 1, figsize=(10, 1.9 * len(TARGET_NAMES)), sharex=True)\n",
"roughness_rows = []\n",
"for ax, name in zip(axes, TARGET_NAMES):\n",
" idx = TARGET_INDEX[name]\n",
" ax.plot(strip_x, strip_targets[:, idx], label=\"truth\", linewidth=1.2)\n",
" ax.plot(strip_x, strip_predictions[:, idx], label=\"prediction\", linewidth=1.2)\n",
" ax.set_ylabel(name)\n",
" ax.grid(alpha=0.25)\n",
" truth_rough = float(np.mean(np.abs(np.diff(strip_targets[:, idx]))))\n",
" pred_rough = float(np.mean(np.abs(np.diff(strip_predictions[:, idx]))))\n",
" roughness_rows.append({\"target\": name, \"truth_mean_abs_diff\": truth_rough, \"pred_mean_abs_diff\": pred_rough})\n",
"axes[0].legend(loc=\"best\")\n",
"axes[-1].set_xlabel(\"x along sampled horizontal strip\")\n",
"fig.suptitle(f\"Spatial line-cut probe: {example_case}\", y=1.01)\n",
"fig.tight_layout()\n",
"roughness_rows\n"
]
},
{
"cell_type": "markdown",
"id": "f972aef3",
"metadata": {},
"source": [
"## 10. Activation representation probes\n",
"\n",
"Manually run the trunk and project hidden states to two PCA dimensions. Points are colored by raw `sdf` to see whether layers separate near-wall/farfield geometry.\n"
]
},
{
"cell_type": "code",
"execution_count": 10,
"id": "4bd0d392",
"metadata": {},
"outputs": [
{
"name": "stderr",
"output_type": "stream",
"text": [
"/tmp/ipykernel_9826/1784321331.py:40: UserWarning: This figure includes Axes that are not compatible with tight_layout, so results might be incorrect.\n",
" fig.tight_layout()\n"
]
},
{
"data": {
"image/png": "iVBORw0KGgoAAAANSUhEUgAABjYAAAFLCAYAAAB1O7U1AAAAOnRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjExLjEsIGh0dHBzOi8vbWF0cGxvdGxpYi5vcmcvctoD+AAAAAlwSFlzAAAPYQAAD2EBqD+naQABAABJREFUeJzs3XdYFMcbB/Dv3dEEQZoiKIgF7L1X7D2WWGPvXWOLxtiNMWoSTSxBYzf23o1dxBbFjmJHBEUBkV7vdn5/8GPjeUcRkaLfz/PwJDc7O/vuKuPevjszCiGEABERERERERERERERUS6gzO4AiIiIiIiIiIiIiIiI0ouJDSIiIiIiIiIiIiIiyjWY2CAiIiIiIiIiIiIiolyDiQ0iIiIiIiIiIiIiIso1mNggIiIiIiIiIiIiIqJcg4kNIiIiIiIiIiIiIiLKNZjYICIiIiIiIiIiIiKiXIOJDSIiIiIiIiIiIiIiyjWY2CAiIiKidHv16hVWrFgBf3//NOtGRkZixYoVuHfvXqbWJaIPc/nyZaxevfqLOzYRERERfb6Y2CAiIiL6Qr19+xYrVqyAl5eX3u337t3DihUr8OLFC7ns8ePHGD58OO7evZtm+2/evMHw4cNx7ty5TK2bXV68eKFzPbKjjS8Zr1/G7Nq1C6NGjfpk7Z8/fx5r1qzJlmMTERER0ZeJiQ0iIiKiL9SLFy8wfPhwHDp0SO/2c+fOYfjw4fDx8ZHL7O3tMXToUDg5OWVVmDmGj4+PzvXIjja+ZLx+GVO7dm0MHjz4k7W/bds2fPvtt9lybCIiIiL6MhlkdwBERERElHsUL14cK1asyO4wiOgDdOrUCZ06dfrijk1EREREny8mNoiIiIgo3V69eoV9+/ahTZs2cHR01Nrm6+uL8+fPw8zMDM2bN0+1nQ+pGxsbC09PTwQEBMDGxgYNGzZEvnz55O3Pnz/HkSNH0KFDB9ja2uLEiRN4/fo1qlativLly6f73BISEuDp6YkXL16gQIECqFatGmxtbQEA9+/fl0e2HDp0CI8fPwYA1K9fH2XLlkVERAS2bNkit2ViYoLixYujTp06UKlU6WojveebmidPnuDixYvydRVCYPPmzWjQoAHKlCnzQdc1WUhICDw9PfH27Vs4OTmhQYMGMDIykre/e/2tra1x/PhxREREoFGjRrC3t5ev7YkTJxASEoL69eujWLFieuNPLaa0rt+7cdjY2OD06dN4/vw56tevj7Nnz6JJkyZwcXHROp4QAqtXr4aLiwsaNmyY6rX9kOvw7vG7desGCwsLCCFw4cIFPH78GEWLFkWDBg3g7e2NCxcuoG/fvsiTJw8A4NatW7h06RIAQKFQwMLCApUrV0apUqW04vHw8MCzZ8/Qt29fBAcH4+TJk1AoFGjcuDEKFCigVffy5cvw9vbGoEGDAABBQUHYs2dPiuc6ePBgqFSqdMVy9OhR3L17F2q1Wivp2bt3b5iZmekc+10+Pj64du0aJElCxYoVUbFixQyfIxERERF9WTgVFRERERGlW0prbPz0009wcXHBihUrsH37dtSrVy/FdTg+pO7+/ftRpEgRjBkzBqdPn8a8efNQpEgR7N+/X65z+/ZtDB8+HBcvXkSTJk2wZs0abNq0CRUqVMCcOXPSdV63bt2Cs7MzRo8ejVOnTmH58uWoUaOG/KD2zZs3ePLkCYCk5MHNmzdx8+ZNBAcHAwDi4+Plsps3b+Lo0aPo1KkTSpUqhefPn6erjfSeb0rmzJkDV1dXrFy5Ejt27ECDBg3k6cTeX7skvcdZunQpHB0d8dNPP+HEiRMYMGAAXFxccOXKFZ3rf+7cOTRq1AgbN27E77//DhcXF5w/fx7Pnz9HgwYNsGHDBri7u6NUqVLYt2+fTvxpxZTW9UuOw8PDA25ubli5ciUWL14MtVqNH374ATNmzNA55rFjxzBkyBAEBgamem0/5Dq8f/ygoCBERkaiSZMmaNmyJfbt24cFCxagXbt2OHLkCIYPH47w8HC5neDgYPncrl27ho0bN6JChQro3r07hBByvb///hsTJkzAyZMn0aJFCxw6dAizZs1C8eLFcfHiRa3431/nIjo6Wuvva/LP5MmTMXz4cKjV6nTH4uvri+DgYGg0Gq22EhMT9R4bSEpgde3aFZUqVcKWLVuwe/du1KlTB61atdK6Fh9yjkRERET0hRFERERE9EW6c+eOACDatGkj3N3ddX569OghAIgTJ07I+3h6egoA4ujRo3LZ0aNHBQAxb948uez169eiUaNGAoBwd3fPUN1r164JQ0NDMWzYMKFWq+Xy77//XuTJk0c8ffpUCCHEwYMHBQBRq1Yt4efnJ9cbP368MDQ0FM+fP0/zWrRo0UJUr15daDQauSw6OlocO3ZM/nzixAmd65GaqKgoUbFiRfHVV1+lq430nq8+hw8fFgDEwoUL5bKQkBDRtGnTDF/X48ePCwBi0qRJcp2IiAhRq1YtUaBAAREWFiaE0L7+AQEBQgghNBqNaNy4sShfvrzo1q2b/GcgSZJo3ry5KF68uNax0xtTatcvOY6qVavK9aOjo0VoaKgYP368MDIyEq9evdLa56uvvhJWVlYiNjY2xWv7oddB3/EHDx4sTExMxM2bN+U2zp07J5ycnAQAERgYmOLxhRDixo0bwsDAQKxZs0YuGzhwoDA1NRWDBg0SCQkJQgghYmJiRNmyZUWtWrW09p8wYYIwNjZO9Rjz5s0TAMTo0aM/OJaRI0cKMzMzvfX1HXvkyJFCpVKJc+fOyWU3b94UpqamomvXrhk6RyIiIiL6snDEBhEREdEX7tWrV3rf3vb390/X/itWrICdnR0mTpwolxUoUAA9e/b8qLq//fYbDA0NsWjRInk6JwCYOXMmhBBYv369Vv0OHTpoLWo+YMAAJCYmwtPTE0DStEMrVqzQ+rl16xaApNEASqUSkiTJ+5uamqY5Tdb7bt++jS1btmDlypX4+++/YWdnJx8/LR96vu9auXIlChYsiHHjxsllNjY2H3Vdly9fDktLS8yaNUuuY25ujp9++glBQUHYtm2bVrsdO3ZEoUKFAABKpRLdu3fHnTt3UKFCBXnaMoVCge7du+PJkyd49uxZppz7+9q2bYuiRYsCSPoztLKywvDhw5GYmIg1a9bI9ZKnjurZsydMTExSbO9Dr8P7xzcxMcHGjRvRs2dPramW6tevj8qVK+s9ZkJCAs6ePYv169dj5cqVuHz5MmxsbHRG3sTExGDixIkwNDQEAOTJkwfffPMNLl++jJiYmHRcrSTbt2/H1KlT0a5dO/z+++8ZiiW94uLisGbNGnTq1An169eXyytWrIiBAwdi586dePXqVaafIxERERF9XrjGBhEREdEXrm3btloPbZOtWLEiXQ/l79y5g7Jly8oPHpPpe2j7IXWvXr2K/PnzY/PmzQCSEhPJP+bm5rh3755W/QoVKmh9dnBwAAAEBAQAADQaDYYPH65V55dffkHFihUxfPhwDB48GC4uLmjfvj3c3NzQuHHjdK9tERwcjA4dOuD27duoX78+HBwcYGBggDdv3iAsLAzx8fEwNjZOtY0PPd933blzB2XKlIGBgfbt/fvX5EOOc+fOHZQuXVpe+yFZ1apV5e3ven89k+T1NVIqf/HiBYoXL/7R5/6+SpUq6ZSVKFECzZs3x8qVK/H9999DqVRi5cqV0Gg0GDBgQKrtfeh1eP/4jx49Qnx8vM76EUDSn8/703+dOXMG33zzDYyMjFCrVi1YWlpCqVRCrVbj9evXWnWNjY3h6uqqVfbu3/v3t+mTvMZHlSpVsHXrViiV/7379iGxpNejR48QFxcnX793Va1aFUIIeHt7o2DBgpl2jkRERET0+WFig4iIiIg+ilqt1klUANBaWDkjdRMSEqDRaODl5aWz7euvv9Z5aG9ubq71Ofk4CQkJAJJGEQwdOlSrTvJD6AEDBqB69erYsWMHzp07hxUrVkCpVGLBggUYPXq0zvHfN336dNy5cwd37tyBs7OzXD5q1Chcu3ZNa22ElHzo+b5Lo9HovYYfc13TalOj0WiVv3/9k5MsKZUn/7l8SEz
"text/plain": [
"<Figure size 1600x320 with 6 Axes>"
]
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"def fourier_features(coordinates: torch.Tensor, scales: torch.Tensor) -> torch.Tensor:\n",
" if scales.numel() == 0:\n",
" return coordinates\n",
" phases = coordinates.unsqueeze(-1) * scales.to(device=coordinates.device, dtype=coordinates.dtype) * math.pi\n",
" return torch.cat((coordinates, torch.sin(phases).flatten(1), torch.cos(phases).flatten(1)), dim=1)\n",
"\n",
"\n",
"def hidden_snapshots(features: np.ndarray, *, layer_indices: tuple[int, ...] = (0, 5, 11)) -> dict[str, np.ndarray]:\n",
" features_norm = torch.from_numpy(normalize_features_raw(features)).to(device)\n",
" coordinates = features_norm.index_select(dim=1, index=model.coordinate_indices.to(device))\n",
" conditions = features_norm.index_select(dim=1, index=model.condition_indices.to(device))\n",
" snapshots: dict[str, np.ndarray] = {}\n",
" with torch.no_grad():\n",
" hidden = model.input(fourier_features(coordinates, model.fourier_scales))\n",
" snapshots[\"input\"] = hidden.detach().cpu().numpy()\n",
" unique_conditions, inverse = torch.unique(conditions, dim=0, return_inverse=True)\n",
" condition_embedding_batch = model.condition_encoder(unique_conditions)\n",
" for index, block in enumerate(model.blocks):\n",
" hidden = block(hidden, condition_embedding_batch, inverse)\n",
" if index in layer_indices:\n",
" snapshots[f\"block_{index}\"] = hidden.detach().cpu().numpy()\n",
" hidden = model.output_norm(hidden)\n",
" hidden = model.activation(hidden)\n",
" snapshots[\"final\"] = hidden.detach().cpu().numpy()\n",
" return snapshots\n",
"\n",
"\n",
"activation_case = load_case(example_case, max_points=FAST_SAMPLE_POINTS)\n",
"snapshots = hidden_snapshots(activation_case[\"features\"])\n",
"sdf_values = activation_case[\"features\"][:, FEATURE_INDEX[\"sdf\"]]\n",
"fig, axes = plt.subplots(1, len(snapshots), figsize=(3.2 * len(snapshots), 3.2))\n",
"for ax, (name, values) in zip(axes, snapshots.items()):\n",
" scores, _, explained = pca2(values.astype(np.float64))\n",
" scatter = ax.scatter(scores[:, 0], scores[:, 1], c=sdf_values, cmap=\"viridis\", s=16)\n",
" ax.set_title(f\"{name}\\nPCA2={explained.sum():.1%}\", fontsize=9)\n",
" ax.set_xticks([])\n",
" ax.set_yticks([])\n",
"fig.colorbar(scatter, ax=list(axes), label=\"raw sdf\", shrink=0.75)\n",
"fig.suptitle(\"Hidden-state geometry organization\", y=1.02)\n",
"fig.tight_layout()\n"
]
},
{
"cell_type": "markdown",
"id": "d3a7524a",
"metadata": {},
"source": [
"## 11. Gradient attribution\n",
"\n",
"Average absolute input gradients answer: “which normalized input dimensions locally move each normalized output?” This is local sensitivity, not causality.\n"
]
},
{
"cell_type": "code",
"execution_count": 11,
"id": "e66b641f",
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"velocity_x [('y', 15.791990280151367), ('sdf', 13.878210067749023), ('x', 2.7461659908294678), ('u_inf', 0.009471949189901352)]\n",
"velocity_y [('y', 4.476894378662109), ('sdf', 3.4993419647216797), ('x', 1.649538516998291), ('aoa_deg', 0.004039796069264412)]\n",
"pressure [('y', 2.3920469284057617), ('sdf', 2.1177940368652344), ('x', 0.7311617136001587), ('aoa_deg', 0.0024339091032743454)]\n",
"turbulent_viscosity [('y', 8.357393264770508), ('sdf', 5.822950839996338), ('x', 0.8091668486595154), ('naca_param_3', 0.0036486415192484856)]\n"
]
},
{
"data": {
"image/png": "iVBORw0KGgoAAAANSUhEUgAAA+4AAAFeCAYAAAAbsUXiAAAAOnRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjExLjEsIGh0dHBzOi8vbWF0cGxvdGxpYi5vcmcvctoD+AAAAAlwSFlzAAAPYQAAD2EBqD+naQAAxTNJREFUeJzs3XVcFPn/B/DX0g0SipQiiqCYiIV1oNiKXdhdp2d3Ynd7enaf3Xq2nq1nnxgoiiCgKCHN7vv3hz/m67q7lOgM5/t5j3k8bj8z+5n3LAPuez4lIyICY4wxxhhjjDHGJElL7AAYY4wxxhhjjDGmGSfujDHGGGOMMcaYhHHizhhjjDHGGGOMSRgn7owxxhhjjDHGmIRx4s4YY4wxxhhjjEkYJ+6MMcYYY4wxxpiEceLOGGOMMcYYY4xJGCfujDHGGGOMMcaYhHHizhhjjDHGGGOMSRgn7owxls+lp6dDJpPhjz/+yNbx8+fPh0wmw/v3779zZNJVvHhxtG/fPsuyH6FOnTqoWrXqDz/vz+Rn/4yDgoIgk8mwdevWTMt+hE+fPkEmkyEwMDDT48SK71v16tULlSpVEjsMxth/ECfujDHG2E/KwcEBAQEB+eKcBw4cgEwmw5s3b75DVIxlz927dyGTybBz506xQ2GM/WR0xA6AMcYYk4Lnz5+LHQLLxKFDh+Dp6QkHBwexQ/lPcHNzAxGJHYZGUo+PMcZ+NG5xZ4wxxpikKRQKHDlyBM2bNxc7FMYYY0wUnLgzxn4q3bp1g62tLaKjo9G2bVuYmZnB3t4ey5cvBwC8e/cO7dq1g4WFBWxsbDBhwgS1rT737t1Dy5YtYW1tDX19fZQqVQqrVq1SOmb27NmQyWTCZm5ujl9++QWnT59WG1NsbCw6d+4Mc3NzFChQAN27d0dCQsI3Xe/OnTvh4eEBAwMDlCpVCvv27cv2e3Ma18mTJ1GrVi2YmJjA2NgY1atXx+HDh9XW+eHDBwQEBKBAgQKoUKECAMDa2hq9evXC1atXUaVKFRgZGaFChQq4evUqAODy5cuoWrUqDA0N4erqiiNHjqjEUKNGDeHz1tbWRuHChdGpUye8fv06y+v9eoz7mDFjlH5+X24NGjRQeu/OnTtRrVo1GBsbw9jYGD4+Prh8+bLSMampqRgzZgzs7OxgbGyMunXr4unTp1nGlUGhUGDevHlwd3eHvr4+rK2t0aZNG6U6kpOTIZPJMGXKFLWfTY0aNQD8b16EsLAwbNu2TbiuL8fmZvw8rly5gsqVK8PQ0BDFihXDokWLlOrNy3NqcuXKFbx79y7LxP1bP+OMa75z5w5q1KgBQ0NDFClSBEuXLlU5Njs/jy/rvH79ulDnhAkTcO3aNchkMhw4cACrVq2Cs7MzTExM4O/vj+joaADAihUrUKxYMRgYGMDHxwcvX75UqjsmJkbpvtTX14erqysmT56M1NTUTK9V3RhyDw8Pjfd8xt9IAEhKSsLEiRPh6uoKfX19FCxYEN26dUNERITSOV69eoWWLVvC1NQUVlZW6N+/P5KSkrL1s1AX35ef2YYNG1C8eHEYGBigUqVKKr9vXx67YsUKFC1aFAYGBqhatSouXLigdOyePXsgk8lw69YttZ/v7NmzAXwerpHx96pDhw7CZ6Pu3meMsbzGiTtj7KdDRBg8eDAGDBiAN2/eYOLEiRg8eDAOHDiAnj17om/fvnj9+jVmzpyJGTNmYPv27Urv//vvv1G1alVoa2vj8uXLeP/+PSZOnIhRo0YpfYEbM2YMiAhEhPT0dNy7dw/u7u5o0qQJHj58qBLToEGD0LlzZ7x58wZbtmzBn3/+ifHjx+f6Onfs2IEOHTqgYcOGCAkJwfHjx3HgwAFcunQpR59VduLas2cPGjZsCHd3dzx+/BjPnj1DlSpV0KxZM2zatEmlzj59+iAgIAAvXrzAkCFDhH2vX7/GokWLsG3bNrx69Qqurq5o1KgRrl69innz5mHz5s0IDQ2Fl5cX2rRpg8jISKW6//77b+EzT0hIwKFDhxAcHIwmTZpkmch8bfbs2UJdGduSJUsAfE5wMkydOhWdO3dGixYt8Pz5c7x48QIVK1aEj48Prly5IhzXvXt3rFixAgsWLEB4eDhmzZqFwYMH49OnT9mKp1evXpg4cSKGDx+OyMhInDt3Dm/evEHVqlURHByco2vT0dEBEcHe3h6dOnUSru/rxOXVq1eYNWsWNm3ahLCwMPz6668YOXIkpk+fnqPz5eSc6hw8eBBFixZF2bJlMz3uWz9jAHjz5g1mzpyJtWvX4u3bt+jSpQuGDBmi8sAtJz+PV69eITAwEGvWrMHz58+VHlZs2bIF79+/x82bN3H9+nU8ePAAXbt2xcqVK/Hu3Ttcv34dd+7cQWhoKLp06aJUr4WFhdL9GRkZiVmzZmHZsmUYN25ctq85w8OHD5XqS09PR4MGDaClpYVSpUoB+PxwpF69etiwYQMWLFiAd+/e4fz58wgODkatWrUQHx8P4HPSW7t2bTx9+hRnz57FixcvULNmTfz22285jutrO3fuxIsXL/D333/j6dOnMDExQfPmzdX+nDPu3evXryMoKAhOTk7w8/PD9evXc3xef39/3LlzB8Dnv68ZnxMn7oyxH4IYY+wn0rVrVwJAR44cUSovW7YsGRsb04EDB5TKK1SoQL/88otSWZkyZah06dKUlpamVB4YGEj6+vr0/v37TGOwtbWlESNGqMR06NAhpeN69+5NRkZGpFAoMq0vLS2NANDatWuFMoVCQU5OTuTt7a1yrLOzMwGgd+/eZVpvduNSKBRUpEgRKleunEqs3t7eZGNjQ6mpqUp17tq1S+V8VlZWZG5uTjExMULZ69evCQAVLlyYPnz4IJSHhYURAFqwYEGm10BEdOvWLQJA58+fF8pcXFyoXbt2SsepK/vSgQMHSEtLi/z8/ISffXBwMGlra9Pw4cNVjq9WrRrVrl2biIju3btHAGjWrFlKx9y5c4cAUJUqVTK9hoz3jxs3Tqn87du3ZGBgQJ06dSIioqSkJAJAkydPVqnD29tb5X6wt7cX3vs1KysrMjY2Vrmfu3XrRgYGBsLPIy/PqYmrqysNGTIk02O+9TMm+nzNZmZm9PHjR6FMLpeTg4MDtWnTRuVcWf08Muo0NDRU+X27evUqAaAmTZoola9cuZIAUKtWrZTKf//9dwJAT548yfI6pkyZQqampsLrx48fEwDasmVLpmVf69u3r8rv2YoVKwgAnTt3TunYsLAw0tfXp3nz5hER0YwZMwgA3b9/X+m4wMBAAkDTp0/P9BrUxZfxmTVq1Ejp2Nu3bxMA2rRpk8qxvr6+SsempKSQvb09+fj4CGW7d+8mAHTz5k2lYz9+/KhyT2XcTzt27FAbd8+ePcnT0zPTa2OMsdzgFnfG2E9HW1sbfn5+SmVubm5ITExU6QLt7u6OFy9eCK9fvXqFBw8eoEWLFtDRUZ7fs27dukhJScG1a9cAAPHx8Rg9ejRKliwJAwMDoVtlRESEykRoWlpaKuf28PBAYmKiSqtydjx//hyvX79Gs2bNlMp1dHTQuHHjbNeTnbiCg4Px6tUrtGjRAjKZTOnY1q1b4927d7h//75SedOmTdWer3r16jA3NxdeOzo6wsjICB4eHihQoIBQbmdnBzMzM6WfDQA8fvwY7du3h52dHXR0dJS6Yn/L5HO3bt1Cx44dUbp0aezevVv42R8/fhxyuRxt2rRReY+vry8uX76M9PR0nDlzBgBUfh7ly5dH0aJFszz/2bNnAQAtW7ZUKre1tYW3t7dQf16rXr06rKyslMr8/f2RnJys1Jvge3r8+DGePn2aZTf5b/2MM3h7e8PCwkJ4raWlpfJ3IKc/j+rVq8Pa2lrt+Ro2bKj02s3NDQBQq1YtpXJ3d3cAULnn9+/fjzp16sDCwgJaWlpC1+34+Phc/e3IMHfuXPz+++8YMGAAhg0bJpQfPnwY1tbWqFOnjtLxdnZ2cHd3F7qhnzlzBs7OzihTpozScf7+/rmOKcPXf8MyesB8/dkAqveDnp4eGjZsiEuXLiEtLe2bY2GMsR+FE3fG2E/Hyso
"text/plain": [
"<Figure size 1100x360 with 2 Axes>"
]
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"def forward_dense_conditions(features_norm: torch.Tensor) -> torch.Tensor:\n",
" \"\"\"Differentiable forward pass without torch.unique, used for input gradients.\"\"\"\n",
" coordinates = features_norm.index_select(dim=1, index=model.coordinate_indices.to(features_norm.device))\n",
" conditions = features_norm.index_select(dim=1, index=model.condition_indices.to(features_norm.device))\n",
" hidden = model.input(fourier_features(coordinates, model.fourier_scales))\n",
" condition_embedding_batch = model.condition_encoder(conditions)\n",
" for block in model.blocks:\n",
" gamma_beta = block.film(condition_embedding_batch)\n",
" gamma, beta = gamma_beta.chunk(2, dim=1)\n",
" update = block.linear(block.activation(block.norm(hidden)))\n",
" update = update * (1.0 + gamma) + beta\n",
" hidden = hidden + update\n",
" hidden = model.output_norm(hidden)\n",
" hidden = model.activation(hidden)\n",
" return model.output(hidden)\n",
"\n",
"\n",
"grad_case = load_case(example_case, max_points=GRAD_SAMPLE_POINTS)\n",
"features_norm_tensor = torch.from_numpy(normalize_features_raw(grad_case[\"features\"])).to(device)\n",
"features_norm_tensor.requires_grad_(True)\n",
"\n",
"grad_rows = []\n",
"with torch.enable_grad():\n",
" outputs = forward_dense_conditions(features_norm_tensor)\n",
" for target_name, target_index in TARGET_INDEX.items():\n",
" grad = torch.autograd.grad(outputs[:, target_index].mean(), features_norm_tensor, retain_graph=True)[0]\n",
" grad_rows.append(grad.detach().abs().mean(dim=0).cpu().numpy())\n",
"grad_matrix = np.stack(grad_rows, axis=0)\n",
"\n",
"fig, ax = plt.subplots(figsize=(11, 3.6))\n",
"image = ax.imshow(grad_matrix, aspect=\"auto\", cmap=\"magma\")\n",
"ax.set_yticks(np.arange(len(TARGET_NAMES)), TARGET_NAMES)\n",
"ax.set_xticks(np.arange(len(FEATURE_NAMES)), FEATURE_NAMES, rotation=60, ha=\"right\")\n",
"ax.set_title(\"mean |d normalized output / d normalized input|\")\n",
"fig.colorbar(image, ax=ax, label=\"absolute gradient\")\n",
"fig.tight_layout()\n",
"\n",
"for target_name, row in zip(TARGET_NAMES, grad_matrix):\n",
" top = np.argsort(row)[-4:][::-1]\n",
" print(target_name, [(FEATURE_NAMES[index], float(row[index])) for index in top])\n"
]
},
{
"cell_type": "markdown",
"id": "52fc0213",
"metadata": {},
"source": [
"## 12. Counterfactual condition swaps\n",
"\n",
"Keep coordinates from one case, swap all condition features from another case, and measure output deltas. This directly probes how strongly the FiLM path conditions a fixed point cloud.\n"
]
},
{
"cell_type": "code",
"execution_count": 12,
"id": "96b0a939",
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"source condition: {'case_id': 'airFoil2D_SST_83.526_0.275_3.114_4.789_1.0_12.835', 'split': 'val', 'u_inf': 83.5260009765625, 'log_re': 15.49339771270752, 'aoa_deg': 0.2750000059604645, 'aoa_sin': 0.004799637012183666, 'aoa_cos': 0.9999884963035583, 'naca_param_3': 12.835000038146973}\n",
"donor condition: {'case_id': 'airFoil2D_SST_52.303_4.27_0.815_7.677_1.0_13.254', 'split': 'test', 'u_inf': 52.303001403808594, 'log_re': 15.025293350219727, 'aoa_deg': 4.269999980926514, 'aoa_sin': 0.0744565948843956, 'aoa_cos': 0.9972242712974548, 'naca_param_3': 13.253999710083008}\n",
"mean absolute target delta: {'velocity_x': 25.142242431640625, 'velocity_y': 11.112140655517578, 'pressure': 667.82080078125, 'turbulent_viscosity': 0.0004072704759892076}\n"
]
},
{
"data": {
"image/png": "iVBORw0KGgoAAAANSUhEUgAABOYAAAEwCAYAAAATo36/AAAAOnRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjExLjEsIGh0dHBzOi8vbWF0cGxvdGxpYi5vcmcvctoD+AAAAAlwSFlzAAAPYQAAD2EBqD+naQAA0N5JREFUeJzs3Xd8FGX+B/DPzO6m90IKqSSEEukdaRJDkS6geCDCD0FQ1BNO1LOdeqeoeHqHDRQQbKhwoqio9N4JHQKEJJTQEtLLZnfn+/tj2Uk2u5vsJrvJJvm+X6+I+8wzM8/Mznz3mWeeeUYgIgJjjDHGGGOMMcYYY6xeiQ1dAMYYY4wxxhhjjDHGmiNumGOMMcYYY4wxxhhjrAFwwxxjjDHGGGOMMcYYYw2AG+YYY4wxxhhjjDHGGGsA3DDHGGOMMcYYY4wxxlgD4IY5xhhjjDHGGGOMMcYaADfMMcYYY4wxxhhjjDHWALhhjjHGGGOMMcYYY4yxBtCoG+ZycnKwadMm5ObmymlZWVnYtGkTysrKrFqGrflZ05Kamopt27bV6zqPHz+Offv21es6DQznTH5+frVpjDHGGHOM4uJibNq0CdevX5fTbP0tbiq/3Q29Hbt378apU6dqTGOMMcYcSSAiauhC1Nbvv/+O4cOHY+vWrRg0aBAA4PPPP8fMmTORnp6OmJgYAMCVK1dw9uxZ9O/fH66urkbLMJefNR9z587FF198gaKionpb59ixY3HhwgWcPHnSqvynT59GVlaW2WmtWrVCq1atrF634ZzZuXMn+vXrZzHNoLi4GKmpqVAoFIiLi4OXl5fJMgsKCnDgwAH5s0qlgo+PD1q1agVfX1+ry2aJTqfDpUuXkJWVhbCwMERHR0OhUNQq7/bt26HRaGpcZ3R0NFq3bm11GbVaLdLT03Hjxg2EhIQgNjYWSqXSbF6NRoMLFy4gPz8fUVFRCA8Pt3odlhqRY2JiEB8fb5JuzfdXGRHhzJkzyM/PR/v27ev8/aWnp+PKlSto2bKlTccpYF3ZL168iIsXL1pcRuvWrREdHQ2gdvvPGrX9Piu7dOkSzp07h/DwcLRv395ueQH9d7p161ZIkoRevXrB29vb5vLVZXm2fEcGtpxPOp0OJ0+eREFBAeLi4mq1/6uydR9XVZfjvi7lsfV8r8reccwR51xmZiYyMzMxYMAAm+d1JidPnkSHDh3w5ZdfYsqUKQDM/xbfunULx44dQ8+ePeHj42O0jOp+uxsTW7ejap2jMldXV/Tv39+m9cfExKBfv3746quvqk0DAEmSkJGRgRs3biA6OtpivNm/fz8KCwsBAKIowtPTEyEhIYiOjoYgCDaVz5zs7GxkZGRAoVAgPj6+2jhcXd4LFy4gIyOjxvUplUr5OqsmJ0+eNGpwNlCpVBg4cKBJurX7tDrXrl1Deno6AgMDERsbCxcXF4t5CwsLcerUKSiVSnTs2LHavLYs1xZnz56VfyPatWtn1TzW7Kfy8nLs2LHD4jLc3d1x991316rMtuw3c27evImMjAwolUq0atUKfn5+FvOq1Wrs378f5eXlGDx4METRtB+RrcdZdey9327cuIG0tDSEhoaa1AFsWZcjv09r9nF16no81KU8dTkvrY03dTm+anN+y6gR27BhAwGgrVu3ymm//fYbJSUl0fXr1+W0xYsXEwC6fPmyyTLM5WfNxxNPPEGenp71us5XXnmFHn30UavzP/LIIwSABg0aRElJSUZ/K1eutGndBw8epKSkJDpx4oScZjiPdu7cKaft37+fxo4dSyqVijp37kwdOnQgV1dXmjt3LhUWFhotc+/evQSAWrduTUlJSXTPPfdQ586dyc3NjRITE+mDDz4gjUZjUzkNPvvsM4qIiKCgoCDq06cPxcTEUGBgIM2bN4/y8vJszjtmzBij/RcVFUUAqG/fvkbpS5YssbqMX3/9NUVFRVFISAj17duXQkNDKTIykr788kujfLdv36annnqKWrRoQR07dqRu3bqRm5sb9enThw4cOFDjenJzcwkAxcXFmRwHn3/+uVFeW74/g9WrV1NERATFxMRQ//79KSIigp566inSarVW7wuDa9eu0T333EM+Pj509913k5+fH/Xv35+uXr1a47y2lH3lypUm+yIpKYlatmxJAIy+R1v2nzXq+n1KkkTvv/8+9enThyIiIggAPfLII3XOW9UHH3xAAAgAHTx40IYttM/ybPmOiKw/n4iIfvzxRwoLC6PQ0FDq3r07ubu709ixYyk/P9/m7arLPjaoy3Ffl/LU5nyvyhFxzN7nHBHRW2+9RYGBgbWa15mcOHGCABjtX3O/zz/++CMBoL1795osw1z+xsjW7aha56j8N3HiRJvXP3nyZHrrrbeM0qKjo2ny5Mny56KiInrhhRcoNDSUQkNDqU+fPuTr60s9evSgffv2mSyzU6dO5O7uLperT58+FBoaSr6+vjR9+nS6ePGizeUkIjp9+jQNGzaMVCoVdezYkXr06EEuLi40bNgwozqctXk///xzo/3Xt29fAkBRUVFG6SNHjrS6jA8++CC5urqafDdjxowxymfrPjXn66+/pk6dOlFkZCT17duXwsLCKCgoiN5++22SJMkk/8KFC8nd3Z26dOlCbdu2pYCAAPruu+/qvFxbZGVlUUBAAAGgGTNm1Jjflv10+/Zts7+3PXv2JADUpUuXWpXZ2v1mTmZmJt13333k4uJC3bt3p44dO5KLiwtNmzaNCgoKjPLu2bOHpkyZQi1atJD3UWlpqdnlWnucWcNe+82wrd7e3tS3b1/q0KED9ezZk44fP16rdTni+7RlH1tSl+OhLuWpy3lpa7yp7fFl6/ldVZNrmDOnuoY51rw1RMOcrQwNc9ZeYNnKXMPc008/Tb169aKTJ0/Kadu3byc3NzcaP3680fyGSvK7775rlF5aWkqLFy8md3d3GjRokM2B/+uvvyYA9Le//Y3Ky8vl9LVr11JQUJBRo4AteSubP38+AaDz58/bVDaDlJQUEkWRJk6cSGq1moiIysvL6S9/+QsJgkCHDh2S86amptKyZcuouLhYTsvKyqJ27dpRYGCgSQWlKsNF7htvvFFjuWz5/oiIVqxYQS4uLkY/rBqNhhYvXmzz9yZJEvXp04dat25NN2/eJCJ95SIxMZF69OhR44+nrWU3p23btuTm5kY5OTlymi37zxp1/T51Oh09/fTTtHv3biooKKi24cWWvJWdPn2a3Nzc5EpcXRvm7Lk8c9+RLefTgQMHSKlU0sMPPyw3/J8/f57CwsJoxIgRNpentvvYoK7HfV3KU9dzxlFxzN7nHFHTbpgzp7qGuebKUp3Dnqo2zJ05c4a8vLxo6dKl8s2qvLw8uueee8jT05PS09ON5u/UqRPFxcWZLPfAgQPUtWtX8vT0pF27dtlUpsLCQgoPD6eEhATKyMiQ0zMzM2nIkCFGF4G25K3s/PnzBIDmz59vU9kqe/DBB6lly5Y15rN1n5rzn//8hy5cuCB/1mq19OKLLxIAWrZsmVHeFStWEABavXq1nPbaa6+RQqEwuTC3Zbm2GjZsmPwbas2Fuz320/vvv08A6L333rO5vLbsN3MMDSBnzpyR0/744w8SRZFmzZpllHfZsmW0atUqys/Pp8mTJ9fYMGfNcVYXtuy3mzdvUmRkJE2cONHod/Dw4cP0xx9/2HVddfk+bdnH5tT1eKhLeepyXtp6HtX2+LL1/K6qVg1zkiRRamoq7dy50+KdYLVaTYcOHaKdO3ea7Y2WmZlJGzdupPLycpIkiY4dO0Z79uwxquxVdfv2bdq1axelpqYSkfmGuatXr9LGjRvlL/XcuXM0d+5cAkDffvstbdy4kTZu3ChfDFTN7+htqEqSJDp37hzt2rWLMjMzjSrvubm5RmU12Lx5s8mdsaysLNq4caO87lu3bsnbumnTJtqzZ4/Jcqpug06no5SUFNq7dy+VlZWZLe/GjRuNKunWbF9djxVr81X9Pk6cOEGbN282ynPx4kXauXMn3bhxg4gsN8wVFxfTsWPHaN++fWb3mzlnzpyR9/nmzZvpyJEjZvfjsWPHTCrZhw4
"text/plain": [
"<Figure size 1240x300 with 8 Axes>"
]
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"def swap_conditions(source_features: np.ndarray, donor_features: np.ndarray) -> np.ndarray:\n",
" swapped = np.array(source_features, copy=True)\n",
" donor_condition = donor_features[0, CONDITION_INDICES]\n",
" swapped[:, CONDITION_INDICES] = donor_condition\n",
" return np.ascontiguousarray(swapped, dtype=np.float32)\n",
"\n",
"\n",
"source_case_id = split_manifest[\"val_ids\"][0]\n",
"donor_case_id = split_manifest[\"test_ids\"][0]\n",
"source = load_case(source_case_id, max_points=FAST_SAMPLE_POINTS)\n",
"donor = load_case(donor_case_id, max_points=1)\n",
"base_pred = predict_targets(source[\"features\"])\n",
"swapped_features = swap_conditions(source[\"features\"], donor[\"features\"])\n",
"swap_pred = predict_targets(swapped_features)\n",
"delta = swap_pred - base_pred\n",
"\n",
"fig, axes = plt.subplots(1, len(TARGET_NAMES), figsize=(3.1 * len(TARGET_NAMES), 3.0), sharex=True, sharey=True)\n",
"x = source[\"features\"][:, FEATURE_INDEX[\"x\"]]\n",
"y = source[\"features\"][:, FEATURE_INDEX[\"y\"]]\n",
"for ax, name in zip(axes, TARGET_NAMES):\n",
" idx = TARGET_INDEX[name]\n",
" scatter = ax.scatter(x, y, c=delta[:, idx], cmap=\"coolwarm\", s=12)\n",
" ax.set_title(f\"Δ {name}\", fontsize=9)\n",
" ax.set_aspect(\"equal\", adjustable=\"box\")\n",
" ax.set_xticks([])\n",
" ax.set_yticks([])\n",
" fig.colorbar(scatter, ax=ax, fraction=0.046, pad=0.02)\n",
"fig.suptitle(f\"condition swap: coords {source_case_id} + conditions {donor_case_id}\", y=1.02)\n",
"fig.tight_layout()\n",
"\n",
"print(\"source condition:\", case_condition_frame([source_case_id])[0])\n",
"print(\"donor condition: \", case_condition_frame([donor_case_id])[0])\n",
"print(\"mean absolute target delta:\", {name: float(np.mean(np.abs(delta[:, idx]))) for name, idx in TARGET_INDEX.items()})\n"
]
},
{
"cell_type": "markdown",
"id": "8b130835",
"metadata": {},
"source": [
"## 13. Embedding nearest neighbors\n",
"\n",
"For each non-train case, find the closest train condition embedding by cosine distance. This helps separate “bad because far from training conditions” from “bad despite nearby train conditions.”\n"
]
},
{
"cell_type": "code",
"execution_count": 13,
"id": "4d6f31e6",
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"{'case_id': 'airFoil2D_SST_83.526_0.275_3.114_4.789_1.0_12.835', 'split': 'val', 'nearest_train': 'airFoil2D_SST_73.459_-1.971_2.825_4.417_0.0_16.891', 'cosine_distance': 0.0287359356880188, 'sampled_norm_mse': 0.10718367993831635, 'worst_channel': 'turbulent_viscosity'}\n",
"{'case_id': 'airFoil2D_SST_89.338_1.322_3.719_4.337_14.632', 'split': 'val', 'nearest_train': 'airFoil2D_SST_87.005_2.047_1.955_5.456_13.769', 'cosine_distance': 0.07479166984558105, 'sampled_norm_mse': 0.036494314670562744, 'worst_channel': 'velocity_x'}\n",
"{'case_id': 'airFoil2D_SST_54.799_-2.718_5.872_6.096_7.675', 'split': 'val', 'nearest_train': 'airFoil2D_SST_61.838_-0.434_5.558_2.045_17.601', 'cosine_distance': 0.14915579557418823, 'sampled_norm_mse': 0.06382805109024048, 'worst_channel': 'turbulent_viscosity'}\n",
"{'case_id': 'airFoil2D_SST_52.303_4.27_0.815_7.677_1.0_13.254', 'split': 'test', 'nearest_train': 'airFoil2D_SST_51.458_1.678_0.559_6.45_1.0_13.58', 'cosine_distance': 0.0238950252532959, 'sampled_norm_mse': 0.0225469172000885, 'worst_channel': 'velocity_x'}\n",
"{'case_id': 'airFoil2D_SST_77.362_11.678_4.986_6.817_10.94', 'split': 'test', 'nearest_train': 'airFoil2D_SST_66.114_8.17_6.33_4.273_10.506', 'cosine_distance': 0.19597971439361572, 'sampled_norm_mse': 0.3469383120536804, 'worst_channel': 'turbulent_viscosity'}\n"
]
}
],
"source": [
"def cosine_distance_matrix(a: np.ndarray, b: np.ndarray) -> np.ndarray:\n",
" a_norm = a / np.maximum(np.linalg.norm(a, axis=1, keepdims=True), 1e-12)\n",
" b_norm = b / np.maximum(np.linalg.norm(b, axis=1, keepdims=True), 1e-12)\n",
" return 1.0 - a_norm @ b_norm.T\n",
"\n",
"\n",
"train_ids = split_manifest[\"train_ids\"]\n",
"probe_ids = split_manifest[\"val_ids\"] + split_manifest[\"test_ids\"]\n",
"train_mask = np.asarray([case_id in train_ids for case_id in all_case_ids])\n",
"probe_mask = np.asarray([case_id in probe_ids for case_id in all_case_ids])\n",
"train_embeddings = condition_embedding[train_mask]\n",
"probe_embeddings = condition_embedding[probe_mask]\n",
"train_case_ids = [case_id for case_id in all_case_ids if case_id in train_ids]\n",
"probe_case_ids = [case_id for case_id in all_case_ids if case_id in probe_ids]\n",
"distances = cosine_distance_matrix(probe_embeddings, train_embeddings)\n",
"\n",
"rows = []\n",
"for row_index, case_id in enumerate(probe_case_ids):\n",
" nearest_index = int(np.argmin(distances[row_index]))\n",
" case_small = load_case(case_id, max_points=FAST_SAMPLE_POINTS)\n",
" pred_small = predict_targets(case_small[\"features\"])\n",
" loss, per_channel = normalized_mse(pred_small, case_small[\"targets\"])\n",
" rows.append(\n",
" {\n",
" \"case_id\": case_id,\n",
" \"split\": split_name_for_case(case_id),\n",
" \"nearest_train\": train_case_ids[nearest_index],\n",
" \"cosine_distance\": float(distances[row_index, nearest_index]),\n",
" \"sampled_norm_mse\": loss,\n",
" \"worst_channel\": max(per_channel, key=per_channel.get),\n",
" }\n",
" )\n",
"\n",
"for row in rows:\n",
" print(row)\n"
]
},
{
"cell_type": "markdown",
"id": "f99696d4",
"metadata": {},
"source": [
"## 14. Channel-specific failures\n",
"\n",
"The recorded run metrics already say turbulent viscosity is hardest. This cell rechecks sampled cases and inspects target-scale `turbulent_viscosity` behavior.\n"
]
},
{
"cell_type": "code",
"execution_count": 14,
"id": "57bd29bc",
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"recorded final test per-channel MSE: {'pressure': 0.2227493229439483, 'turbulent_viscosity': 0.34920740170930153, 'velocity_x': 0.13785394670886558, 'velocity_y': 0.17241358693725012}\n",
"sampled mean per-channel MSE: {'velocity_x': 0.08825000996390979, 'velocity_y': 0.05321387698252996, 'pressure': 0.049152967849901565, 'turbulent_viscosity': 0.22080567106604576}\n",
"negative predicted nut fraction: 1.82%\n"
]
},
{
"data": {
"image/png": "iVBORw0KGgoAAAANSUhEUgAABEIAAAFyCAYAAADxgH9zAAAAOnRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjExLjEsIGh0dHBzOi8vbWF0cGxvdGxpYi5vcmcvctoD+AAAAAlwSFlzAAAPYQAAD2EBqD+naQAAhFRJREFUeJzs3XVYVdn7NvD70KKE0iiCYrdioY7dgT12x8yIzoyJjt2Oo19bsQtrjLEbA2wMVERsJSxAumO9f/hjvx7KAwIHOPfnurhmztrP2vtZ+2zh8LD3WjIhhAARERERERERkQpQU3YCRERERERERER5hYUQIiIiIiIiIlIZLIQQERERERERkcpgIYSIiIiIiIiIVAYLIURERERERESkMlgIISIiIiIiIiKVwUIIEREREREREakMFkKIiIiIiIiISGWwEEJEREREREREKoOFEKIs8vLygkwmw/79+3Nsn6GhoZDJZFiyZEmO7bOgWbJkCWQyGSIjIzNtywsuLi6QyWTw8vLK0+PmlqlTp0Imkyk7DQD5K5ecoqzrlIgKttjYWMhkMsyZMyfH962jo4NJkybl+H7zSkHJPz/kWa5cOfTq1UuuLa/zSn283Ly2s5IHUWZYCCEiyiZnZ2fIZDLIZDKcO3cu3ZhmzZpBJpPB2Ng4zbarV6+ic+fOsLKygq6uLipVqoSRI0fCw8Mjw+Ok97V27dpcGR8RkSqTyWSYOnWqstPI13iO0irI56Qg5F4QcqSCQUPZCRARZWTq1KkF4odd0aJFsWPHDrRr106u/fXr13B3d0fRokXT9HFxccGgQYMwcOBAXLx4EdbW1vD29sasWbNQv359JCQkQEND/lv0hQsX0Lp161wdCxERkbLFxsYqOwWF5Nc88zqv/HIe8kseVDDwjhAioh/Uo0cPHD16FGFhYXLtO3bsgKmpKZo0aZKmz6JFi1C6dGns3LkTFStWhI6ODurUqYOTJ09i5syZhe7xESIiIiKi/IKFEFKKpKQkLFiwAJUqVYKuri6srKzQr18/PHv2TIoJCgqSu/1fW1sbFStWxNy5c5GQkCDFXbt2DTKZDCdPnsTatWthY2MDPT099OzZE1++fAEArF69GmXKlIGOjg5at26Nd+/eyeXz7T5Wr14Na2trFClSBPb29nB3d1doTNHR0Zg+fTrKly8PbW1tmJqaYvjw4fj8+bNc3Nu3b9G9e3cUK1YMRkZGcHR0RFxcnELH+DbPrVu3wtbWFjo6Oqhfvz5u3ryZJv7jx48YOXIkLCwsoKWlhTJlymDq1KmIiYlJd58bNmxAuXLloK6ujrt370rzOURFRWHEiBEoXrw4TE1NMW/ePABAZGQkRowYASMjIxgaGmLs2LFITEyUy2HHjh1y76Oenh4aNWqEI0eOfHe8qedeePnyZaaPiPj7+0t9X758iQEDBsDMzAxaWlooX748lixZguTkZLljnD17FnZ2dtDR0YGtrS22bt2q0HvxrX79+kEIITdvjBACO3fuxKBBg9Lc2QEAX758gbm5OdTU0n4bnjdvHtTV1bOcx/dERkZi+PDhMDQ0hIGBAfr164ePHz9K23fu3AmZTIYrV66k6btv3z7IZDKcP38+02O8ePECgwcPhqWlJYoUKYIaNWrA2dkZSUlJcnGxsbFwdHSEkZER9PT00Lt3bwQFBcnFKHrtpFyniuwzK7GA4tcRERUeb9++lYrRf//9t/Q9qFu3bgD+/8+iLVu2pOlrY2ODvn37Sq+/nSvh3LlzsLOzg7a2NpydneX6nTp1CtWrV4eOjg4qV66MXbt2yW3PyjEzIoTAhg0bULt2bRQpUgQGBgbo3Lmz3HxY3+Z76dIl1KlTBzo6OqhQoQL27t2r8DlKLTk5GTY2Nhne1Vi1alU0btxYep3eXA+7d+9G3bp1oa+vDxMTE7Rt2xZXr15NM0ZnZ2fUq1cPRYsWhZmZGfr06YMXL17IjXHGjBmwtbWFtrY2zM3NMXToULx//z7Lx/s2z8zOSXJyMmxtbdG8efN0x1+pUiU0atQo3W0pIiIi8Ntvv8HExAR6enro2rUrPnz4kG5sVs/f997P713Hmc3NkVPXtiLXXHp5KPJ+K3rdU+HCQggpxYIFC7B06VKsWLECgYGB8PDwQI8ePbBy5UopxtjYGEII6evjx49YsGABVqxYgZkzZ6bZ544dOxAeHo67d+/ixo0buH//PoYPH441a9YgLCwMHh4euHfvHt68eYOhQ4emm9fWrVulfLy9vWFpaYk2bdrg7t27mY4nLi4OrVq1wu7du7Fy5UoEBQXh8uXL8PHxQdOmTaVf5L98+YKffvoJr169wtWrV/H69Ws0bNgQEyZMyNL527NnD/z8/HDjxg08e/YMOjo66Nq1K6Kjo6WY0NBQNG7cGFevXsWhQ4cQFBSElStXYvPmzejYsWOaX0y3bduGgIAAuLu748qVKyhSpIi0bcKECejevTvevXuHtWvXYu7cudiyZQtGjRoFBwcHvHnzBlu3bsXGjRuxatUquf0OHTpUeg+TkpLg4+ODli1bonfv3mk+UHxPuXLl5K4JIQRev34NU1NTlCxZEoaGhgCAJ0+eoG7duvj8+TMuXLiAL1++YMWKFVi+fDnGjBkj7e/ixYvo3LkzatasiefPn8Pd3R2PHj3CoUOHspRX8eLF0bVrV+zYsUNqc3V1ha+vb4bXWpMmTXDv3r0sH+tHjB07Fp06dYKvry/OnDkDDw8PtGjRQrpu+vTpAyMjI6xfvz5N3/Xr16N06dKZPprz6NEj1K1bF2/evMHRo0cRGBiIvXv34tGjR3jw4IFc7IQJE9CqVSu8efMGp06dwuXLlzF27Fi5mKxeO4rsMyuxil5HRFS42NjYQAgBAHBycpK+Dx09ejTb+7x//z62b9+OAwcOwNvbG2XLlpW23bt3Dy4uLjh27Bh8fX3Rs2dPDBkyBNu3b//RocgZPXo0Jk+ejDFjxsDf3x9eXl4oUaIEGjVqhOfPn8vFPnz4EDt27MChQ4fg7++P5s2bY+DAgfD29gaQ9XOkpqaGYcOG4dKlS3jz5o3ctlu3bsHb2xsjRozIMPdTp05h8ODBGD58OHx9ffHy5UtMnToVy5Ytk4sbPnw4JkyYgIEDB+L58+d48uQJevfuLc2lJYRAt27dsHbtWvz9998IDAzEsWPHcPv2bdjb20sFcUWP963Mzomamhp+++03XL16VTqHKVxdXfHs2bNMx5+cnAwHBwccPnwYO3fuREBAAP7880+MHj06zR+gsnP+FH0/M7uO05OT13Z2/l0q+n6n+N51T4WMIFKC5s2bixYtWmSr74wZM0Tx4sWl1+7u7gKA6Natm1zc6tWrhUwmEz///LNc+7p16wQA8erVqzT7aNeunVxsTEyMMDc3F23btpXaHj9+LACIffv2SW2rVq0SAIS7u7tcf19fX6GpqSlWrFghhBBi7ty5QiaTCW9vb7m4OXPmCABi8eLFmY49JU8HBwe59tu3bwsAYs+ePVLbvHnzBABx+/Ztudjt27cLAOLgwYNy+2zTpk2a4zk5OQkAYsuWLXLtbdu2FUWLFhXOzs5y7R06dBCVK1fOdAwpatWqJQYOHCi9Xrx4sQAgIiIiMm371pcvX0SlSpWEnp6e8PT0lNrbtGkjSpYsKSIjI+Xit2zZImQymXj+/LkQQoh69eqJChUqiKSkJLm4unXrCgDi8ePHmY5hw4YNAoC4efOmOH36tAAgfHx8hBBC9O/fX9SrV08IIUSnTp2EkZGRXN+AgADRqFEjAUBYWlqKXr16icWLF8uNI/VxMvr6Xp4p7+OaNWvk2q9duyYAiNWrV0ttkydPFpqamuLDhw9SW8o1P3v27EyP07x5c2FqairCw8O/m8umTZvk2mfOnCnU1NRESEhIpscQIu21k5V9ZiVW0evoe9cpERVMAISTk1Oa9hcvXggAYvPmzWm2WVtbiz59+kivY2JiBABhbm4uYmNj5WJTtllaWoq4uDi5ba1btxZmZmYiISEhy8cUQghtbW0xceJE6fX169cFAOnzSIr4+HhRtmxZMWDAALmcrK2tRXx8vBQXEREhdHV
"text/plain": [
"<Figure size 1100x380 with 2 Axes>"
]
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"channel_rows = []\n",
"for case_id in [split_manifest[\"train_ids\"][0], *split_manifest[\"val_ids\"], *split_manifest[\"test_ids\"]]:\n",
" case = load_case(case_id, max_points=FAST_SAMPLE_POINTS)\n",
" pred = predict_targets(case[\"features\"])\n",
" loss, per_channel = normalized_mse(pred, case[\"targets\"])\n",
" channel_rows.append({\"case_id\": case_id, \"split\": split_name_for_case(case_id), \"loss\": loss, **per_channel})\n",
"\n",
"means = {name: float(np.mean([row[name] for row in channel_rows])) for name in TARGET_NAMES}\n",
"fig, axes = plt.subplots(1, 2, figsize=(11, 3.8))\n",
"axes[0].bar(list(means), list(means.values()))\n",
"axes[0].set_yscale(\"log\")\n",
"axes[0].set_title(\"sampled normalized MSE by channel\")\n",
"axes[0].tick_params(axis=\"x\", rotation=35)\n",
"axes[0].grid(axis=\"y\", alpha=0.25)\n",
"\n",
"nut_idx = TARGET_INDEX[\"turbulent_viscosity\"]\n",
"nut_truth = PREDICTION_RESULT[\"case\"][\"targets\"][:, nut_idx]\n",
"nut_pred = PREDICTION_RESULT[\"predictions\"][:, nut_idx]\n",
"axes[1].hist(nut_truth, bins=40, alpha=0.55, label=\"truth\")\n",
"axes[1].hist(nut_pred, bins=40, alpha=0.55, label=\"prediction\")\n",
"axes[1].set_title(\"turbulent viscosity distribution\")\n",
"axes[1].set_xlabel(\"target-scale nut\")\n",
"axes[1].legend()\n",
"fig.tight_layout()\n",
"\n",
"print(\"recorded final test per-channel MSE:\", final_metrics[\"test_mse_per_channel\"])\n",
"print(\"sampled mean per-channel MSE:\", means)\n",
"print(f\"negative predicted nut fraction: {np.mean(nut_pred < 0):.2%}\")\n"
]
},
{
"cell_type": "markdown",
"id": "b39e9290",
"metadata": {},
"source": [
"## 15. Model surgery: zero one block's FiLM modulation\n",
"\n",
"Run the same trunk while setting `gamma=0` and `beta=0` for one residual block. Large deltas indicate that block's condition modulation matters for the selected points.\n"
]
},
{
"cell_type": "code",
"execution_count": 15,
"id": "3cc282ea",
"metadata": {},
"outputs": [
{
"data": {
"text/plain": [
"[{'block': 0,\n",
" 'velocity_x': 7.872861385345459,\n",
" 'velocity_y': 6.204006195068359,\n",
" 'pressure': 264.695068359375,\n",
" 'turbulent_viscosity': 0.00039422939880751073},\n",
" {'block': 6,\n",
" 'velocity_x': 1.0614341497421265,\n",
" 'velocity_y': 1.109315276145935,\n",
" 'pressure': 62.921085357666016,\n",
" 'turbulent_viscosity': 3.732619734364562e-05},\n",
" {'block': 11,\n",
" 'velocity_x': 1.616136908531189,\n",
" 'velocity_y': 1.489053726196289,\n",
" 'pressure': 63.22951889038086,\n",
" 'turbulent_viscosity': 1.955496918526478e-05}]"
]
},
"execution_count": 15,
"metadata": {},
"output_type": "execute_result"
},
{
"data": {
"image/png": "iVBORw0KGgoAAAANSUhEUgAAA0gAAAFyCAYAAADPmfo5AAAAOnRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjExLjEsIGh0dHBzOi8vbWF0cGxvdGxpYi5vcmcvctoD+AAAAAlwSFlzAAAPYQAAD2EBqD+naQAAYHFJREFUeJzt3XdUFNfbB/DvLmXpSBEFRVAJ9pKC2FCx19iILSqWGGyxhcRYkmjUYIuaoEl+0cTeO/YIgr1GYxQURZogItIVqXvfP4R9XSku68qKfD/neM7unTt3nhlndJ69d+5IhBACREREREREBKm2AyAiIiIiInpbMEEiIiIiIiLKxwSJiIiIiIgoHxMkIiIiIiKifEyQiIiIiIiI8jFBIiIiIiIiyscEiYiIiIiIKB8TJCIiIiIionxMkIiIiIiIiPIxQSKiYq1atQpOTk7Q1dVF1apVAQAZGRkYN24c7OzsIJFIMGLECO0GWUq+vr6F9ultd/v2bUgkEmzatEnbobwVLly4AIlEgn379r2y7ps+dqWJpSytW7cOEokEYWFh2g6FirFixQpIJBI8fPhQ26EQ0UuYIBFRkS5evIiJEyfi22+/RWZmpuI/8eXLl2PLli0ICAiAEALr1q3T+LYbNmyInj17arzdM2fOYNKkSfjuu++QlZXFGxMqM5GRkZBIJDh69Ohrt/Xw4UNIJJJi/2RmZha53r59+16ZLBbUkUgkWL9+fZF1Bg4cqKiTm5v72vtTkTRs2BDjxo17Y+3//vvvkEgkiIyMfGPbIKoImCARUZECAwMBAIMHD4aurq5S+UcffYR69eppKzS1nTx5EgAwaNAg6OjoaDka1dWtWxdCCAwdOlTbodBbZMKECRBCFPpjYGCAESNGQAgBJycntdo2NjYu8seP5ORk7N+/H8bGxq8ZfcUTHh6O4OBg9O7dW9uhENErMEEioiI9evQIOjo60NfXL1RuaGiopahez+PHj4vcJyJS1q9fP5w8ebJQT8TWrVshlUrRvXt37QRWju3fvx+mpqZwd3fXdihE9ApMkIgqoP3796NNmzYwNTWFkZERWrdujYCAAABAbm4uJBIJli9fjry8vELDd27cuIFDhw4pvl+4cEGldl/k5+cHd3d3mJubo1KlSujcuTPOnz8PADAxMUFwcLDSNl71K7gQAitXrkSjRo1gYGAAS0tL9O7dGzdu3FDapxUrVijt05w5c4psb+HChcUOX2revLnKx7KAk5MTPDw88N9//6F9+/YwNjbGxIkTATxPOD///HPY2dlBX18fNWvWxNdff42MjAzF+kU9R/Pisy9r166Fk5MTDAwM8NFHH+Hs2bOF9ikiIgJ9+/aFqakprKysMG7cOCQkJEAikWD+/PklHl9V4yxtTM+ePcO3334LZ2dnyGQy2NjYYMSIEaUa+rhq1So4OjrCwMAAzZs3V/QSamJ/CmL8/vvvUa9ePRgYGMDBwQETJ07E48ePi2372bNn8PDwgJmZGQ4fPlxsPblcjh9//BH16tWDsbEx7O3tMWjQINy+fVu1nS/B6z6D1KtXL1hZWRUaZrd27Vr0798fZmZmKrd1+vRptGzZEoaGhrC3t8dPP/0Ef39/SCQS+Pv7K+odPXpU6VozMjLCRx99hLVr1yq1V/Dszv379zF16lRUrlwZlpaWmDJlCvLy8pCdnY2pU6fCxsYGJiYmGD58eKG/V020oWq8Bfbv34+uXbtCJpMplWdnZ2PKlCmoXLkyTExM8PHHHyM8PLzQ+rGxsRgzZgyqVasGfX19ODo6YubMmcjKygIAfPPNN4rhezVr1lTEVTCss7TxElVogogqlBUrVgiJRCLmzJkjYmJiREJCgpgzZ47Q0dERR48eVdSbPHmy0NHRKbR+gwYNRI8ePdRud+nSpUIikYivvvpK3Lt3T6Smporjx4+LAQMGvHIbxZk6darQ1dUVP//8s0hMTBS3bt0SHTp0EMbGxuK///575T6pYt26dQKAGD16dKn3uXbt2qJNmzaiW7du4tq1ayIuLk7s3LlTpKamCicnJ1GrVi1x6tQpkZqaKvz8/ISVlZVwc3MTubm5Qgghbt26JQCIjRs3Kto8f/68ACAGDhwoZs+eLeLi4kRUVJRo27atsLKyEunp6Yq6SUlJwt7eXjRq1EhcuXJFpKSkiG3btolhw4YJAGLevHkl7ruqcZYmpqysLNGqVStRrVo14efnJ1JTU0VwcLBo3bq1eO+990RaWlqx8RRsp0+fPmLGjBni4cOHIiIiQnzyySdCX19fXLhwQVG3qGOn6v5kZmaK5s2biypVqojt27eLxMREER0dLVatWiWWLl2qFMvevXuFEELExcWJjz76SNjb24vr168rthkRESEAiCNHjijK5s+fL4yNjcXBgwfFkydPxMOHD8XOnTvFmDFjSvz7iIuLEwDEhAkTiq2zdu1aAUDcvXtXUbZ3795Cx+JlBXX27t0rJk2aJGrVqiXkcrkQQogbN24IAMLf31+MHj1aABA5OTklxnrlyhUhk8lE3759xb1798SjR4/EggULRN++fQUAcfz48SLXk8vl4uHDh2LJkiVCKpWKHTt2KJYtX75cABDDhw8XW7ZsEampqeLIkSPCyMhIzJs3T4wdO1Zs3LhRpKSkiL///lsYGxsLb29vpfY10Yaq8QohRGJiotDR0RGbNm0qFMOgQYPE2rVrRXJysrh27Zpo2rSpsLOzEwkJCYq60dHRwtbWVjRr1kxcvHhRpKeni8DAQFGjRg3Ru3dvRb3ffvtNABAREREl/r28Kl6iio4JElEFEh8fLwwMDJRu8gt07dpVNGnSRPG9NAmSqu0+ePBA6Ovri+HDh5cYZ2kSpIiICCGVSsX48eOVylNSUoS5ublSO+omSAEBAUJPT0+0aNFCPHv2TAhRumNZu3ZtoaurKyIjI5Xq/fjjjwKAOHv2rFL5xo0bBQCxdetWIUTJCVL37t2V1v3nn38EALF+/XpF2bx584REIhE3b95Uqrtw4UKVEiRV4yxNTKtWrRIARGBgoFLd2NhYIZPJxJIlS4qNp2A7HTp0UCrPysoS1apVE+3bt1eUFXXsVN2fFStWCADi2LFjr4xl79694vr168Le3l589NFH4sGDB8WuU6Bjx47Czc3tlfVeVpAgFfXn+++/F0JoJkG6evWqACCCgoKEEEJMmzZNODg4CLlcrnKC1KNHD2FjY6O4bgr06dOnxATpRV27dhUdO3ZUfC9ILObPn69Ub/jw4cLY2FjMmTNHqXzkyJHCwsJCqUwTbagarxBCrF+/Xujq6oqkpKRCMcyaNUupbmhoqJBKpWLmzJmKsmHDhgkzMzPx8OFDpboHDx4UAMTp06eFEKonSK+Kl6ii4xA7ogrk+PHjyMzMxCeffFJoWceOHXH9+nWkpqa+sXb9/f2RnZ2NIUOGqBV/UYKCgiCXy9GvXz+lcnNzc3Ts2BEnTpyAEELt9kNCQtC/f3/UqFEDfn5+MDAwAFD6Y9mkSRM4ODgo1QsICEDVqlXRsmVLpfL+/fsrlr9Kjx49lL43bNgQAJSG6AQGBqJmzZpo0KCBUt2PP/74le2rE6cqMR04cADW1tZo166dUl07OzvUq1dPpaFyL8evr6+Pbt264fTp08jJyXnt/Tl8+DCsra3RuXPnV8Zy5MgRtGrVCi4uLjh58iRsbW1fuU6TJk1w9uxZzJo1Czdu3Cj1eVrUJA3FDRtVx/vvv48mTZpg3bp1yM3NxaZNm+Dp6QmJRKJyG4GBgejYsaPiuilQ1Lknl8uxbNkyfPDBBzA2NlYaIlbUUMFu3bopfa9bty6ePn1aqLxevXpITk5GSkqKRtsoTbwFQ3EtLCwKxfDysXB2dkb9+vVx4sQJRdmBAwfQpk0bVKlSRalu+/btIZFIVLpeSnt8iSoyJkhEFUjBsx09evSArq4udHR0IJVKIZVK4e3tDQBITEx8Y+0+evQIAFCtWjVN7I5SvEW906hq1ap49uxZoWcHVBUfH48ePXpAKpUqbpYLlPZYFrXPiYmJRcZtaGgIc3PzEp9zKfDyjbi+vj5kMpnSjVxiYiJsbGwKrVt
"text/plain": [
"<Figure size 850x380 with 1 Axes>"
]
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"def forward_with_film_surgery(features: np.ndarray, *, zero_film_block: int | None = None) -> np.ndarray:\n",
" features_norm = torch.from_numpy(normalize_features_raw(features)).to(device)\n",
" coordinates = features_norm.index_select(dim=1, index=model.coordinate_indices.to(device))\n",
" conditions = features_norm.index_select(dim=1, index=model.condition_indices.to(device))\n",
" with torch.no_grad():\n",
" hidden = model.input(fourier_features(coordinates, model.fourier_scales))\n",
" unique_conditions, inverse = torch.unique(conditions, dim=0, return_inverse=True)\n",
" condition_embedding_batch = model.condition_encoder(unique_conditions)\n",
" for block_index, block in enumerate(model.blocks):\n",
" gamma_beta = block.film(condition_embedding_batch).index_select(dim=0, index=inverse)\n",
" gamma, beta = gamma_beta.chunk(2, dim=1)\n",
" if zero_film_block == block_index:\n",
" gamma = torch.zeros_like(gamma)\n",
" beta = torch.zeros_like(beta)\n",
" update = block.linear(block.activation(block.norm(hidden)))\n",
" update = update * (1.0 + gamma) + beta\n",
" hidden = hidden + update\n",
" hidden = model.output_norm(hidden)\n",
" hidden = model.activation(hidden)\n",
" return denormalize_targets(model.output(hidden).detach().cpu().numpy())\n",
"\n",
"\n",
"surgery_features = source[\"features\"]\n",
"baseline_manual = forward_with_film_surgery(surgery_features, zero_film_block=None)\n",
"baseline_api = predict_targets(surgery_features)\n",
"assert np.allclose(baseline_manual, baseline_api, atol=1e-4, rtol=1e-4)\n",
"\n",
"block_choices = [0, len(model.blocks) // 2, len(model.blocks) - 1]\n",
"surgery_rows = []\n",
"for block_index in block_choices:\n",
" edited = forward_with_film_surgery(surgery_features, zero_film_block=block_index)\n",
" abs_delta = np.mean(np.abs(edited - baseline_api), axis=0)\n",
" surgery_rows.append({\"block\": block_index, **{name: float(abs_delta[idx]) for name, idx in TARGET_INDEX.items()}})\n",
"\n",
"fig, ax = plt.subplots(figsize=(8.5, 3.8))\n",
"width = 0.18\n",
"x_positions = np.arange(len(block_choices))\n",
"for offset, name in enumerate(TARGET_NAMES):\n",
" ax.bar(x_positions + (offset - 1.5) * width, [row[name] for row in surgery_rows], width=width, label=name)\n",
"ax.set_xticks(x_positions, [f\"block {index}\" for index in block_choices])\n",
"ax.set_ylabel(\"mean absolute target-scale Δ\")\n",
"ax.set_title(\"effect of zeroing one block's FiLM gamma/beta\")\n",
"ax.legend(fontsize=8)\n",
"ax.grid(axis=\"y\", alpha=0.25)\n",
"fig.tight_layout()\n",
"\n",
"surgery_rows\n"
]
},
{
"cell_type": "markdown",
"id": "c8474ce4",
"metadata": {},
"source": [
"## 16. Cleanup\n",
"\n",
"The notebook intentionally avoids claiming scientific validity. It is a compact probe suite for one checkpoint: load, predict, localize errors, inspect the FiLM condition path, perturb conditions, and do small surgery experiments.\n",
"\n",
"Useful next moves after the first clean run:\n",
"\n",
"- Increase `CASE_SAMPLE_POINTS` for smoother maps.\n",
"- Switch `CHECKPOINT_NAME` between `checkpoint_best.pt` and `checkpoint_final.pt`.\n",
"- Set `USE_CUDA = True` if the GPU has enough memory.\n",
"- Replace the example case IDs with specific train/val/test outliers from the nearest-neighbor table.\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
}