{"nbformat":4,"nbformat_minor":0,"metadata":{"colab":{"provenance":[{"file_id":"1mvUfawVVitSfaJpT6JmvzSsGy8JlrwIP","timestamp":1768385467682},{"file_id":"1NXE8w1nzHaG59A5riAHsNmsjda3BLaqw","timestamp":1759336686509},{"file_id":"1C1gL0bU0COqXpQRFRBfF-EdNkpnEOZar","timestamp":1757541772790},{"file_id":"1HqPNMMjNXcX8IzJyfu9V5oOzzZtNZ32f","timestamp":1757344418096},{"file_id":"19ZRywJ3yT4WsFS8OJ0Zyr6ZAeM4ndvQy","timestamp":1757335170914},{"file_id":"1YolB-97MAW_ty298MzULvz_66YS36QU2","timestamp":1757114162968}],"collapsed_sections":["TMBlpf-F9YjP"]},"kernelspec":{"name":"python3","display_name":"Python 3"},"language_info":{"name":"python"}},"cells":[{"cell_type":"markdown","source":["# Utilities for PHYS 2208 (Spring 2026)\n","Created by: Carlos Kaskoun, Ashley Kim, and Natasha Holmes"],"metadata":{"id":"0XtyP78GypK9"}},{"cell_type":"markdown","source":["### Load packages"],"metadata":{"id":"HqUXIofzyzD2"}},{"cell_type":"code","source":["# This must go first (it needs to be ran before any matplotlib package is imported)\n","try:\n"," # Install packages for interactivity\n"," %pip -q install \"matplotlib==3.8.4\" \"ipympl==0.9.3\" \"ipywidgets==8.1.2\" shapely\n","\n"," # Change matplotlib backend\n"," from google.colab import output\n"," output.enable_custom_widget_manager()\n"," %matplotlib widget\n"," print('Utilities sucessfully loaded!')\n","except:\n"," print('Please click Runtime -> Restart Session, choose Yes, and run again.')"],"metadata":{"id":"VStjvndjZZ60"},"execution_count":null,"outputs":[]},{"cell_type":"code","execution_count":null,"metadata":{"id":"JYnXafYwyeDU"},"outputs":[],"source":["import pandas as pd\n","import numpy as np\n","\n","import matplotlib\n","import matplotlib.pyplot as plt\n","import matplotlib.cm as cm\n","from matplotlib.patches import Polygon as MplPolygon\n","from matplotlib.path import Path as MplPath\n","\n","import os\n","import seaborn as sns\n","import re\n","\n","from scipy import stats\n","from scipy.optimize import curve_fit\n","from sklearn.metrics import r2_score\n","\n","import plotly.graph_objects as go\n","import plotly.express as px\n","\n","from pathlib import Path\n","from typing import List, Optional, Dict, Iterable"]},{"cell_type":"markdown","source":["### Define functions to load and parse data"],"metadata":{"id":"jCSstqPfy1nL"}},{"cell_type":"code","source":["def load_data(file_path):\n"," \"\"\"Load CSV file from Google Drive and return as pandas DataFrame\"\"\"\n"," if not file_path.startswith('/content/drive'):\n"," file_path = f'/content/drive/MyDrive/{file_path}'\n","\n"," if not os.path.exists(file_path):\n"," print(f\"File not found: {file_path}\")\n"," print(\"Please check file_path and try again.\")\n"," return None\n","\n"," return pd.read_csv(file_path)\n","\n","def fill_in_missing(df):\n"," return df.interpolate(method='linear').bfill().ffill()\n","\n","# Helper to find bot IDs based on x-columns\n","def _discover_bot_ids(columns, x_prefix):\n"," bot_ids = []\n"," prefix_len = len(x_prefix)\n"," for col in columns:\n"," if col.startswith(x_prefix):\n"," suffix = col[prefix_len:]\n"," bot_ids.append(suffix)\n"," def _sort_key(val):\n"," try:\n"," return int(val)\n"," except ValueError:\n"," return val\n"," return sorted(set(bot_ids), key=_sort_key)"],"metadata":{"id":"zPmVGxPiy7J9"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":["# Code to create dataframes"],"metadata":{"id":"ZAAyZRsXgKvO"}},{"cell_type":"code","source":["def get_bot_position(file_path):\n"," \"\"\"\n"," Create a data frame with just the bot positions at each time stamp\n","\n"," Returns:\n"," x_positions: pandas data frame of each bot's x position at each time stamp\n"," y_ positions: pandas data farme of each bot's y position at each time stamp\n"," \"\"\"\n","\n"," df = pd.read_csv(file_path)\n"," df_filled = fill_in_missing(df)\n","\n"," bot_data_columns = df_filled.drop(['frame_number', 'timestamp'], axis=1)\n","\n"," columns = bot_data_columns.columns\n"," num_bots = (len(columns)) // 3\n","\n"," x_positions = df_filled.iloc[:, [0,1]].copy()\n"," y_positions = df_filled.iloc[:, [0,1]].copy()\n","\n"," for bot_id in range(num_bots):\n"," start = (bot_id * 3)\n"," label = \"Bot \" + str(bot_id+1)\n"," x_positions[label] = bot_data_columns.iloc[:,start].copy()\n"," y_positions[label] = bot_data_columns.iloc[:,start + 1].copy()\n","\n"," return x_positions, y_positions\n","\n","\n","def get_bot_CM(positions):\n"," \"\"\"\n"," Create a data frame of the bot center of mass position at each time stamp\n","\n"," Returns:\n"," positions: pandas data frame of each bot's displacements at each time stamp,\n"," output from get_bot_positions, such that positions[0] = x_positions and positions[1] = y_positions\n"," \"\"\"\n","\n"," x_pos = positions[0]\n"," y_pos = positions[1]\n","\n"," # Create center of mass position\n"," x_positions = x_pos.drop(['frame_number', 'timestamp'], axis=1)\n"," y_positions = y_pos.drop(['frame_number', 'timestamp'], axis=1)\n","\n"," CM_x = x_pos.iloc[:, [0,1]].copy()\n"," CM_y = y_pos.iloc[:, [0,1]].copy()\n"," CM_x[\"CM\"] = x_positions.mean(axis=1)\n"," CM_y[\"CM\"] = y_positions.mean(axis=1)\n","\n"," return CM_x, CM_y\n","\n","def get_bot_displacement(positions, is_CM = False):\n"," \"\"\"\n"," Create a data frame of the bot displacements at each time stamp\n","\n"," Returns:\n"," positions: pandas data frame of each bot's displacements at each time stamp,\n"," output from get_bot_positions, such that positions[0] = x_positions and positions[1] = y_positions\n"," \"\"\"\n","\n"," x_pos = positions[0]\n"," y_pos = positions[1]\n","\n"," x_displacements = x_pos.iloc[:, [0,1]].copy()\n"," y_displacements = y_pos.iloc[:, [0,1]].copy()\n","\n"," bot_data_columns_x = x_pos.drop(['frame_number', 'timestamp'], axis=1)\n"," bot_data_columns_y = y_pos.drop(['frame_number', 'timestamp'], axis=1)\n"," columns = bot_data_columns_x.columns\n"," num_bots = (len(columns))\n","\n"," if is_CM == False:\n"," for bot_id in range(num_bots):\n"," label = \"Bot \" + str(bot_id+1)\n"," x_displacements[label] = bot_data_columns_x.iloc[:,bot_id].diff()\n"," y_displacements[label] = bot_data_columns_y.iloc[:,bot_id].diff()\n"," elif is_CM == True:\n"," for bot_id in range(num_bots):\n"," label = \"CM\"\n"," x_displacements[label] = bot_data_columns_x.iloc[:,0].diff()\n"," y_displacements[label] = bot_data_columns_y.iloc[:,0].diff()\n","\n"," #drop the first row\n"," # x_displacements = x_displacements.drop(index=0)\n"," #y_displacements = y_displacements.drop(index=0)\n","\n"," return x_displacements, y_displacements\n","\n","\n","def get_bot_velocity(positions, is_CM = False):\n"," \"\"\"\n"," Create a data frame of the bot velocities at each time stamp\n","\n"," Returns:\n"," positions: pandas data frame of each bot's displacements at each time stamp,\n"," output from get_bot_positions, such that positions[0] = x_positions and positions[1] = y_positions\n"," \"\"\"\n","\n"," x_displacements, y_displacements = get_bot_displacement(positions)\n","\n"," x_velocity = x_displacements.iloc[:, [0,1]].copy()\n"," y_velocity = y_displacements.iloc[:, [0,1]].copy()\n","\n"," x_velocity['time_interval'] = x_velocity['timestamp'].diff()#.fillna(0)\n"," y_velocity['time_interval'] = y_velocity['timestamp'].diff()#.fillna(0)\n","\n"," bot_data_columns_x = x_displacements.drop(['frame_number', 'timestamp'], axis=1)\n"," bot_data_columns_y = y_displacements.drop(['frame_number', 'timestamp'], axis=1)\n"," columns = bot_data_columns_x.columns\n"," num_bots = (len(columns))\n","\n"," if is_CM == False:\n"," for bot_id in range(num_bots):\n"," label = \"Bot \" + str(bot_id+1)\n"," x_velocity[label] = bot_data_columns_x.iloc[:,bot_id]/x_velocity['time_interval']\n"," y_velocity[label] = bot_data_columns_y.iloc[:,bot_id]/x_velocity['time_interval']\n"," elif is_CM == True:\n"," for bot_id in range(num_bots):\n"," label = \"CM\"\n"," x_velocity[label] = bot_data_columns_x.iloc[:,0]/x_velocity['time_interval']\n"," y_velocity[label] = bot_data_columns_y.iloc[:,0]/x_velocity['time_interval']\n","\n"," return x_velocity, y_velocity\n","\n","def get_bot_speed(positions):\n"," \"\"\"\n"," Create a data frame of the bot speeds at each time stamp\n","\n"," Returns:\n"," speeds\n"," \"\"\"\n"," x_velocities, y_velocities = get_bot_velocity(positions)\n"," x_speed = x_velocities.abs()\n"," y_speed = y_velocities.abs()\n"," return x_speed, y_speed\n","\n","def trim_data(df, start_time, end_time):\n"," df_trimmed = df[(df['timestamp'] >= start_time) & (df['timestamp'] <= end_time)]\n"," return df_trimmed\n","\n","def trim_data_tuple(df, start_time, end_time):\n"," df_x = df[0]\n"," df_y = df[1]\n"," df_trimmed_x = df_x[(df_x['timestamp'] >= start_time) & (df_x['timestamp'] <= end_time)]\n"," df_trimmed_y = df_y[(df_y['timestamp'] >= start_time) & (df_y['timestamp'] <= end_time)]\n"," df_trimmed = [df_trimmed_x, df_trimmed_y]\n"," return df_trimmed"],"metadata":{"id":"-da39iDGgKbf"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":["# Create plots"],"metadata":{"id":"ZQ8lQkNPM6rb"}},{"cell_type":"code","source":["def plot_histogram(df, bot_id, min, max, variable):\n","\n"," \"\"\"\n"," df: pandas data frame of each bot's position, displacement, or velocity at each time stamp,\n"," bot_id: plot one bot at a time by number.\n"," For Center of Mass data, indicate bot_id = \"CM\"\n"," min: minimum horizontal axis value\n"," max: maximum horizontal axis value\n"," variable: position, displacement, velocity\n","\n","\n"," \"\"\"\n"," # Assign dataframes to the x and y histograms\n"," x = df[0]\n"," y = df[1]\n","\n"," # Organize labels with variables being displayed\n"," if bot_id == \"CM\":\n"," label = \"CM\"\n"," else:\n"," label = \"Bot \" + str(bot_id)\n","\n"," x_title = \"x \" + str(variable)\n"," y_title = \"y \" + str(variable)\n","\n"," # Set up the figure\n"," fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(12,5))\n","\n"," # x values histogram\n"," n1, bins1, patches1 = ax1.hist(x[label], bins=26, range=(min, max), color='skyblue', alpha=0.7, edgecolor='black')\n"," ax1.set_title(x_title, fontsize=14)\n"," ax1.set_xlabel(x_title)\n"," ax1.set_ylabel(\"Count\")\n","\n"," # Y values histogram\n"," n2, bins2, patches2 = ax2.hist(y[label], bins=26, range=(min, max), color='skyblue', alpha=0.7, edgecolor='black')\n"," ax2.set_title(y_title, fontsize=14)\n"," ax2.set_xlabel(y_title)\n"," ax2.set_ylabel(\"Count\")\n","\n"," plt.suptitle(str(variable) + \" Histograms for Bot \" + str(bot_id), fontsize=15, fontweight='bold')\n"," plt.tight_layout()\n"," plt.style.use('seaborn-v0_8')\n"," plt.show()\n","\n","\n","def scatter_over_time(df, variable, bots=None, domain=None, yrange=None):\n"," \"\"\"\n"," df: pandas data frame of each bot's position at each time stamp,\n"," bots: indicate which bots you want to plot:\n"," NONE to mean all\n"," Individual numbers in square brackets to indicate individual bugs, e.g., [1,2]\n"," \"CM\" to indicate Center of Mass\n"," variable: position, displacement, velocity\n","\n"," \"\"\"\n"," # Assign dataframes to the x and y histograms based on the variable assigned in the function call\n"," x, y = df\n","\n"," # Organize labels with variables being displayed\n"," x_title = \"x \" + str(variable)\n"," y_title = \"y \" + str(variable)\n"," num_bots = len(x.columns)-2\n","\n"," # Initialize plot\n"," fig, axs = plt.subplots(2, 1, figsize=(10,6))\n"," if bots is None:\n"," for i in range(num_bots):\n"," label = \"Bot \" + str(i+1)\n"," axs[0].scatter(x['timestamp'], x[label], 1.5, label=f'Bot {i+1}')\n"," axs[1].scatter(y['timestamp'], y[label], 1.5, label=f'Bot {i+1}')\n"," elif bots == \"CM\":\n"," label = \"CM\"\n"," axs[0].scatter(x['timestamp'], x[\"CM\"], 1.5, label='CM')\n"," axs[1].scatter(y['timestamp'], y[\"CM\"], 1.5, label='CM')\n"," else:\n"," for i in bots:\n"," # Scatter\n"," label = \"Bot \" + str(i)\n"," axs[0].scatter(x['timestamp'], x[label], 1.5, label=f'Bot {i+1}')\n"," axs[1].scatter(y['timestamp'], y[label], 1.5, label=f'Bot {i+1}')\n","\n"," # Add labels\n"," axs[0].set_xlabel(\"Time\")\n"," axs[0].set_ylabel(\"X \" + str(variable))\n"," axs[0].set_title(\"X \" + str(variable) + \" Over Time\")\n"," axs[0].legend(bbox_to_anchor=(1.03, 1), loc='upper left')\n","\n"," axs[1].set_xlabel(\"Time\")\n"," axs[1].set_ylabel(\"Y \" + str(variable))\n"," axs[1].set_title(\"Y \" + str(variable) + \" Over Time\")\n"," axs[1].legend(bbox_to_anchor=(1.03, 1), loc='upper left')\n","\n"," # Set domain and range if provided\n"," if domain is not None:\n"," axs[0].set_xlim(domain)\n"," axs[1].set_xlim(domain)\n"," if yrange is not None:\n"," axs[0].set_ylim(yrange)\n"," axs[1].set_ylim(yrange)\n","\n"," plt.tight_layout()\n"," plt.show()\n","\n"],"metadata":{"id":"1pDNxDRd3G_P"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":["# Basic stats functions"],"metadata":{"id":"9Des0GYInGeE"}},{"cell_type":"code","source":["def standard_unc_of_mean(data):\n"," n = len(data)\n"," return np.std(data) / np.sqrt(n)\n","\n","def effect_size(A, unc_A, B, unc_B):\n"," return abs(A-B) / np.sqrt(unc_A**2 + unc_B**2)\n","\n","def chiSquared(x, y, dy, f, args):\n"," '''Function Chi-Squared.\n"," x, y and dy are numpy arrays, referring to x, y and the uncertainty in y respectively.\n"," f is the function we are fitting.\n"," args are the arguments of the function we have fit.\n"," '''\n"," return 1/(len(x))*np.sum((f(x, args)-y)**2/dy**2)\n","\n","\n","def linearFit(x, m,b):\n"," return m*x+b\n","\n","def autoFit(x=[], y=[], dy=[], title=\"Use title= in your call\", xaxis=\"Use xaxis= in your call\", yaxis=\"Use yaxis= in your call\", showPlot=True):\n"," x0=[0,0]\n"," x=np.array(x)\n"," y=np.array(y)\n"," dy=np.array(dy)\n","\n"," args,cov=curve_fit(linearFit, x, y, x0, dy)\n","\n"," m=args[0]\n"," b=args[1]\n"," dm = np.sqrt(cov[0][0])\n"," db = np.sqrt(cov[1][1])\n","\n"," if showPlot == True:\n"," fig, ax=plt.subplots(1,1, figsize=(8,6), dpi=75)\n"," line,=ax.plot(x,linearFit(x, m,b))\n"," data=ax.errorbar(x,y, dy, fmt='.k', capsize=5, markersize=10)\n"," ax.set_title(title)\n"," ax.set_xlabel(xaxis)\n"," ax.set_ylabel(yaxis)\n"," plt.show()\n","\n"," print(\"Best fit line : \")\n"," print(\"y=\", round(args[0],4),\"x\", \"+\", round(args[1],4))\n"," print(r\"m={}+\\-{}\".format(m.round(3),dm.round(3),))\n"," print(r\"b={}+\\-{}\".format(b.round(3),db.round(3),))\n","\n"," return m, dm, b, db"],"metadata":{"id":"_XBfwH0gnF1E"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":[],"metadata":{"id":"X33azaexnfOs"}},{"cell_type":"markdown","source":["# Define arena section"],"metadata":{"id":"TMBlpf-F9YjP"}},{"cell_type":"code","source":["def section_data(\n"," csv_path: list[str],\n"," output_dir: str | None = None,\n"," headless_polygons: list[list[tuple[float, float]]] | None = None,\n"," x_cols_prefix: str = \"center_x_\",\n"," y_cols_prefix: str = \"center_y_\",\n"," tracked_cols_prefix: str = \"tracked_\",\n"," show_grid: bool = True,\n"," point_markersize: float = 2.5\n",") -> list[str]:\n"," \"\"\"Load a CSV, let the user draw polygons over a combined bug scatter plot,\n"," and save one CSV per polygon where points outside that polygon are set to NaN/False.\n","\n"," When headless_polygons is provided, skip the interactive UI and apply\n"," those polygons directly.\n"," \"\"\"\n"," try:\n"," from shapely.geometry import Polygon as ShapelyPolygon\n"," _SHAPELY_AVAILABLE = True\n"," except Exception:\n"," _SHAPELY_AVAILABLE = False\n","\n"," # Load CSV\n"," csv_path = csv_path[0]\n","\n"," df = pd.read_csv(csv_path)\n"," all_columns = list(df.columns)\n"," bug_ids = _discover_bug_ids(all_columns, x_cols_prefix)\n"," if not bug_ids:\n"," raise ValueError(f\"No columns beginning with '{x_cols_prefix}' found.\")\n","\n"," # Ensure y and tracked columns exist\n"," for bug_id in bug_ids:\n"," x_col = f\"{x_cols_prefix}{bug_id}\"\n"," y_col = f\"{y_cols_prefix}{bug_id}\"\n"," t_col = f\"{tracked_cols_prefix}{bug_id}\"\n"," if y_col not in df.columns:\n"," df[y_col] = np.nan\n"," print(f\"Warning: {y_col} was missing; created with NaNs.\")\n"," if t_col not in df.columns:\n"," df[t_col] = ~(df[x_col].isna() | df[y_col].isna())\n"," print(f\"Info: {t_col} was missing; derived from non-NaN x/y values.\")\n","\n"," # Helper used by both interactive callback and headless mode\n"," def _split_and_save(polygons: list[list[tuple[float, float]]]) -> list[str]:\n"," if not polygons:\n"," print(\"No polygons were supplied or drawn; no output files will be created.\")\n"," return []\n"," # Prepare output directory\n"," out_dir = output_dir\n"," if out_dir is None:\n"," out_dir = os.path.dirname(os.path.abspath(csv_path)) or '.'\n"," os.makedirs(out_dir, exist_ok=True)\n","\n"," out_paths: list[str] = []\n"," for poly_idx, poly_coords in enumerate(polygons):\n"," if len(poly_coords) < 3:\n"," print(f\"Skipping polygon {poly_idx+1}: fewer than 3 vertices.\")\n"," continue\n"," # clean polygon if shapely available\n"," if _SHAPELY_AVAILABLE:\n"," try:\n"," shapely_poly = ShapelyPolygon(poly_coords).buffer(0)\n"," coords_for_path = list(shapely_poly.exterior.coords)\n"," except Exception:\n"," coords_for_path = poly_coords[:]\n"," else:\n"," coords_for_path = poly_coords[:]\n"," polygon_path = MplPath(coords_for_path)\n"," df_out = df.copy()\n"," for bug_id in bug_ids:\n"," x_col = f\"{x_cols_prefix}{bug_id}\"\n"," y_col = f\"{y_cols_prefix}{bug_id}\"\n"," t_col = f\"{tracked_cols_prefix}{bug_id}\"\n"," xs = df_out[x_col].to_numpy(dtype=float)\n"," ys = df_out[y_col].to_numpy(dtype=float)\n"," pts = np.column_stack((xs, ys))\n"," valid = ~(np.isnan(xs) | np.isnan(ys))\n"," inside = np.zeros_like(xs, dtype=bool)\n"," if valid.any():\n"," inside[valid] = polygon_path.contains_points(pts[valid])\n"," outside = ~inside\n"," df_out.loc[outside, [x_col, y_col]] = np.nan\n"," df_out.loc[outside, t_col] = False\n"," # if original tracked column did not exist, set inside ones True\n"," if t_col not in df.columns:\n"," df_out.loc[inside, t_col] = True\n"," csv_name = os.path.basename(csv_path)\n"," csv_name_truncated, _ = os.path.splitext(csv_name)\n"," out_name = f\"{csv_name_truncated}_section{poly_idx+1}.csv\"\n"," out_path = os.path.join(out_dir, out_name)\n"," df_out.to_csv(out_path, index=False)\n"," out_paths.append(out_path)\n"," # summary\n"," any_inside = np.zeros(len(df_out), dtype=bool)\n"," for bug_id in bug_ids:\n"," x_col = f\"{x_cols_prefix}{bug_id}\"\n"," y_col = f\"{y_cols_prefix}{bug_id}\"\n"," any_inside |= ~(df_out[[x_col, y_col]].isna().any(axis=1))\n"," print(\n"," f\"Polygon {poly_idx+1}: {int(any_inside.sum())} rows contain at least one \"\n"," f\"point inside; output saved to {out_path}\"\n"," )\n","\n"," print(f\"Generated {len(out_paths)} section file(s) in {os.path.abspath(out_dir)}.\")\n"," return out_paths\n","\n"," # HEADLESS MODE: process immediately and return paths\n"," if headless_polygons is not None:\n"," return _split_and_save(headless_polygons)\n","\n"," # INTERACTIVE MODE: register callbacks and return immediately;\n"," # actual splitting happens inside the 'finish' callback.\n"," fig, ax = plt.subplots(figsize=(12, 8))\n"," for bug_id in bug_ids:\n"," x = df[f\"{x_cols_prefix}{bug_id}\"].values\n"," y = df[f\"{y_cols_prefix}{bug_id}\"].values\n"," mask = ~(np.isnan(x) | np.isnan(y))\n"," ax.scatter(x[mask], y[mask], s=point_markersize, c='tab:blue', alpha=0.7)\n"," ax.set_aspect('equal', adjustable='datalim')\n"," if show_grid:\n"," ax.grid(True, linestyle='--', alpha=0.5)\n"," ax.set_title(\n"," \"Draw polygons: click to add vertices, 'z' to undo, 'Enter' to close; \"\n"," \"press Enter again (with no active polygon) to FINISH & SAVE CSVs\"\n"," )\n","\n"," current_poly: list[tuple[float, float]] = []\n"," vertex_labels = []\n"," current_lines = []\n"," final_polygons: list[list[tuple[float, float]]] = []\n","\n"," def _draw_segment(pt1, pt2):\n"," (x1, y1), (x2, y2) = pt1, pt2\n"," line, = ax.plot([x1, x2], [y1, y2], color='red')\n"," current_lines.append(line)\n","\n"," def _remove_last():\n"," if not current_poly:\n"," print(\"Nothing to undo.\")\n"," return\n"," current_poly.pop()\n"," label = vertex_labels.pop()\n"," label.remove()\n"," if current_lines:\n"," line = current_lines.pop()\n"," line.remove()\n"," fig.canvas.draw_idle()\n","\n"," def onclick(event):\n"," if event.button != 1 or event.inaxes != ax:\n"," return\n"," x, y = float(event.xdata), float(event.ydata)\n"," current_poly.append((x, y))\n"," idx = len(current_poly)\n"," label = ax.text(x, y, str(idx), color='red', fontsize=8, va='bottom', ha='left')\n"," vertex_labels.append(label)\n"," if len(current_poly) > 1:\n"," _draw_segment(current_poly[-2], current_poly[-1])\n"," fig.canvas.draw_idle()\n","\n"," def onkey(event):\n"," if event.key == 'z':\n"," _remove_last()\n"," elif event.key == 'enter': # (add 'return' here too if you want Mac 'Return' key)\n"," if current_poly:\n"," if len(current_poly) < 3:\n"," print(\"Polygon requires at least 3 vertices; add more points.\")\n"," return\n"," # close polygon\n"," _draw_segment(current_poly[-1], current_poly[0])\n"," patch = MplPolygon(current_poly, closed=True, fill=True, alpha=0.3, color='red')\n"," ax.add_patch(patch)\n"," for lbl in vertex_labels:\n"," lbl.set_visible(False)\n"," final_polygons.append(current_poly.copy())\n"," current_poly.clear()\n"," vertex_labels.clear()\n"," current_lines.clear()\n"," fig.canvas.draw_idle()\n"," else:\n"," # FINISH: split and save *here* in the callback, then close\n"," print(\"Finished polygon drawing session.\")\n"," fig.canvas.mpl_disconnect(cid_click)\n"," fig.canvas.mpl_disconnect(cid_key)\n"," _split_and_save(final_polygons)\n"," plt.close(fig)\n","\n"," cid_click = fig.canvas.mpl_connect('button_press_event', onclick)\n"," cid_key = fig.canvas.mpl_connect('key_press_event', onkey)\n","\n"," # Show and return immediately; splitting happens later in the callback.\n"," plt.show()\n"," print(\"Interactive mode active: draw polygons now. When done, press Enter with no active polygon \"\n"," \"to save CSVs.\")\n"," return"],"metadata":{"id":"DH6Y_Dw89ajW"},"execution_count":null,"outputs":[]}]}