diff --git a/_toc.yml b/_toc.yml index 803646ae0..4919b0d26 100644 --- a/_toc.yml +++ b/_toc.yml @@ -51,6 +51,7 @@ parts: - file: docs/software/mouse_task_test.md - caption: Dev - Software package Documentation chapters: + - file: docs/software_package/experiment_lifecycle - file: docs/software_package/active_sensing_task - file: docs/software_package/dlc_processor - caption: Experiments - Training protocol and parameters diff --git a/dj_pipeline/gui_transfer/README.md b/dj_pipeline/gui_transfer/README.md index 68e932b91..ca826ed12 100644 --- a/dj_pipeline/gui_transfer/README.md +++ b/dj_pipeline/gui_transfer/README.md @@ -381,11 +381,34 @@ With camera prefix from `IMG_SRC` (default `Imagingsource`), a full session typi The server-side mirror is `vr4mice/actions/populate_rig.py` → `get_files_paths()`. +### Multi-camera rigs + +Some rigs run more than one camera at once (currently 3). Each camera's timestamp/video file +carries its camera index as a suffix directly after the keyword, with no underscore in between: + +| Role | Example filename | GUI key | +|------|------------------|---------| +| Camera timestamps (camera 3) | `TS_vr4mice_Testmouse_2023-02-22_2_CAMERA3.npy` | `camera_path` | +| Video (camera 3) | `vr4mice_Testmouse_2023-02-22_2_VIDEO3.avi` | `video_path` | + +`camera_number_from_filename()` in `utils/session_files.py` extracts that index (`CAMERA\d+` / +`VIDEO\d+`, case-insensitive). `find_related_files()` uses it so autofill stays consistent +across a single camera's files: + +- Pick a numbered `camera_path`/`video_path` file by hand → the sibling of the *other* role is + constrained to the same camera number, instead of grabbing whichever camera sorts first. +- Pick a file with no camera number of its own (DLC, PROC, teensy) → which camera to autofill is + ambiguous, so it defaults to `DEFAULT_CAMERA_NUMBER` (currently `3`) rather than camera 1. +- Single-camera rigs (no `CAMERA`/`VIDEO` suffix anywhere) are unaffected — this logic only + kicks in once more than one camera number is present for a session. + ### How the GUI classifies files 1. **Validation** — `modules/transfer.py` → `_set_path_format()` (glob patterns for the file picker). 2. **Type tag** — `get_type()` scans for keywords: `VIDEO`, `TS`, `DLC`, `PROC`; otherwise `teensy_path`. -3. **Sibling search** — `find_related_files()` lists configured rig folders and keeps files whose stem matches the selected session. +3. **Sibling search** — `find_related_files()` lists configured rig folders, keeps files whose + stem matches the selected session, and (see *Multi-camera rigs* above) disambiguates by + camera number when more than one camera's files are present. ### If formats change @@ -395,6 +418,7 @@ The server-side mirror is `vr4mice/actions/populate_rig.py` → `get_files_paths | Different date format | Stem parsing, auto-fill | Enter mouse/date/attempt manually | | New file category | Not shown in transfer section | Requires new GUI key + populate path | | Mouse names with `_` | Wrong stem split | Avoid underscores in mouse names or update regex | +| Rig's default camera count/index changes | Ambiguous DLC/PROC/teensy pick defaults to the wrong camera | Update `DEFAULT_CAMERA_NUMBER` in `utils/session_files.py`, or select camera/video files by hand | ### Code to update (checklist) @@ -402,7 +426,7 @@ When changing rig naming, edit **together**: | File | What to change | |------|----------------| -| `gui_transfer/utils/session_files.py` | `SESSION_RE`, validation helpers | +| `gui_transfer/utils/session_files.py` | `SESSION_RE`, `CAMERA_NUMBER_RE`, `DEFAULT_CAMERA_NUMBER`, validation helpers | | `gui_transfer/modules/transfer.py` | `_set_path_format()`, `get_type()`, transfer keys | | `vr4mice/actions/populate_rig.py` | `get_files_paths()` | | `tests/unit/test_gui_transfer.py` | Golden filename examples | diff --git a/dj_pipeline/gui_transfer/modules/transfer.py b/dj_pipeline/gui_transfer/modules/transfer.py index f5d661398..2de8f1be7 100644 --- a/dj_pipeline/gui_transfer/modules/transfer.py +++ b/dj_pipeline/gui_transfer/modules/transfer.py @@ -14,6 +14,7 @@ from utils.utils import check_files from utils.session_files import ( PATH_KEYS_FOR_SEARCH, + camera_number_from_filename, dataset_stem_from_filename, find_related_files, ) @@ -89,7 +90,8 @@ def _set_labels(): def _path_is_remote(key): """ Determines whether the specified file type is expected to have a remote path or not. - Currently, it's the case only of video_path-typed file + Currently, it's the case only of video_path-typed file. In this module, + "remote" means the file is not part of the transfer/move set (it stays on rig). Args: key (str): The file type (key) to check. @@ -211,8 +213,10 @@ def get_transfer_files(self, key=None, send=False): Returns: dict or None: The transfer file for the specified key, or all transfer files. """ - if key is not None and key in self.get_keys(): - return self.transfer_file[key] + if key is not None: + if key in self.transfer_file: + return self.transfer_file[key] + return None if send is True: ret = dict() @@ -226,10 +230,17 @@ def get_transfer_files(self, key=None, send=False): def get_processed_files(self): """ Get files that should be moved to processed_path after a successful submit. + + Every file that was actually transferred (i.e. not remote-only, see + _path_is_remote) moves to processed_path once submit succeeds - + that's teensy/dlc/camera/proc plus the GUI-generated gui_output. + video_path is excluded: videos stay on the rig, they're never + transferred, so there's nothing to move. """ ret = list() - for key in ("gui_output", "teensy_path"): - info = self.get_transfer_files(key=key) + for key, info in self.transfer_file.items(): + if _path_is_remote(key): + continue if info: ret.append(info) return ret @@ -402,16 +413,25 @@ def _check_video(self, keys, video_label="video_path"): def _pre_fetch_files(self, filenames, skip_path=None): """ Find sibling session files across configured rig folders. + + On multi-camera rigs, the picked file's camera index (if any, e.g. + "..._CAMERA3.npy" / "..._VIDEO3.avi") constrains which sibling + camera/video file gets auto-filled, so it matches the camera the + user actually selected instead of whichever camera sorts first. """ dataset_stem = dataset_stem_from_filename(filenames) if not dataset_stem: logger.warning(f"Could not parse session from filename: {filenames}") return [] + camera_number = camera_number_from_filename(filenames) + path_by_key = { path_key: config.get_path(path_key) for path_key in PATH_KEYS_FOR_SEARCH } - related = find_related_files(dataset_stem, path_by_key, get_type) + related = find_related_files( + dataset_stem, path_by_key, get_type, camera_number=camera_number + ) skip_resolved = Path(skip_path).resolve() if skip_path else None processed_keys = list() diff --git a/dj_pipeline/gui_transfer/utils/session_files.py b/dj_pipeline/gui_transfer/utils/session_files.py index 11d4d033b..5eb191167 100644 --- a/dj_pipeline/gui_transfer/utils/session_files.py +++ b/dj_pipeline/gui_transfer/utils/session_files.py @@ -9,6 +9,11 @@ from pathlib import Path SESSION_RE = re.compile(r"([A-Za-z0-9]+)_(\d{4}-\d{2}-\d{2})_(\d+)") +CAMERA_NUMBER_RE = re.compile(r"(?:CAMERA|VIDEO)(\d+)", re.IGNORECASE) + +# Rig always has exactly 3 cameras; this is the camera to default to when a +# picked file (e.g. DLC output) carries no camera number of its own. +DEFAULT_CAMERA_NUMBER = 3 PATH_KEYS_FOR_SEARCH = ( "teensy_path", @@ -30,6 +35,18 @@ def dataset_stem_from_filename(filename): return f"{match.group(1)}_{match.group(2)}_{match.group(3)}" +def camera_number_from_filename(filename): + """ + Extract the camera index from a rig filename (e.g. "..._CAMERA3.npy" or + "..._VIDEO3.avi" -> 3). Returns None for filenames with no camera suffix + (single-camera rigs, or non-camera files like DLC/PROC/teensy). + """ + match = CAMERA_NUMBER_RE.search(Path(filename).stem) + if not match: + return None + return int(match.group(1)) + + def parse_session_from_filename(filename): """ Parse mouse name, attempt, and date from a filename. @@ -85,7 +102,7 @@ def check_file_format(key, filename, format_spec, current_mouse=None): return mouse_name, attempt, date -def find_related_files(dataset_stem, path_by_key, get_type_fn): +def find_related_files(dataset_stem, path_by_key, get_type_fn, camera_number=None): """ Find one file per transfer type that belongs to the same session. @@ -93,6 +110,18 @@ def find_related_files(dataset_stem, path_by_key, get_type_fn): dataset_stem: e.g. Testmouse_2023-02-22_2 path_by_key: mapping config key -> directory path string get_type_fn: callable(filename) -> transfer key string + camera_number: if set, on a rig with multiple cameras (files + suffixed "..._CAMERA3.npy" / "..._VIDEO3.avi"), only match + candidate files for that camera index. Files with no camera + suffix (single-camera rigs, DLC/PROC/teensy) are unaffected. + If None (the file that was picked has no camera suffix, e.g. + DLC/PROC/teensy), a role with several different camera numbers + present is ambiguous; DEFAULT_CAMERA_NUMBER is used to resolve + it. If DEFAULT_CAMERA_NUMBER isn't among the candidates, there is + no safe default to autocomplete to, so that role is left out of + the result entirely (no fallback to e.g. the max camera number). + A role where every match shares the same number (or none has a + number at all) is unaffected. Returns: dict mapping transfer key -> Path @@ -100,7 +129,7 @@ def find_related_files(dataset_stem, path_by_key, get_type_fn): if not dataset_stem: return {} - found = {} + candidates = {} seen_dirs = set() for path_key in PATH_KEYS_FOR_SEARCH: @@ -121,8 +150,28 @@ def find_related_files(dataset_stem, path_by_key, get_type_fn): continue if dataset_stem_from_filename(filepath.name) != dataset_stem: continue + file_camera_number = camera_number_from_filename(filepath.name) + if ( + camera_number is not None + and file_camera_number is not None + and file_camera_number != camera_number + ): + continue file_key = get_type_fn(filepath.name) - if file_key not in found: - found[file_key] = filepath + candidates.setdefault(file_key, []).append((file_camera_number, filepath)) + + found = {} + for file_key, matches in candidates.items(): + if camera_number is None: + distinct_numbers = {n for n, _ in matches if n is not None} + if len(distinct_numbers) > 1: + if DEFAULT_CAMERA_NUMBER not in distinct_numbers: + continue + matches = [m for m in matches if m[0] == DEFAULT_CAMERA_NUMBER] + else: + exact = [m for m in matches if m[0] == camera_number] + if exact: + matches = exact + found[file_key] = matches[0][1] return found diff --git a/docs/software/install_dj_pipeline.md b/docs/software/install_dj_pipeline.md index 240350685..80faae022 100644 --- a/docs/software/install_dj_pipeline.md +++ b/docs/software/install_dj_pipeline.md @@ -539,6 +539,12 @@ The rig GUI and **`populate_rig`** on the server assume the same session filenam Classification uses keyword tags (`TS`, `DLC`, `VIDEO`, `PROC`) and glob patterns in `gui_transfer/modules/transfer.py`; parsing lives in `gui_transfer/utils/session_files.py`. The server mirror is `vr4mice/actions/populate_rig.py` → `get_files_paths()`. +**Multi-camera rigs:** on rigs with more than one camera, timestamp/video filenames carry a +camera index suffix (`..._CAMERA3.npy`, `..._VIDEO3.avi`). `camera_number_from_filename()` + +`find_related_files()` in `gui_transfer/utils/session_files.py` keep sibling autofill on the same +camera; when the picked file has no camera number of its own (DLC/PROC/teensy) and several +cameras are present, autofill defaults to `DEFAULT_CAMERA_NUMBER`. Details: `dj_pipeline/gui_transfer/README.md` → *Multi-camera rigs*. + **If naming changes**, update GUI + populate + tests in one change set — patterns are **not** configurable in `config.json`. Full checklist and limitations: `dj_pipeline/gui_transfer/README.md` → *Rig filename contract*. Further GUI module details: `dj_pipeline/gui_transfer/README.md`. diff --git a/docs/software_package/experiment_lifecycle.md b/docs/software_package/experiment_lifecycle.md new file mode 100644 index 000000000..997f1431b --- /dev/null +++ b/docs/software_package/experiment_lifecycle.md @@ -0,0 +1,112 @@ +# What happens when you run a session + +This page is the developer-facing counterpart to the [step-by-step session guide](../software_installation/run_a_session.md). +It walks through what actually happens in the code across a session, from opening the two GUIs to the data landing on disk. + +## The two processes, and what each one owns + +A session always involves two independent processes that never share memory — only a live socket connects them: + +| | **vr4mice** (`teensyexp/teensy_experiment.py`) | **DeepLabCut-live-GUI** (external repo) | +|---|---|---| +| Owns | The Teensy (water/reward, trial logic), the Unity game, trial/session bookkeeping | The camera, DLC pose inference, the photodiode Teensy | +| Talks to | `DLCClient` — reads position data | `MyProcessor_socket`/`dlc_inference_w_pd(_sync)` — a `Listener` that streams position data | +| Saves | Trial/Teensy/Unity data, via the **"Save Task Data"** button | PROC/HDF5/timestamp/video files, via its own **"Stop"/"Save Video"** buttons | + +The socket (`("localhost", 6000)`, `multiprocessing.connection`) only carries **live pose/kinematics data one way**, DLCLiveGUI → vr4mice, +for driving the Unity game in real time. It is not used for control messages, and it is not the authoritative record of anything — +see [Saving the data](#saving-the-data) below for why. + +## Startup + +1. **vr4mice**: `Connect` opens `Teensy(...)` ([teensy.py](../../teensyexp/teensy.py)), which starts a background thread polling the serial + port for reward/lick inputs. `Ready` constructs the task (e.g. `ActiveSensingTask`, [task_active_sensing.py](../../mouse_task/task_active_sensing.py)), + which in turn constructs `DLCClient` ([dlc_deque_socket.py](../../teensyexp/tasks_abc/dlc_deque_socket.py)) — this immediately tries to + connect to `("localhost", 6000)` on a background thread, and opens the Unity build via `UnityEnvironment`. +2. **DLCLiveGUI**: `Init Cam` → `Set Proc` (loads e.g. `dlc_inference_w_pd_sync`, which opens the `Listener` on port 6000 and, if + `use_teensy=1`, its own `TeensyLatency`/`TeensyLatencySync` serial connection to the photodiode Teensy) → `Init DLC`. +3. There's no explicit handshake beyond the `multiprocessing.connection` auth key. In the current implementation, + the task-side `DLCClient` attempts its socket connect once on a background thread during task init; if the + listener is not available yet, that connect attempt fails and the task must be re-initialized to reconnect. + +## During the run + +Every game frame, `UnityTask.loop()` ([unity_task.py](../../teensyexp/tasks_abc/unity_task.py)): +1. Steps the Unity environment and reads back observations/reward. +2. Calls `ActiveSensingTask._get_dlc_on_frame()`, which does `self.dlcClient.read()` to get the latest position/heading from the DLC + processor (used as an input to the agent/game logic). Note `DLCClient.read()` clears its buffer on every call — it's a live sample, + not an accumulating log. +3. On a trial boundary (`self.terminal`), increments `self.episode`; once `self.episode` exceeds the current epoch's trial count + (`epochs` config, default `[250]` — see [common.yaml](../../mouse_task/configs/common.yaml)), advances to the next epoch, or ends the + task if there isn't one. + +Meanwhile, on every camera frame, the DLC processor's `process()` computes position/heading/TTL-signal, buffers it in its own deques +(`self.center_x`, `self.time_stamp`, ...), and streams it to whichever client is connected — tolerating a client that isn't there yet +or has disconnected (see [Failure handling](#failure-handling)). + +## Stopping + +There are two independent stop actions, and reaching one does **not** trigger the other: + +- **The task stops** (250-trial cap reached, or the experimenter hits vr4mice's "Stop"): `run_task_on_thread` exits its loop and calls + `task.stop()` ([teensy_experiment.py](../../teensyexp/teensy_experiment.py)), which for `ActiveSensingTask` closes the Teensy serial + connection, the Unity env, and the `DLCClient` socket/thread. +- **The DLC processor stops**: only when the experimenter hits DLCLiveGUI's own "Stop"/"Save Video" — this closes the photodiode + Teensy and flushes the processor's buffered data to disk. It is not aware of, and does not react to, the vr4mice task stopping. + +This is why the [session guide](../software_installation/run_a_session.md#saving-data) has you stop/save on **both** GUIs, in a specific +order — they are not automatically linked. + +## Saving the data + +Two entirely separate save paths, triggered by two separate manual actions: + +- **vr4mice**: "Save Task Data" → `save_data()` → `task.get_data()` (trial params, Teensy inputs/outputs, Unity states) → + pickled to `/__.pickle`. +- **DLCLiveGUI**: "Stop" + "Save Video" → `on_recording_stopped()` hook on the processor → `save()` (PROC pickle), + `save_legacy_dlc_h5()` (`.h5`), `save_legacy_timestamp_npy()` (`_TS.npy`), plus the GUI's own `.avi` video save. + +Neither side has incremental/periodic autosave — both buffer an entire session in RAM and flush once, on that manual trigger. A crash +or force-quit before that trigger loses whatever hasn't been flushed yet on that side. + +vr4mice does warn about unsaved data at two points: clicking **"Ready"** to initialize a new task (which replaces `self.task`, making +the previous task's in-memory data unreachable) shows a one-time, dismissible reminder if the current task hasn't been saved yet; and +closing the window (via the "Close" button or the window's `[X]`) shows a blocking "did you save?" confirmation if `saved_ok` is still +`False`. Ctrl+C in the terminal is handled separately: it shows a warning dialog telling the experimenter to use the GUI buttons, and +keeps the GUI running. Neither unsaved-data warning is a hard requirement — you can proceed either way — they're just there so an +unsaved session isn't discarded purely by accident. + +## Failure handling + +A few things worth knowing about how this stack behaves under partial failure: + +- **Closing a serial port while a background reader thread is blocked on it** is a known hazard on Windows/pyserial (a blocked + `readline()` racing a `close()` from another thread raises `TypeError: byref() argument must be a ctypes instance, not 'NoneType'`). + Both `TeensyLatency.close_serial()` and `Teensy.close()` avoid this by using a read timeout and joining the reader thread before + closing the port. +- **`DLCClient`/socket disconnects are expected and handled on the processor side** — `MyProcessor_socket.process()` catches send + failures and just resets `self.conn`; it re-`accept()`s a fresh client on the next frame. `ActiveSensingTask.stop()` closes its + `dlcClient` (and joins its reader thread) so the socket/thread don't linger past the task's lifetime; this is safe precisely because + the processor side already tolerates a client disconnecting at any time. +- **Task-side DLC socket teardown is race-safe around startup** — `DLCClient.read_on_thread()` now owns the connection lifecycle and + closes the local connection in a `finally` block. This prevents leaking a socket/FD if `close()` is called while the background + thread is still establishing the connection. +- **Neither side's save is atomic** — both write pickle/HDF5/npy files directly to their final path. A crash or disk-full condition + mid-write can leave a truncated file at that path, including overwriting a previously-good one if re-saving to the same filename. +- **`save_legacy_timestamp_npy()`** (DLC processor side) depends on timestamp JSON files written by DLCLiveGUI's video recorder, a + separate component. If ever called before that recorder has finished flushing, it degrades gracefully — logs a warning and returns + `0` — rather than raising, so it's safe to call speculatively, just possibly a no-op in that case. + +## Where to look for what + +| Concern | File | +|---|---| +| GUI shell, session start/stop/save wiring | `teensyexp/teensy_experiment.py` | +| Generic task lifecycle (`loop`/`stop`/`get_data` contract) | `teensyexp/tasks_abc/task.py` | +| Unity-specific task base (epoch/trial counting, env step) | `teensyexp/tasks_abc/unity_task.py` | +| The concrete task used in practice | `mouse_task/task_active_sensing.py` | +| Task-variant config (per-task YAML overrides) | `mouse_task/configs/` (see `configs/README.md`) | +| Rig Teensy (reward/lick I/O) | `teensyexp/teensy.py` | +| Photodiode Teensy (latency capture) | `mouse_task/latency_tests/Teensy_latency/TeensyLatency*.py` | +| Position-data socket client (task side) | `teensyexp/tasks_abc/dlc_deque_socket.py` | +| Position-data socket server + saving (processor side) | `mouse_task/dlc_utils/dlc_processor_socket*.py` | diff --git a/mouse_task/dlc_utils/__init__.py b/mouse_task/dlc_utils/__init__.py index f707829a5..0f3f1d11d 100644 --- a/mouse_task/dlc_utils/__init__.py +++ b/mouse_task/dlc_utils/__init__.py @@ -1,7 +1,35 @@ -from .dlcProcessor_dlconly import dlc_only -from .dlc_processor_socket import MyProcessor_socket -from .dlc_processor_socket_pd import dlc_inference_w_pd -from .dlc_processor_socket_pd_sync import dlc_inference_w_pd_sync -from .simple_processor import TeensyLaser -from .processor_with_signal import ProcessorWithSignal \ No newline at end of file +"""Optional exports for DLC processor plugins. + +This package may be imported in environments that do not install +`dlclivegui` (for example, DataJoint-only runtime images). In that case, +skip exporting dlclivegui-backed processors so unrelated imports continue +to work. +""" + +from __future__ import annotations + +import importlib + +__all__: list[str] = [] + + +def _export_if_available(module_name: str, symbol_name: str) -> None: + try: + module = importlib.import_module(f".{module_name}", __name__) + except ModuleNotFoundError as exc: + missing = (exc.name or "").split(".", 1)[0] + if missing == "dlclivegui": + return + raise + + globals()[symbol_name] = getattr(module, symbol_name) + __all__.append(symbol_name) + + +_export_if_available("dlcProcessor_dlconly", "dlc_only") +_export_if_available("dlc_processor_socket", "MyProcessor_socket") +_export_if_available("dlc_processor_socket_pd", "dlc_inference_w_pd") +_export_if_available("dlc_processor_socket_pd_sync", "dlc_inference_w_pd_sync") +_export_if_available("simple_processor", "TeensyLaser") +_export_if_available("processor_with_signal", "ProcessorWithSignal") \ No newline at end of file diff --git a/mouse_task/dlc_utils/dlcProcessor_dlconly.py b/mouse_task/dlc_utils/dlcProcessor_dlconly.py index 012b3d3c1..c8b3bc335 100644 --- a/mouse_task/dlc_utils/dlcProcessor_dlconly.py +++ b/mouse_task/dlc_utils/dlcProcessor_dlconly.py @@ -1,23 +1,44 @@ import numpy as np from dlclive.processor.processor import Processor + +from dlclivegui.processors import PROCESSOR_REGISTRY, register_processor from math import sqrt, acos, atan2, copysign, degrees import pickle +PROCESSOR_REGISTRY.pop("dlc_only", None) + + +@register_processor class dlc_only(Processor): - def __init__(self, con = 50, com=2): + PROCESSOR_NAME = "DLCOnly" + PROCESSOR_DESCRIPTION = "Runs DLC inference and computes head/body kinematics." + PROCESSOR_PARAMS = { + "con": { + "type": "int", + "default": 50, + "description": "Reserved parameter (currently unused).", + }, + "com": { + "type": "int", + "default": 2, + "description": "Reserved parameter (currently unused).", + }, + } + + def __init__(self, con=50, com=2): super().__init__() self.x = [] def process(self, pose, **kwargs): xy = pose[:, :2] conf = pose[:, 2] - head_xy = xy [[0, 1, 2, 3, 4, 5, 6, 26],:] - head_conf = conf [[0, 1, 2, 3, 4, 5, 6, 26]] + head_xy = xy[[0, 1, 2, 3, 4, 5, 6, 26], :] + head_conf = conf[[0, 1, 2, 3, 4, 5, 6, 26]] center = np.average(head_xy, axis=0, weights=head_conf) body_axis = xy[7] - xy[13] # tail_base -> neck - body_axis /= sqrt(np.sum(body_axis ** 2)) + body_axis /= sqrt(np.sum(body_axis**2)) head_axis = xy[0] - xy[7] # neck -> nose - head_axis /= sqrt(np.sum(head_axis ** 2)) + head_axis /= sqrt(np.sum(head_axis**2)) cross = body_axis[0] * head_axis[1] - head_axis[0] * body_axis[1] sign = copysign(1, cross) # Positive when looking left try: @@ -30,22 +51,28 @@ def process(self, pose, **kwargs): vals = *center, heading % (360), head_angle self.x.append(center) return pose - + def save(self, filename): ### save stim on and stim off times - + filename += ".npy" try: - np.savez( - filename, out_time=self.x) + np.savez(filename, out_time=self.x) save_code = True except Exception: print("not saved") save_code = False return save_code - - - - \ No newline at end of file + + +def get_available_processors(): + return { + "dlc_only": { + "class": dlc_only, + "name": getattr(dlc_only, "PROCESSOR_NAME", "dlc_only"), + "description": getattr(dlc_only, "PROCESSOR_DESCRIPTION", ""), + "params": getattr(dlc_only, "PROCESSOR_PARAMS", {}), + } + } diff --git a/mouse_task/dlc_utils/dlc_processor_socket.py b/mouse_task/dlc_utils/dlc_processor_socket.py index c0fca5397..8c4df4817 100644 --- a/mouse_task/dlc_utils/dlc_processor_socket.py +++ b/mouse_task/dlc_utils/dlc_processor_socket.py @@ -1,27 +1,104 @@ import pickle import time +import importlib.util +import sys import warnings from collections import deque from math import acos, atan2, copysign, degrees, sqrt from multiprocessing.connection import Listener +from pathlib import Path from typing import Any, Dict, Optional import numpy as np from numpy.typing import NDArray -from dlc_utils.processor_with_signal import ProcessorWithSignal +from dlclivegui.processors import PROCESSOR_REGISTRY, register_processor # type: ignore[import-not-found] +try: + from dlc_utils.processor_with_signal import ProcessorWithSignal +except ModuleNotFoundError: + _local_path = Path(__file__).with_name("processor_with_signal.py") + _local_name = "dlclivegui_plugins.local_processor_with_signal" + _spec = importlib.util.spec_from_file_location(_local_name, _local_path) + if _spec is None or _spec.loader is None: + raise ImportError(f"Could not import ProcessorWithSignal from {_local_path}") + _module = sys.modules.get(_local_name) + if _module is None: + _module = importlib.util.module_from_spec(_spec) + sys.modules[_local_name] = _module + _spec.loader.exec_module(_module) + ProcessorWithSignal = _module.ProcessorWithSignal +PROCESSOR_REGISTRY.pop("MyProcessor_socket", None) + + +@register_processor class MyProcessor_socket(ProcessorWithSignal): + """DLC-live processor that streams mouse kinematics to a local socket client. + + Each `process()` call converts a DLC pose into center position, heading, + and head angle, appends it to in-memory buffers (drained on `save()`), + and pushes the same values to a connected client over a + `multiprocessing.connection.Listener`. Also owns the generic + resource-cleanup contract (`stop`/`close`) used by all socket-based + processors in this module, since it's the class that owns the listener + and connection. + """ + + HEAD_CONF_THRESHOLD = 0.6 + + # Legacy initialization ensures compatibility with old DLCLiveGUI processors: + # sockets / serial / side-effect-heavy resources are created inside DLCLiveWorker. + PROCESSOR_BUILD_IN_WORKER = True + + PROCESSOR_NAME = "SocketProcessor" + PROCESSOR_DESCRIPTION = "Sends DLC-derived kinematics over a local socket." + PROCESSOR_PARAMS = { + "bind": { + "type": "tuple", + "default": ("127.0.0.1", 6000), + "description": "Server bind address as (host, port).", + }, + "authkey": { + "type": "bytes", + "default": b"secret password", + "description": "Authentication key for socket clients.", + }, + "signal_delay": { + "type": "float", + "default": 10, + "description": "Delay in seconds before TTL signal starts.", + }, + "signal_type": { + "type": "str", + "default": "pulse_geo", + "description": "Signal mode: pulse, pulse_geo, sin, or flip.", + }, + "freq": { + "type": "float", + "default": 5, + "description": "Signal frequency in Hz.", + }, + } + def __init__( - self, signal_delay: float = 10, signal_type: str = "pulse_geo", freq: float = 5 + self, + bind: tuple[str, int] = ("127.0.0.1", 6000), + authkey: bytes = b"secret password", + signal_delay: float = 10, + signal_type: str = "pulse_geo", + freq: float = 5, ) -> None: super().__init__(signal_delay=signal_delay, signal_type=signal_type, freq=freq) - self.address = ("localhost", 6000) # family is deduced to be 'AF_INET' - self.listener = Listener(self.address, authkey=b"secret password") - self.conn = self.listener.accept() - print("Connection accepted from", self.listener.last_accepted) + self.address = bind + self.authkey = authkey + self.listener = Listener(self.address, authkey=self.authkey) + self.conn = None + try: + self.listener._listener._socket.settimeout(0.0) + except Exception: + pass self.center_x = deque() self.center_y = deque() @@ -35,7 +112,63 @@ def __init__( self.curr_step = 0 # frame counter self.previous = np.array([0, 0]) + def _ensure_connection(self) -> None: + if self.conn is not None: + return + try: + self.conn = self.listener.accept() + print("Connection accepted from", self.listener.last_accepted) + except Exception: + self.conn = None + + @staticmethod + def _select_single_pose(pose: Any) -> np.ndarray: + """Return one pose with shape (K, 3). + + Accepts: + (K, 3): already a single pose. + (N, K, 3): one or more detections. + + For multiple detections, selects the detection with the highest + mean keypoint likelihood. + """ + poses = np.asarray(pose) + + if poses.ndim == 2: + if poses.shape[1] != 3: + raise ValueError(f"Expected pose shape (K, 3), got {poses.shape}") + return poses + + if poses.ndim == 3: + if poses.shape[0] == 0 or poses.shape[2] != 3: + raise ValueError(f"Expected pose shape (N, K, 3), got {poses.shape}") + + if poses.shape[0] == 1: + return poses[0] + + scores = np.nanmean(poses[..., 2], axis=1) + + if not np.isfinite(scores).any(): + warnings.warn( + "No detection has a finite confidence score; selecting detection 0" + ) + return poses[0] + + index = int(np.nanargmax(scores)) + return poses[index] + + raise ValueError(f"Expected pose shape (K, 3) or (N, K, 3), got {poses.shape}") + def process(self, pose: NDArray[np.float64], **kwargs: Any) -> NDArray[np.float64]: + """Derive kinematics from a DLC pose and stream them to the socket client. + + Computes head-weighted center position (falling back to the previous + center when head-keypoint confidence is low), body/head heading, and + head angle; appends each to the processor's buffers and sends them + over `self.conn` if a client is connected. Returns the single-animal + pose used for the computation (see `_select_single_pose`). + """ + pose = self._select_single_pose(pose) # IMPORTANT: not multi-animal friendly xy = pose[:, :2] conf = pose[:, 2] @@ -43,7 +176,7 @@ def process(self, pose: NDArray[np.float64], **kwargs: Any) -> NDArray[np.float6 head_xy = xy[[0, 1, 2, 3, 4, 5, 6, 26], :] head_conf = conf[[0, 1, 2, 3, 4, 5, 6, 26]] - if np.mean(head_conf) < 0.6: + if np.mean(head_conf) < self.HEAD_CONF_THRESHOLD: center = self.previous else: center = np.average(head_xy, axis=0, weights=head_conf) @@ -80,11 +213,22 @@ def process(self, pose: NDArray[np.float64], **kwargs: Any) -> NDArray[np.float6 self.signal.append(self.curr_signal) self.frame_time.append(kwargs.get("frame_time", self.curr_time)) - self.conn.send([time.time(), vals[0], vals[1], vals[2], vals[3], vals[4]]) + self._ensure_connection() + if self.conn is not None: + try: + self.conn.send( + [time.time(), vals[0], vals[1], vals[2], vals[3], vals[4]] + ) + except Exception: + self.conn = None self.previous = center return pose def save(self, file: Optional[str] = None) -> int: + """Pickle `save_latency_data()` to `file`. + + Returns 1 on success, -1 on error, 0 if no file given. + """ save_code = 0 if file: try: @@ -102,6 +246,7 @@ def save(self, file: Optional[str] = None) -> int: return save_code def save_latency_data(self) -> Dict[str, Any]: + """Collect buffered kinematics/timing arrays for saving. Subclasses extend this dict.""" save_dict = dict() save_dict["start_time"] = np.array(self.start_time) save_dict["frame_time"] = np.array(self.frame_time) @@ -114,3 +259,62 @@ def save_latency_data(self) -> Dict[str, Any]: save_dict["head_angle"] = np.array(self.head_angle) return save_dict + + # ------------------------------------------------------------------ + # Cleanup + # ------------------------------------------------------------------ + + def stop(self, save: bool = False, file: Optional[str] = None) -> None: + """Cleanly stop processor resources. + + Subclasses with extra resources to release (e.g. a serial device) + should override `_close_extra_resources` rather than `stop` itself, + so they don't need to re-implement the save/socket/listener sequence. + """ + if save: + try: + self.save(file) + except Exception: + warnings.warn("Processor save during stop failed") + + self._close_extra_resources() + self._close_socket_connection() + self._close_listener() + + def close(self) -> None: + """Alias for generic cleanup.""" + self.stop(save=False) + + def _close_extra_resources(self) -> None: + """Hook for subclasses to close resources beyond the socket/listener. No-op by default.""" + + def _close_socket_connection(self) -> None: + try: + conn = getattr(self, "conn", None) + if conn is not None: + conn.close() + except Exception: + warnings.warn("Failed to close processor socket connection") + finally: + self.conn = None + + def _close_listener(self) -> None: + try: + listener = getattr(self, "listener", None) + if listener is not None: + listener.close() + except Exception: + warnings.warn("Failed to close processor listener") + finally: + self.listener = None + + +def get_available_processors() -> Dict[str, Dict[str, Any]]: + return { + "MyProcessor_socket": { + "class": MyProcessor_socket, + "name": getattr(MyProcessor_socket, "PROCESSOR_NAME", "MyProcessor_socket"), + "description": getattr(MyProcessor_socket, "PROCESSOR_DESCRIPTION", ""), + "params": getattr(MyProcessor_socket, "PROCESSOR_PARAMS", {}), + } + } diff --git a/mouse_task/dlc_utils/dlc_processor_socket_pd.py b/mouse_task/dlc_utils/dlc_processor_socket_pd.py index ae2869eed..9203f15c1 100644 --- a/mouse_task/dlc_utils/dlc_processor_socket_pd.py +++ b/mouse_task/dlc_utils/dlc_processor_socket_pd.py @@ -1,13 +1,94 @@ import numpy as np +import importlib.util +import sys +import warnings +from pathlib import Path from typing import Any, Dict -from latency_tests.Teensy_latency.TeensyLatency import TeensyLatency -from dlc_utils.dlc_processor_socket import MyProcessor_socket +from dlclivegui.processors import PROCESSOR_REGISTRY, register_processor +try: + from latency_tests.Teensy_latency.TeensyLatency import TeensyLatency +except ModuleNotFoundError: + TeensyLatency = None +try: + from dlc_utils.dlc_processor_socket import MyProcessor_socket +except ModuleNotFoundError: + _local_path = Path(__file__).with_name("dlc_processor_socket.py") + _local_name = "dlclivegui_plugins.local_dlc_processor_socket" + _spec = importlib.util.spec_from_file_location(_local_name, _local_path) + if _spec is None or _spec.loader is None: + raise ImportError(f"Could not import MyProcessor_socket from {_local_path}") + _module = sys.modules.get(_local_name) + if _module is None: + _module = importlib.util.module_from_spec(_spec) + sys.modules[_local_name] = _module + _spec.loader.exec_module(_module) + MyProcessor_socket = _module.MyProcessor_socket + +PROCESSOR_REGISTRY.pop("dlc_inference_w_pd", None) + + +@register_processor class dlc_inference_w_pd(MyProcessor_socket): + """`MyProcessor_socket` plus optional Teensy-driven photodiode (PD) capture. + + When `use_teensy=1`, opens a `TeensyLatency` serial connection at + construction time and records its TTL/photodiode readings alongside the + socket kinematics from the parent class. Set `use_teensy=0` to run the + socket streaming behavior alone, e.g. when no rig is attached. + """ + + HEAD_CONF_THRESHOLD = 0.6 + + # Legacy initialization ensures compatibility with old DLCLiveGUI processors: + # sockets / serial / side-effect-heavy resources are created inside DLCLiveWorker. + PROCESSOR_BUILD_IN_WORKER = True + + PROCESSOR_NAME = "SocketProcessorWithPD" + PROCESSOR_DESCRIPTION = "Socket processor with optional Teensy photodiode capture." + PROCESSOR_PARAMS = { + "com": { + "type": "str", + "default": "COM3", + "description": "Serial port used for Teensy.", + }, + "baudrate": { + "type": "int", + "default": 9600, + "description": "Teensy serial baudrate.", + }, + "signal_delay": { + "type": "float", + "default": 10, + "description": "Delay in seconds before TTL signal starts.", + }, + "signal_type": { + "type": "str", + "default": "pulse_geo", + "description": "Signal mode: pulse, pulse_geo, sin, or flip.", + }, + "freq": { + "type": "float", + "default": 5, + "description": "Signal frequency in Hz.", + }, + "use_teensy": { + "type": "bool", + "default": True, + "description": "Enable Teensy photodiode acquisition.", + }, + } + def _create_teensy(self, com, baudrate): + """Construct the Teensy handle. Overridden by subclasses to swap in other Teensy classes.""" + if TeensyLatency is None: + raise ImportError( + "TeensyLatency dependency is unavailable. Ensure mouse_task is on PYTHONPATH " + "and Teensy latency modules are installed." + ) return TeensyLatency(com, baudrate=baudrate) def __init__( @@ -45,3 +126,39 @@ def save_latency_data(self) -> Dict[str, Any]: save_dict["photodiode_time"] = np.array(self.teensy.input_data_time) return save_dict + + # ------------------------------------------------------------------ + # Cleanup + # ------------------------------------------------------------------ + + def _close_extra_resources(self) -> None: + """Close the Teensy connection as part of the base class's `stop()` sequence.""" + self._close_teensy() + super()._close_extra_resources() + + def _close_teensy(self) -> None: + try: + teensy = getattr(self, "teensy", None) + if teensy is not None: + close_serial = getattr(teensy, "close_serial", None) + if callable(close_serial): + close_serial() + else: + close = getattr(teensy, "close", None) + if callable(close): + close() + except Exception as e: + warnings.warn(f"Failed to close Teensy cleanly: {e}") + finally: + self.teensy = None + + +def get_available_processors() -> Dict[str, Dict[str, Any]]: + return { + "dlc_inference_w_pd": { + "class": dlc_inference_w_pd, + "name": getattr(dlc_inference_w_pd, "PROCESSOR_NAME", "dlc_inference_w_pd"), + "description": getattr(dlc_inference_w_pd, "PROCESSOR_DESCRIPTION", ""), + "params": getattr(dlc_inference_w_pd, "PROCESSOR_PARAMS", {}), + } + } diff --git a/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py b/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py index 152b1dcba..9f8ce912b 100644 --- a/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py +++ b/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py @@ -1,42 +1,846 @@ """Sync photodiode processor that reuses the shared DLC/socket behavior.""" -from typing import Any, Dict +from __future__ import annotations + +from datetime import datetime +import importlib.util +import json +import logging +import pickle +import re +import shutil +import sys +import time +import warnings +from pathlib import Path +from typing import Any, Dict, Optional import numpy as np +import pandas as pd +from numpy.typing import NDArray + +from dlclivegui.processors import PROCESSOR_REGISTRY, register_processor # type: ignore[import-not-found] + + +try: + from dlc_utils.dlc_processor_socket_pd import dlc_inference_w_pd +except ModuleNotFoundError: + _local_path = Path(__file__).with_name("dlc_processor_socket_pd.py") + _local_name = "dlclivegui_plugins.local_dlc_processor_socket_pd" + _spec = importlib.util.spec_from_file_location(_local_name, _local_path) + if _spec is None or _spec.loader is None: + raise ImportError(f"Could not import dlc_inference_w_pd from {_local_path}") + + _module = sys.modules.get(_local_name) + if _module is None: + _module = importlib.util.module_from_spec(_spec) + sys.modules[_local_name] = _module + _spec.loader.exec_module(_module) + + dlc_inference_w_pd = _module.dlc_inference_w_pd + + +import_issues = None +try: + from latency_tests.Teensy_latency.TeensyLatencySync import TeensyLatencySync +except ModuleNotFoundError as e: + TeensyLatencySync = None + import_issues = e + + +PROCESSOR_REGISTRY.pop("dlc_inference_w_pd_sync", None) +logger = logging.getLogger(__name__) + -from dlc_utils.dlc_processor_socket_pd import dlc_inference_w_pd -from latency_tests.Teensy_latency.TeensyLatencySync import TeensyLatencySync +def build_session_stem( + mouse: str, + date: str, + attempt: str, + namespace: str | None = None, +) -> str: + """Build __, optionally with a namespace prefix.""" + stem = f"{mouse}_{date}_{attempt}" + return f"{namespace}_{stem}" if namespace else stem +@register_processor class dlc_inference_w_pd_sync(dlc_inference_w_pd): + """`dlc_inference_w_pd` with Teensy sync timing and legacy DLCLiveGUI file outputs. + + Swaps in `TeensyLatencySync` (adds TTL-read timestamps to the photodiode + capture) and, via the `on_recording_started`/`on_recording_stopped` + hooks, reproduces the old DLCLiveGUI on-disk layout the DataJoint + pipeline expects: + - `_PROC` pickle, + - `_DLC.hdf5` pose file, + - `_TS.npy` timestamp files, + - and DB-compatible copies of video/proc/DLC outputs named `vr4mice___*`. + Poses are buffered per-frame in `process()` while recording so `save_legacy_dlc_h5()` can + write them out in one shot at the end. + """ + + HEAD_CONF_THRESHOLD = 0.6 + + # Legacy initialization ensures compatibility with old DLCLiveGUI processors: + # sockets / serial / side-effect-heavy resources are created inside DLCLiveWorker. + PROCESSOR_BUILD_IN_WORKER = True + + PROCESSOR_NAME = "SocketProcessorWithPDSync" + PROCESSOR_DESCRIPTION = "Photodiode processor with Teensy sync timing capture." + + PROCESSOR_PARAMS = { + "com": { + "type": "str", + "default": "COM3", + "description": "Serial port used for Teensy.", + }, + "baudrate": { + "type": "int", + "default": 9600, + "description": "Teensy serial baudrate.", + }, + "signal_delay": { + "type": "float", + "default": 10, + "description": "Delay in seconds before TTL signal starts.", + }, + "signal_type": { + "type": "str", + "default": "pulse_geo", + "description": "Signal mode: pulse, pulse_geo, sin, or flip.", + }, + "freq": { + "type": "float", + "default": 5, + "description": "Signal frequency in Hz.", + }, + "use_teensy": { + "type": "bool", + "default": True, + "description": "Enable Teensy photodiode acquisition.", + }, + } + def __init__( self, - com="COM3", - baudrate=9600, - signal_delay=10, - signal_type="pulse_geo", - freq=5, - use_teensy=1, - ): - super().__init__( - com=com, - baudrate=baudrate, - signal_delay=signal_delay, - signal_type=signal_type, - freq=freq, - use_teensy=use_teensy, - ) + com: str = "COM3", + baudrate: int = 9600, + signal_delay: float = 10, + signal_type: str = "pulse_geo", + freq: float = 5, + use_teensy: int | bool = 1, + ) -> None: + """Initialize legacy-output state before the parent opens the Teensy/socket connections.""" + self.recording_context: dict[str, Any] = {} + + self.save_path: Optional[Path] = None + self.dlc_h5_path: Optional[Path] = None + self.legacy_timestamp_path: Optional[Path] = None + + self._legacy_recording_active = False + self._legacy_poses: list[np.ndarray] = [] + self._legacy_pose_times: list[float] = [] + self._legacy_frame_times: list[float] = [] + + self.dlc_cfg = None + + try: + super().__init__( + com=com, + baudrate=baudrate, + signal_delay=signal_delay, + signal_type=signal_type, + freq=freq, + use_teensy=use_teensy, + ) + + try: + logger.info( + "Listener status: %s with authkey: %r", + self.listener._listener._socket.getsockname(), + self.authkey, + ) + except Exception: + logger.info( + "Listener initialized with authkey: %r", + getattr(self, "authkey", None), + ) + + except Exception as e: + self.stop(save=False) + raise RuntimeError( + f"Failed to initialize dlc_inference_w_pd_sync: {e}." + ) from e + + # ------------------------------------------------------------------ + # DLCLive processor API + # ------------------------------------------------------------------ + def process(self, pose: NDArray[np.float64], **kwargs: Any) -> NDArray[np.float64]: + """Run the parent processor and buffer pose data for legacy DLC HDF5 saving.""" + processed_pose = super().process(pose, **kwargs) + + # Parent should return pose, but guard just in case. + pose_to_buffer = processed_pose if processed_pose is not None else pose + + if self._legacy_recording_active: + try: + # Use float32 to reduce RAM pressure during long recordings. + self._legacy_poses.append( + np.asarray(pose_to_buffer, dtype=np.float32).copy() + ) + self._legacy_pose_times.append(time.time()) + self._legacy_frame_times.append( + float(kwargs.get("frame_time", time.time())) + ) + except Exception: + logger.exception("Failed to buffer pose for legacy DLC save") + + return pose_to_buffer def _create_teensy(self, com, baudrate): + """Use `TeensyLatencySync` instead of the parent's `TeensyLatency` to also capture TTL-read timing.""" + if TeensyLatencySync is None: + raise ImportError( + "TeensyLatencySync dependency is unavailable. Ensure mouse_task is on PYTHONPATH " + "and Teensy latency modules are installed." + ) from import_issues + return TeensyLatencySync(com, baudrate=baudrate) + def set_dlc_cfg(self, dlc_cfg): + """Store the DLC model config (used to label bodyparts when saving the legacy DLC h5).""" + self.dlc_cfg = dlc_cfg + + # ------------------------------------------------------------------ + # Recording lifecycle hooks from DLCLiveGUI + # ------------------------------------------------------------------ + + def on_recording_started(self, context: dict) -> None: + """Receive recording context from DLCLiveGUI.""" + self.recording_context = dict(context or {}) + + base_path = self._context_processor_base_path() + if base_path is None: + self.save_path = None + self.dlc_h5_path = None + self.legacy_timestamp_path = None + logger.warning("Processor recording started without processor_base_path") + return + + # Primary/native processor outputs. + self.save_path = base_path.parent / f"{base_path.name}_PROC" + self.dlc_h5_path = base_path.parent / f"{base_path.name}_DLC.hdf5" + + # Single/base timestamp target. Per-video timestamp targets are derived later. + self.legacy_timestamp_path = base_path.parent / f"{base_path.name}_TS.npy" + + self._legacy_recording_active = True + self._legacy_poses.clear() + self._legacy_pose_times.clear() + self._legacy_frame_times.clear() + + logger.info("Processor save path set to %s", self.save_path) + logger.info("Processor DLC h5 path set to %s", self.dlc_h5_path) + logger.info("Processor timestamp path set to %s", self.legacy_timestamp_path) + + def _save_legacy_outputs(self): + """Save all legacy outputs, attempting each independently so one failure + does not prevent the remaining outputs from being written.""" + for name, method in [ + ("PROC save", self.save), + ("DLC h5 save", self.save_legacy_dlc_h5), + ("timestamp npy save", self.save_legacy_timestamp_npy), + ("video copy", self.copy_legacy_video_files), + ("output alignment", self.copy_processor_outputs_to_primary_legacy_base), + ]: + try: + result = method() + logger.info("Processor legacy %s result: %r", name, result) + except Exception: + logger.exception("Processor %s failed during legacy save", name) + + self._clear_legacy_pose_buffers() + + def on_recording_stopped(self, context: dict) -> None: + """Save all custom legacy outputs after GUI recording stops.""" + previous_context = dict(getattr(self, "recording_context", {}) or {}) + previous_context.update(context or {}) + self.recording_context = previous_context + + self._legacy_recording_active = False + self._save_legacy_outputs() + + def stop(self, save: bool = False, file=None): + """Save all buffered data before tearing down resources. + + This guards against data loss when ``stop()`` is called + before ``on_recording_stopped`` (e.g. during DLCLiveWorker + shutdown, socket disconnect, or task auto-stop). + """ + if self._legacy_recording_active and self.save_path is not None: + self._save_legacy_outputs() + self._legacy_recording_active = False + + super().stop(save=False, file=file) + + # ------------------------------------------------------------------ + # Primary PROC save + # ------------------------------------------------------------------ + def save_latency_data(self) -> Dict[str, Any]: + """Extend parent's save_latency_data with Teensy TTL-read timestamps.""" save_dict = super().save_latency_data() - if self.use_teensy == 1: + if getattr(self, "use_teensy", 0) == 1: save_dict["ttl_read"] = np.array(getattr(self.teensy, "input_data_ttl", [])) save_dict["teensy_time"] = np.array( getattr(self.teensy, "input_data_teensy_time", []) ) return save_dict + + def save(self, file: str | Path | None = None) -> int: + """Save processor PROC-style data. + + If `file` is not provided, uses `self.save_path`. + """ + target = Path(file) if file is not None else getattr(self, "save_path", None) + + if target is None: + warnings.warn("Processor save skipped: no file or save_path was provided.") + return 0 + + try: + target = Path(target) + target.parent.mkdir(parents=True, exist_ok=True) + save_dict = self.save_latency_data() + + with target.open("wb") as f: + pickle.dump(save_dict, f) + + logger.info("Processor data saved to: %s", target) + return 1 + + except Exception as e: + warnings.warn(f"Proc file was not saved, an exception occurred: {e}") + logger.exception("Processor PROC save failed") + return -1 + + # ------------------------------------------------------------------ + # Legacy DLC HDF5 saving + # ------------------------------------------------------------------ + + def save_legacy_dlc_h5(self) -> int: + """Write old-style _DLC.hdf5 from buffered poses.""" + target = getattr(self, "dlc_h5_path", None) + if target is None: + logger.warning("Skipping DLC h5 save: no dlc_h5_path") + return 0 + + poses = np.asarray(getattr(self, "_legacy_poses", [])) + if poses.size == 0: + logger.warning("Skipping DLC h5 save: no buffered poses") + return 0 + + try: + target = Path(target) + target.parent.mkdir(parents=True, exist_ok=True) + + if poses.ndim == 2: + poses = poses[None, :, :] + + if poses.ndim != 3 or poses.shape[-1] != 3: + logger.warning( + "Skipping DLC h5 save: unexpected pose shape %s", poses.shape + ) + return 0 + + flat = poses.reshape((poses.shape[0], poses.shape[1] * poses.shape[2])) + + bodyparts = self._get_bodyparts_for_pose_width(flat.shape[1]) + if bodyparts: + pdindex = pd.MultiIndex.from_product( + [bodyparts, ["x", "y", "likelihood"]], + names=["bodyparts", "coords"], + ) + pose_df = pd.DataFrame(flat, columns=pdindex) + else: + logger.warning( + "Bodyparts information not found or mismatched; saving DLC h5 without labels." + ) + pose_df = pd.DataFrame(flat) + + pose_df["frame_time"] = list(self._legacy_frame_times) + pose_df["pose_time"] = list(self._legacy_pose_times) + + pose_df.to_hdf(target, key="df_with_missing", mode="w") + + logger.info("Legacy DLC h5 saved to: %s", target) + return 1 + + except Exception: + logger.exception("Failed to save legacy DLC h5") + return -1 + + def _get_bodyparts_for_pose_width(self, flat_width: int) -> list[str] | None: + """Return bodypart names from `self.dlc_cfg` if their count matches the flattened pose width.""" + dlc_cfg = getattr(self, "dlc_cfg", None) + bodyparts = None + + if isinstance(dlc_cfg, dict): + bodyparts = dlc_cfg.get("all_joints_names") or dlc_cfg.get( + "metadata", {} + ).get("bodyparts") + + if bodyparts and len(bodyparts) * 3 == flat_width: + return list(bodyparts) + + return None + + # ------------------------------------------------------------------ + # Timestamp JSON -> legacy NPY + # ------------------------------------------------------------------ + + def save_legacy_timestamp_npy(self) -> int: + # Reads timestamp JSON files written by DLCLiveGUI's video recorder, which is + # a separate component. If this is ever called from a teardown path that can + # run before the video recorder has finished flushing (e.g. before/without + # DLCLiveGUI's own "Stop"/"Save Video"), the JSON files may not exist yet or + # may be incomplete -- this degrades gracefully to a logged warning and + # `return 0` rather than raising, so that's safe, just possibly a no-op. + json_paths = self._find_timestamp_json_files() + + if not json_paths: + logger.warning( + "Skipping legacy timestamp npy save: no timestamp JSON files found" + ) + return 0 + + saved = 0 + total = len(json_paths) + compat_base = self._db_compat_base() + + for index, json_path in enumerate(json_paths): + try: + timestamps = self._extract_timestamps_from_json(json_path) + if timestamps.size == 0: + logger.warning("No timestamps extracted from %s", json_path) + continue + + for out_path in self._timestamp_output_paths( + compat_base, + index=index, + total=total, + ): + self._save_npy(out_path, timestamps) + logger.info( + "DB-compatible timestamp npy saved to %s with %d timestamps", + out_path, + len(timestamps), + ) + saved += 1 + + except Exception: + logger.exception( + "Failed to convert timestamp JSON to npy: %s", json_path + ) + + return 1 if saved else 0 + + def _extract_timestamps_from_json(self, json_path: Path) -> np.ndarray: + """Extract timestamps from the new VideoRecorder JSON format. + + Old GUI saved `np.save(..., write_frame_ts)`, i.e. a 1D numeric array. + This returns the same shape/type. + """ + json_path = Path(json_path) + + with json_path.open("r", encoding="utf-8") as f: + data = json.load(f) + + if isinstance(data, dict) and isinstance(data.get("frame_timestamps"), list): + timestamps = [ + float(rec["software_timestamp"]) + for rec in data["frame_timestamps"] + if isinstance(rec, dict) and "software_timestamp" in rec + ] + return np.asarray(timestamps, dtype=float) + + if isinstance(data, dict): + for key in ("timestamps", "frame_times", "times"): + values = data.get(key) + if isinstance(values, list): + return np.asarray(values, dtype=float) + + if isinstance(data, list): + values = [] + for item in data: + if isinstance(item, (int, float)): + values.append(float(item)) + elif isinstance(item, dict): + value = self._first_present( + item, ("software_timestamp", "timestamp", "frame_time", "time") + ) + if value is not None: + values.append(float(value)) + return np.asarray(values, dtype=float) + + return np.asarray([], dtype=float) + + def _find_timestamp_json_files(self) -> list[Path]: + """Resolve timestamp JSON files from recording context.""" + paths = self._paths_from_context_value( + self.recording_context.get("timestamp_json_files") + or self.recording_context.get("timestamp_files") + ) + + paths = [p for p in paths if p.exists()] + if paths: + return sorted(paths) + + run_dir = self.recording_context.get("run_dir") + if run_dir is not None: + return sorted(Path(run_dir).glob("*_timestamps.json")) + + return [] + + def _timestamp_output_paths( + self, + compat_base: Path, + *, + index: int, + total: int, + ) -> list[Path]: + camera_token = "CAMERA" if total == 1 else f"CAMERA{index + 1}" + + return self._unique_paths( + [ + compat_base.parent / f"TS_{compat_base.name}_{camera_token}.npy", + # compat_base.parent / f"TIMESTAMP_{compat_base.name}_{camera_token}.npy", + ] + ) + + # ------------------------------------------------------------------ + # Legacy video/sidecar compatibility copies + # ------------------------------------------------------------------ + def copy_legacy_video_files(self) -> int: + copied = 0 + video_files = self._find_video_files() + compat_base = self._db_compat_base() + total = len(video_files) + + for index, video_path in enumerate(video_files): + try: + video_token = "VIDEO" if total == 1 else f"VIDEO{index + 1}" + out_path = ( + compat_base.parent + / f"{compat_base.name}_{video_token}{video_path.suffix}" + ) + + if self._copy_file_if_needed(video_path, out_path): + copied += 1 + + except Exception: + logger.exception( + "Failed to copy DB-compatible video file %s", video_path + ) + + return 1 if copied else 0 + + def copy_processor_outputs_to_primary_legacy_base(self) -> int: + compat_base = self._db_compat_base() + copied = 0 + + src_proc = getattr(self, "save_path", None) + if src_proc is not None: + dst_proc = compat_base.parent / f"{compat_base.name}_PROC" + if self._copy_file_if_needed(Path(src_proc), dst_proc): + copied += 1 + + src_h5 = getattr(self, "dlc_h5_path", None) + if src_h5 is not None: + dst_h5 = compat_base.parent / f"{compat_base.name}_DLC.hdf5" + if self._copy_file_if_needed(Path(src_h5), dst_h5): + copied += 1 + + return 1 if copied else 0 + + # ------------------------------------------------------------------ + # Legacy base / path helpers + # ------------------------------------------------------------------ + def _video_prefix(self) -> str: + videos = self._find_video_files() + + if videos: + return videos[0].stem.split("_", 1)[0] + + filename_stem = self.recording_context.get("filename_stem") + if filename_stem: + return str(filename_stem).split("_", 1)[0] + + return "recording" + + def _db_compat_base(self) -> Path: + """Return DB-GUI-compatible base path. + + New Live-GUI layout is usually: + + MouseA/run_/ + + This returns: + + /MouseA_YYYY-MM-DD_1 + + so files parse correctly as: + mouse_name = MouseA + date = YYYY-MM-DD + attempt = 1 + """ + context = getattr(self, "recording_context", {}) or {} + + run_dir = context.get("run_dir") + run_dir = Path(run_dir) if run_dir is not None else self._fallback_output_dir() + + prefix = self._video_prefix() + # mouse = self._mouse_from_context_or_run_dir(run_dir) + date = self._date_from_context_or_run_dir(run_dir) + attempt = self._attempt_from_context(default="1") + + return run_dir / build_session_stem( + prefix, + date, + attempt, + namespace="vr4mice", + ) + + def _fallback_output_dir(self) -> Path: + base_path = self._context_processor_base_path() + if base_path is not None: + return base_path.parent + + save_path = getattr(self, "save_path", None) + if save_path is not None: + return Path(save_path).parent + + return Path.cwd() + + def _mouse_from_context_or_run_dir(self, run_dir: Path) -> str: + context = getattr(self, "recording_context", {}) or {} + + for key in ("mouse", "mouse_name", "subject", "session_name"): + value = context.get(key) + if value: + return self._sanitize(str(value)) + + # New Live-GUI layout: MouseA/run_/ + parent_name = getattr(run_dir.parent, "name", "") + if parent_name: + return self._sanitize(parent_name) + + return "Mouse" + + def _date_from_context_or_run_dir(self, run_dir: Path) -> str: + context = getattr(self, "recording_context", {}) or {} + + for key in ("date", "recording_date", "session_date"): + value = context.get(key) + if value: + parsed = self._normalize_date(str(value)) + if parsed: + return parsed + + parsed = self._date_from_run_dir_name(run_dir.name) + if parsed: + return parsed + + try: + return datetime.fromtimestamp(run_dir.stat().st_mtime).strftime("%Y-%m-%d") + except Exception: + return datetime.now().strftime("%Y-%m-%d") + + def _attempt_from_context(self, default: str = "1") -> str: + context = getattr(self, "recording_context", {}) or {} + + for key in ("attempt", "trial", "run_index"): + value = context.get(key) + if value not in (None, ""): + return self._sanitize(str(value)) + + filename_stem = context.get("filename_stem") + if filename_stem: + parts = str(filename_stem).split("_") + for part in reversed(parts): + if part.isdigit(): + return self._sanitize(part) + + return default + + @staticmethod + def _sanitize(value: str) -> str: + value = str(value).strip() + value = value.replace(" ", "") + value = value.replace("_", "") + return value or "unknown" + + @staticmethod + def _normalize_date(value: str) -> str | None: + value = str(value) + + # Already YYYY-MM-DD + match = re.search(r"(20\d{2}-\d{2}-\d{2})", value) + if match: + return match.group(1) + + # YYYYMMDD + match = re.search(r"(20\d{2})(\d{2})(\d{2})", value) + if match: + return f"{match.group(1)}-{match.group(2)}-{match.group(3)}" + + return None + + def _date_from_run_dir_name(self, run_name: str) -> str | None: + return self._normalize_date(run_name) + + def _context_processor_base_path(self) -> Path | None: + base_path = self.recording_context.get("processor_base_path") + return Path(base_path) if base_path is not None else None + + def _primary_legacy_base(self) -> Path | None: + return self._db_compat_base() + + def _legacy_base_for_timestamp_json(self, json_path: Path) -> Path: + """Return legacy base path inferred from a timestamp JSON file.""" + json_path = Path(json_path) + run_dir = Path(self.recording_context.get("run_dir") or json_path.parent) + + video_name = self._video_name_from_timestamp_json(json_path) + if video_name: + return run_dir / Path(video_name).stem + + return run_dir / self._strip_timestamp_json_suffix(json_path.name) + + def _legacy_base_for_video(self, video_path: Path) -> Path: + """Return legacy base path inferred from a video path.""" + video_path = Path(video_path) + return video_path.parent / video_path.stem + + def _video_name_from_timestamp_json(self, json_path: Path) -> str | None: + try: + with Path(json_path).open("r", encoding="utf-8") as f: + data = json.load(f) + if isinstance(data, dict): + video_name = data.get("video_file") + return str(video_name) if video_name else None + except Exception: + return None + + return None + + @staticmethod + def _strip_timestamp_json_suffix(name: str) -> str: + for suffix in ( + ".avi_timestamps.json", + ".mp4_timestamps.json", + "_timestamps.json", + ): + if name.endswith(suffix): + return name[: -len(suffix)] + return Path(name).stem + + def _find_video_files(self) -> list[Path]: + paths = self._paths_from_context_value( + self.recording_context.get("video_files") + ) + paths = [p for p in paths if p.exists()] + if paths: + return sorted(paths) + + run_dir = self.recording_context.get("run_dir") + if run_dir is None: + return [] + + run_dir = Path(run_dir) + return sorted([*run_dir.glob("*.avi"), *run_dir.glob("*.mp4")]) + + @staticmethod + def _paths_from_context_value(value: Any) -> list[Path]: + if value is None: + return [] + + if isinstance(value, (str, Path)): + return [Path(value)] + + if isinstance(value, dict): + return [Path(v) for v in value.values() if isinstance(v, (str, Path))] + + if isinstance(value, (list, tuple, set)): + return [Path(v) for v in value if isinstance(v, (str, Path))] + + return [] + + @staticmethod + def _unique_paths(paths: list[Path]) -> list[Path]: + unique: list[Path] = [] + seen: set[str] = set() + for path in paths: + key = str(path) + if key not in seen: + unique.append(path) + seen.add(key) + return unique + + @staticmethod + def _copy_file_if_needed(src: Path, dst: Path) -> bool: + src = Path(src) + dst = Path(dst) + + if not src.exists(): + return False + + try: + if src.resolve() == dst.resolve(): + return False + except Exception: + if str(src) == str(dst): + return False + + dst.parent.mkdir(parents=True, exist_ok=True) + shutil.copy2(src, dst) + logger.info("Copied compatibility file %s -> %s", src, dst) + return True + + @staticmethod + def _save_npy(path: Path, values: np.ndarray) -> None: + path = Path(path) + path.parent.mkdir(parents=True, exist_ok=True) + np.save(path, values) + + @staticmethod + def _first_present(mapping: dict, keys: tuple[str, ...]) -> Any: + for key in keys: + if key in mapping: + return mapping[key] + return None + + def _clear_legacy_pose_buffers(self) -> None: + try: + self._legacy_poses.clear() + self._legacy_pose_times.clear() + self._legacy_frame_times.clear() + except Exception: + logger.warning("Failed to clear legacy pose buffers after recording stop") + + +def get_available_processors() -> Dict[str, Dict[str, Any]]: + return { + "dlc_inference_w_pd_sync": { + "class": dlc_inference_w_pd_sync, + "name": getattr( + dlc_inference_w_pd_sync, "PROCESSOR_NAME", "dlc_inference_w_pd_sync" + ), + "description": getattr( + dlc_inference_w_pd_sync, "PROCESSOR_DESCRIPTION", "" + ), + "params": getattr(dlc_inference_w_pd_sync, "PROCESSOR_PARAMS", {}), + } + } diff --git a/mouse_task/dlc_utils/simple_processor.py b/mouse_task/dlc_utils/simple_processor.py index 303c355ba..91e4f980c 100644 --- a/mouse_task/dlc_utils/simple_processor.py +++ b/mouse_task/dlc_utils/simple_processor.py @@ -1,23 +1,38 @@ from dlclive.processor.processor import Processor -import serial -import struct import pickle import time +from dlclivegui.processors import PROCESSOR_REGISTRY, register_processor # type: ignore[import-not-found] +PROCESSOR_REGISTRY.pop("TeensyLaser", None) + + +@register_processor class TeensyLaser(Processor): - def __init__( - self, com = 50, conn=2): + PROCESSOR_NAME = "TeensyLaser" + PROCESSOR_DESCRIPTION = "Simple processor that logs stimulation timestamps." + PROCESSOR_PARAMS = { + "com": { + "type": "int", + "default": 50, + "description": "Reserved COM parameter (currently unused).", + }, + "conn": { + "type": "int", + "default": 2, + "description": "Reserved connection parameter (currently unused).", + }, + } + + def __init__(self, com=50, conn=2): super().__init__() self.stim_on_time = [] - def process(self, pose, **kwargs): # define criteria to stimulate (e.g. if first point is in a corner of the video) self.stim_on_time.append(time.time()) - return pose @@ -34,4 +49,15 @@ def save(self, file=None): save_code = 1 except Exception: save_code = -1 - return save_code \ No newline at end of file + return save_code + + +def get_available_processors(): + return { + "TeensyLaser": { + "class": TeensyLaser, + "name": getattr(TeensyLaser, "PROCESSOR_NAME", "TeensyLaser"), + "description": getattr(TeensyLaser, "PROCESSOR_DESCRIPTION", ""), + "params": getattr(TeensyLaser, "PROCESSOR_PARAMS", {}), + } + } diff --git a/mouse_task/latency_tests/Teensy_latency/TeensyLatency.py b/mouse_task/latency_tests/Teensy_latency/TeensyLatency.py index 07a0993c0..9908cc3b4 100644 --- a/mouse_task/latency_tests/Teensy_latency/TeensyLatency.py +++ b/mouse_task/latency_tests/Teensy_latency/TeensyLatency.py @@ -23,7 +23,20 @@ def _handle_line(self, line: str, now: float): def read_on_thread(self): while self.reading_teensy and not self.stop_event.is_set(): - line = self.ser.readline().decode("utf-8").rstrip() + try: + raw_line = self.ser.readline() + except (serial.SerialException, OSError): + # Serial port was closed (e.g. by close_serial()) while readline() was blocked. + break + except TypeError as err: + # On some Windows/pyserial versions, close() racing readline() + # raises: "byref() argument must be a ctypes instance, not 'NoneType'". + if "byref() argument must be a ctypes instance" in str(err): + break + raise + if not raw_line: + continue # timeout with no data + line = raw_line.decode("utf-8").rstrip() now = time.time() # Current time try: self._handle_line(line, now) @@ -32,15 +45,16 @@ def read_on_thread(self): def start_read_buffer(self): """Start the reader thread for serial buffer, writer for `input_data`, save start time.""" - self.ser = serial.Serial(self.com, self.baudrate) + self.ser = serial.Serial(self.com, self.baudrate, timeout=0.5) self.start_read_time = time.time() - threading.Thread(target=self.read_on_thread, daemon=True).start() - + self._reader_thread = threading.Thread(target=self.read_on_thread, daemon=True) + self._reader_thread.start() def _stop_reading(self): """Stop reading from teensy and close serial connection.""" self.reading_teensy = False self.stop_event.set() - def close_serial(self): self._stop_reading() + if hasattr(self, '_reader_thread') and self._reader_thread.is_alive(): + self._reader_thread.join(timeout=2.0) self.ser.close() diff --git a/mouse_task/task_active_sensing.py b/mouse_task/task_active_sensing.py index 3b155f7cd..2a8584a11 100644 --- a/mouse_task/task_active_sensing.py +++ b/mouse_task/task_active_sensing.py @@ -283,8 +283,12 @@ def __init__( self.dlc_time_step = [] self.trial_mouse_report_delay = [] - # self.set_channel() - # self.reset_environment() + def stop(self): + """Stop the task: parent's teardown, plus closing the DLC socket client if one was opened.""" + super().stop() + dlc_client = getattr(self, "dlcClient", None) + if dlc_client is not None: + dlc_client.close() def _get_dlc_on_frame(self): """ diff --git a/teensyexp/tasks_abc/dlc_deque_socket.py b/teensyexp/tasks_abc/dlc_deque_socket.py index f766df9cb..875d2ad97 100644 --- a/teensyexp/tasks_abc/dlc_deque_socket.py +++ b/teensyexp/tasks_abc/dlc_deque_socket.py @@ -1,13 +1,23 @@ +"""Socket-based DLC client using deque-backed buffers. + +Starts a background reader thread that receives DLC frames from DLCLiveGUI +and keeps the newest values in deques for low-overhead access. +""" + import numpy as np import time import threading -from multiprocessing.connection import Client +import socket +from multiprocessing.connection import Connection, answer_challenge, deliver_challenge from collections import deque class DLCClient(object): def __init__(self, address=("localhost", 6000)): self.address = address + self.authkey = b"secret password" + self.connect_timeout = 0.2 + self._stop_event = threading.Event() self.reading = True self.input_data = deque() self.save_input_data = deque() @@ -15,19 +25,62 @@ def __init__(self, address=("localhost", 6000)): self.start_read_buffer() self.start_time = time.time() + def _connect_with_timeout(self): + sock = None + conn = None + + try: + if isinstance(self.address, tuple): + sock = socket.create_connection(self.address, timeout=self.connect_timeout) + else: + sock = socket.socket(socket.AF_UNIX) + sock.settimeout(self.connect_timeout) + sock.connect(self.address) + sock.settimeout(None) + conn = Connection(sock.detach()) + answer_challenge(conn, self.authkey) + deliver_challenge(conn, self.authkey) + return conn + except Exception: + if conn is not None: + conn.close() + elif sock is not None: + sock.close() + raise + def read_on_thread(self): - self.conn = Client(self.address, authkey=b"secret password") - while self.reading: - try: - this_read = self.conn.recv() - self.input_data.append(this_read) + conn = None + try: + # Keep connect attempts bounded so close() can wait synchronously. + while self.reading and not self._stop_event.is_set(): + try: + conn = self._connect_with_timeout() + break + except (TimeoutError, socket.timeout): + continue + + if conn is None: + return + + self.conn = conn + while self.reading: + try: + this_read = conn.recv() + self.input_data.append(this_read) - except EOFError: - self.reading = False - break + except (EOFError, OSError): + # EOFError: remote side closed cleanly. OSError: local conn.close() + # ran while recv() was blocked (e.g. from close()). + self.reading = False + break + finally: + if conn is not None: + conn.close() + self.conn = None def start_read_buffer(self): - threading.Thread(target=self.read_on_thread, daemon=True).start() + self._read_thread = threading.Thread(target=self.read_on_thread, daemon=True) + self._read_thread.start() def read(self): if len(self.input_data) >= 1: @@ -51,6 +104,7 @@ def read(self): def stop(self): self.reading = False + self._stop_event.set() def get_input_data(self): return np.array(list(self.input_data)) @@ -60,4 +114,9 @@ def reset(self): def close(self): self.stop() - self.conn.close() + conn = getattr(self, "conn", None) + if conn is not None: + conn.close() + read_thread = getattr(self, "_read_thread", None) + if read_thread is not None: + read_thread.join() diff --git a/teensyexp/tasks_abc/dlc_socket.py b/teensyexp/tasks_abc/dlc_socket.py index 665692602..68d9b812b 100644 --- a/teensyexp/tasks_abc/dlc_socket.py +++ b/teensyexp/tasks_abc/dlc_socket.py @@ -1,37 +1,87 @@ +"""Socket-based DLC client using a list-backed read buffer. + +Starts a background reader thread that receives DLC frames from DLCLiveGUI +and stores them for task consumption. +""" + import numpy as np import time import threading -from multiprocessing.connection import Client -import numpy as np -import time +import socket +from multiprocessing.connection import Connection, answer_challenge, deliver_challenge + class DLCClient(object): def __init__(self, address = ('localhost', 6000)): # start read buffer self.address = address + self.authkey = b'secret password' + self.connect_timeout = 0.2 + self._stop_event = threading.Event() self.reading = True self.input_data = [] self.start_read_buffer() - + + + def _connect_with_timeout(self): + sock = None + conn = None + + try: + if isinstance(self.address, tuple): + sock = socket.create_connection(self.address, timeout=self.connect_timeout) + else: + sock = socket.socket(socket.AF_UNIX) + sock.settimeout(self.connect_timeout) + sock.connect(self.address) + sock.settimeout(None) + conn = Connection(sock.detach()) + answer_challenge(conn, self.authkey) + deliver_challenge(conn, self.authkey) + return conn + except Exception: + if conn is not None: + conn.close() + elif sock is not None: + sock.close() + raise def read_on_thread(self): - # start connection to the socket - self.conn = Client(self.address, authkey=b'secret password') - # start reading and add data to a list - while self.reading == True: - try: - this_read = self.conn.recv() - self.input_data.append(list((time.time(),this_read))) - # if the connection on the DLClivegui is closed, stop the thread reading in a clean way - except EOFError: - self.reading == False - break + conn = None + try: + # Keep connect attempts bounded so close() can wait synchronously. + while self.reading and not self._stop_event.is_set(): + try: + conn = self._connect_with_timeout() + break + except (TimeoutError, socket.timeout): + continue + + if conn is None: + return + + self.conn = conn + # start reading and add data to a list + while self.reading == True: + try: + this_read = conn.recv() + self.input_data.append(list((time.time(),this_read))) + # if the connection on the DLCLiveGUI is closed, or close() ran while + # recv() was blocked, stop the thread reading in a clean way + except (EOFError, OSError): + self.reading = False + break + finally: + if conn is not None: + conn.close() + self.conn = None def start_read_buffer(self): # start reading from DLClivegui in thread self.start_read_time = time.time() self.reading = True - threading.Thread(target=self.read_on_thread, daemon=True).start() + self._read_thread = threading.Thread(target=self.read_on_thread, daemon=True) + self._read_thread.start() def read(self, index=-1, input=None): """ @@ -42,11 +92,9 @@ def read(self, index=-1, input=None): return({"time": vals[0], "vals": vals [1]}) def stop(self): - - """ - change the reading class attribute to False (switch flag) - """ + """Change the reading class attribute to False (switch flag).""" self.reading = False + self._stop_event.set() def get_input_data(self, format='array'): """ @@ -58,17 +106,21 @@ def get_input_data(self, format='array'): return np.array(self.input_data) def reset(self): - """ - method reset to empty list input_data and output_data attributes - """ + """Reset input_data to an empty list.""" self.input_data = [] def close(self): - """ - method to stop communication and update reading state attribute to False via stop() - """ + """Stop communication and update reading state attribute.""" self.stop() - self.conn.close() + conn = getattr(self, "conn", None) + if conn is not None: + conn.close() + read_thread = getattr(self, "_read_thread", None) + if read_thread is not None: + read_thread.join() + + + diff --git a/teensyexp/teensy.py b/teensyexp/teensy.py index b9be1d394..655c2bf47 100644 --- a/teensyexp/teensy.py +++ b/teensyexp/teensy.py @@ -63,17 +63,26 @@ def read_on_thread(self): buffer = None delta = 1 while self.reading: - if self.ser.inWaiting() > delta: + try: + waiting = self.ser.inWaiting() > delta + except (serial.SerialException, OSError): + # Serial port was closed (e.g. by close()) while this thread was running. + break + if waiting: + try: + new_bytes = self.ser.read() + except (serial.SerialException, OSError): + break if buffer: - buffer = buffer + self.ser.read() + buffer = buffer + new_bytes else: - buffer = self.ser.read() + buffer = new_bytes if self.end_bytes in buffer: lines = buffer.split(self.end_bytes) buffer = lines[-1] this_read = struct.unpack('h' * self.n_inputs, lines[-2]) self.input_data.append(list((time.time(),) + this_read)) - + def start_read_buffer(self): """ method that starts the reader thread (reader for serial buffer), writer for (input_data) @@ -81,7 +90,8 @@ def start_read_buffer(self): """ self.start_read_time = time.time() self.reading = True - threading.Thread(target=self.read_on_thread, daemon=True).start() + self._read_thread = threading.Thread(target=self.read_on_thread, daemon=True) + self._read_thread.start() def read(self, index=-1, input=None): """ @@ -179,4 +189,7 @@ def close(self): stop serial communication and update reading state attribute to False via stop() """ self.stop() + read_thread = getattr(self, "_read_thread", None) + if read_thread is not None: + read_thread.join(timeout=2) self.ser.close() diff --git a/teensyexp/teensy_experiment.py b/teensyexp/teensy_experiment.py index 17b9a9aae..1e28314eb 100644 --- a/teensyexp/teensy_experiment.py +++ b/teensyexp/teensy_experiment.py @@ -1,10 +1,10 @@ -""" -GUI to run teensy experiments - - system setup information taken from system_setup.json (which is written by system_setup.py) +"""Teensy experiment GUI. -GK 05/07/2019 +Loads rig and task setup from JSON configuration files and runs experiment +sessions. -Note(mary): API documentation added 11/08/2022 +Original implementation: GK (2019-05-07) +API documentation additions: mary (2022-11-08) """ import os @@ -602,6 +602,20 @@ def init_task(self): parent=self.window) self.task_on.set(1) else: + # Re-initializing replaces self.task, so any unsaved data from the + # previous task becomes unreachable. Let experimenters cancel to + # save first, or explicitly proceed and discard old task access. + if self.task is not None and not self.saved_ok: + proceed = messagebox.askokcancel( + "Unsaved Data", + "The previous task's data has not been saved.\n" + "Click Cancel to save first, or OK to initialize a new task and discard access to the previous task data.", + parent=self.window, + ) + if not proceed: + self.task_on.set(0) + return + task_object = getattr(self.task_module, self.task_name.get()) task_params = copy.deepcopy(self.task_params[self.task_name.get()]) try: @@ -609,21 +623,24 @@ def init_task(self): except Exception as err: self.task = None self.task_info = {} - self.task_label["text"] = "No Task" + self.task_label["text"] = "No Task" self.task_on.set(0) - try: - self._reset_progress_labels() - except Exception: - pass - finally: - self.info_labels = [] - self.value_labels = [] + try: + self._reset_progress_labels() + except Exception: + pass + finally: + self.info_labels = [] + self.value_labels = [] messagebox.showerror( "Task Initialization Failed", f"Could not initialize task '{self.task_name.get()}'.\n{err}", parent=self.window, ) return + # This is a fresh task with nothing saved yet, regardless of whether + # the previous task's data was ever saved. + self.saved_ok = False parent_class = [c.__name__ for c in self.task.__class__.__mro__] self.gui_task = True if 'GuiTask' in parent_class else False self.unity_task = True if 'UnityTask' in parent_class else False @@ -742,8 +759,21 @@ def _dump_data(self, data_to_save, filename): Args: data_to_save: output form task (return of self.task.get_data()) filename(str): path and name of file to save + + Note: + `self.saved_ok` is only set on success. """ - pickle.dump(data_to_save, open(filename, 'wb')) + try: + with open(filename, 'wb') as f: + pickle.dump(data_to_save, f) + except Exception as e: + messagebox.showerror( + "Save Failed", + "Failed to save data to %s:\n%s" % (filename, e), + parent=self.window, + ) + return + messagebox.showinfo("File Saved", "File saved to %s" % filename, parent=self.window) self.saved_ok = True @@ -800,16 +830,21 @@ def save_data(self): def check_close(self): """ - method used for close bottom callback + method used for close button callback (and the window's X button) checks if there is a running task and if all data saved """ if self.task_on.get(): messagebox.showerror("Task Open", "Task is currently open. Please stop task before closing.", parent=self.window) + elif not self.saved_ok: + if messagebox.askokcancel( + "Exit", + "ARE YOU SURE YOU SAVED YOUR Data?", + parent=self.window, + ): + self.gui_on = False else: - if not self.saved_ok: - if messagebox.askokcancel("Exit", "ARE YOU SURE YOU SAVED YOUR Data?"): - self.gui_on = False + self.gui_on = False def close_window(self): """ @@ -969,6 +1004,10 @@ def create_gui(self): Button(window, text="Close", command=self.check_close).grid(sticky="nsew", row=cur_row, column=1, columnspan=1) cur_row += 1 + # route the window's own [X] close button through the same "did you save?" check + # instead of letting Tkinter destroy the window unprompted + window.protocol("WM_DELETE_WINDOW", self.check_close) + # configure size of empty rows col_count, row_count = window.grid_size() for r in range(row_count): @@ -985,20 +1024,45 @@ def run_experiment(self): print_delay = .01 last_print = time.time() - while self.gui_on: - curr_time = time.time() - if self.task_on_button: - if curr_time - last_print > print_delay: - self.check_task_progress() - last_print = curr_time - elif self.task_on.get() == 1: - self.task_on.set(0) - if self.gui_task: - self.task.window.destroy() - - self.window.update() + def _warn_use_gui_close(): + try: + messagebox.showwarning( + "Use the GUI to Close", + "Ctrl+C does not safely close this program.\n" + "Please use the \"Stop\"/\"Close\" buttons in the GUI instead.", + parent=self.window, + ) + except KeyboardInterrupt: + # Repeated Ctrl+C while the modal warning is focused should not abort + # cleanup handling. + pass - self.close_window() + while self.gui_on: + try: + curr_time = time.time() + if self.task_on_button: + if curr_time - last_print > print_delay: + self.check_task_progress() + last_print = curr_time + elif self.task_on.get() == 1: + self.task_on.set(0) + if self.gui_task: + self.task.window.destroy() + + self.window.update() + except KeyboardInterrupt: + # Ctrl+C is not a supported way to close this GUI (see run_a_session.md): + # it can skip Teensy/Unity/socket cleanup, so just warn and keep running + # instead of exiting -- the experimenter should use "Close"/"Stop" instead. + _warn_use_gui_close() + + while True: + try: + self.close_window() + break + except KeyboardInterrupt: + # Keep trying to shut down even if Ctrl+C is pressed during teardown. + continue def main(): diff --git a/tests/unit/test_dlc_socket_behavior.py b/tests/unit/test_dlc_socket_behavior.py new file mode 100644 index 000000000..751da564c --- /dev/null +++ b/tests/unit/test_dlc_socket_behavior.py @@ -0,0 +1,249 @@ +import threading +import time +import unittest +from collections import deque +from unittest.mock import MagicMock, patch + +import numpy as np + +from teensyexp.tasks_abc.dlc_deque_socket import DLCClient as DequeSocketClient +from teensyexp.tasks_abc.dlc_socket import DLCClient as ListSocketClient + + +class TestDlcSocketCloseBehavior(unittest.TestCase): + """Regression tests for deterministic close behavior in DLC socket clients.""" + + def _assert_close_waits_for_reader_shutdown(self, client_cls): + connect_started = threading.Event() + + def _blocking_connect(self): + connect_started.set() + time.sleep(2.2) + raise TimeoutError("simulated connect stall") + + with patch.object(client_cls, "_connect_with_timeout", new=_blocking_connect): + client = client_cls(address=("localhost", 6000)) + self.assertTrue(connect_started.wait(timeout=1), "reader never reached connect") + + read_thread = client._read_thread + self.assertTrue(read_thread.is_alive()) + + start = time.monotonic() + client.close() + elapsed = time.monotonic() - start + + self.assertGreaterEqual(elapsed, 2.0) + self.assertFalse(read_thread.is_alive()) + self.assertIsNone(getattr(client, "conn", None)) + + def test_list_buffer_client_close_waits_for_thread_exit(self): + self._assert_close_waits_for_reader_shutdown(ListSocketClient) + + def test_deque_buffer_client_close_waits_for_thread_exit(self): + self._assert_close_waits_for_reader_shutdown(DequeSocketClient) + + +class _FakeConn: + def __init__(self, recv_side_effects): + self._recv_side_effects = list(recv_side_effects) + self.closed = False + + def recv(self): + if not self._recv_side_effects: + raise EOFError() + value = self._recv_side_effects.pop(0) + if isinstance(value, Exception): + raise value + return value + + def close(self): + self.closed = True + + +class TestDlcSocketReadBehavior(unittest.TestCase): + def _make_client_without_thread(self, client_cls): + with patch.object(client_cls, "start_read_buffer", return_value=None): + client = client_cls(address=("localhost", 6000)) + if isinstance(client, DequeSocketClient): + client.input_data = deque() + return client + + def _assert_one_payload_received(self, client, payload): + if isinstance(client, ListSocketClient): + self.assertEqual(len(client.input_data), 1) + self.assertEqual(client.input_data[0][1], payload) + self.assertIsInstance(client.input_data[0][0], float) + else: + self.assertEqual(list(client.input_data), [payload]) + + def _assert_happy_path(self, client_cls): + payload = {"x": 1, "y": 2} + fake_conn = _FakeConn([payload, EOFError()]) + client = self._make_client_without_thread(client_cls) + + with patch.object(client, "_connect_with_timeout", return_value=fake_conn): + client.read_on_thread() + + self._assert_one_payload_received(client, payload) + self.assertFalse(client.reading) + self.assertTrue(fake_conn.closed) + self.assertIsNone(getattr(client, "conn", None)) + + def test_list_buffer_happy_path_receives_payload(self): + self._assert_happy_path(ListSocketClient) + + def test_deque_buffer_happy_path_receives_payload(self): + self._assert_happy_path(DequeSocketClient) + + def _assert_timeout_retries_then_receives(self, client_cls): + payload = "frame" + fake_conn = _FakeConn([payload, EOFError()]) + client = self._make_client_without_thread(client_cls) + + attempts = {"count": 0} + + def _connect_attempt(): + attempts["count"] += 1 + if attempts["count"] < 3: + raise TimeoutError("retry") + return fake_conn + + with patch.object(client, "_connect_with_timeout", side_effect=_connect_attempt): + client.read_on_thread() + + self.assertEqual(attempts["count"], 3) + self._assert_one_payload_received(client, payload) + self.assertTrue(fake_conn.closed) + + def test_list_buffer_retries_timeouts_then_connects(self): + self._assert_timeout_retries_then_receives(ListSocketClient) + + def test_deque_buffer_retries_timeouts_then_connects(self): + self._assert_timeout_retries_then_receives(DequeSocketClient) + + def _assert_recv_oserror_stops_reader(self, client_cls): + fake_conn = _FakeConn([OSError("recv interrupted")]) + client = self._make_client_without_thread(client_cls) + + with patch.object(client, "_connect_with_timeout", return_value=fake_conn): + client.read_on_thread() + + self.assertFalse(client.reading) + self.assertTrue(fake_conn.closed) + self.assertIsNone(getattr(client, "conn", None)) + + def test_list_buffer_recv_oserror_stops_reader(self): + self._assert_recv_oserror_stops_reader(ListSocketClient) + + def test_deque_buffer_recv_oserror_stops_reader(self): + self._assert_recv_oserror_stops_reader(DequeSocketClient) + + def _assert_non_timeout_connect_exception_propagates(self, client_cls): + client = self._make_client_without_thread(client_cls) + + with patch.object(client, "_connect_with_timeout", side_effect=ValueError("bad connect")): + with self.assertRaises(ValueError): + client.read_on_thread() + + self.assertIsNone(getattr(client, "conn", None)) + + def test_list_buffer_non_timeout_connect_exception_propagates(self): + self._assert_non_timeout_connect_exception_propagates(ListSocketClient) + + def test_deque_buffer_non_timeout_connect_exception_propagates(self): + self._assert_non_timeout_connect_exception_propagates(DequeSocketClient) + + +class TestDlcSocketPublicApi(unittest.TestCase): + def _make_client_without_thread(self, client_cls): + with patch.object(client_cls, "start_read_buffer", return_value=None): + client = client_cls(address=("localhost", 6000)) + if isinstance(client, DequeSocketClient): + client.input_data = deque() + return client + + def _assert_close_is_idempotent(self, client_cls): + client = self._make_client_without_thread(client_cls) + client.conn = MagicMock() + client._read_thread = None + + client.close() + client.close() + + self.assertFalse(client.reading) + self.assertTrue(client._stop_event.is_set()) + self.assertEqual(client.conn.close.call_count, 2) + + def test_list_buffer_close_is_idempotent(self): + self._assert_close_is_idempotent(ListSocketClient) + + def test_deque_buffer_close_is_idempotent(self): + self._assert_close_is_idempotent(DequeSocketClient) + + def test_list_buffer_read_returns_latest_item(self): + client = self._make_client_without_thread(ListSocketClient) + client.input_data = [[1.0, "old"], [2.0, "new"]] + + out = client.read() + + self.assertEqual(out["time"], 2.0) + self.assertEqual(out["vals"], "new") + + def test_list_buffer_read_returns_none_when_empty(self): + client = self._make_client_without_thread(ListSocketClient) + client.input_data = [] + self.assertIsNone(client.read()) + + def test_list_buffer_reset_clears_input_data(self): + client = self._make_client_without_thread(ListSocketClient) + client.input_data = [[1.0, "frame"]] + + client.reset() + + self.assertEqual(client.input_data, []) + + def test_list_buffer_get_input_data_returns_numpy_array(self): + client = self._make_client_without_thread(ListSocketClient) + client.input_data = [[1.0, "frame1"], [2.0, "frame2"]] + + out = client.get_input_data() + + self.assertIsInstance(out, np.ndarray) + self.assertEqual(out.shape[0], 2) + + def test_deque_buffer_read_returns_none_when_empty(self): + client = self._make_client_without_thread(DequeSocketClient) + client.input_data = deque() + self.assertIsNone(client.read()) + + def test_deque_buffer_read_pops_latest_and_clears_queue(self): + client = self._make_client_without_thread(DequeSocketClient) + client.input_data = deque(["old", "new"]) + + out = client.read() + + self.assertEqual(out["vals"], "new") + self.assertEqual(out["previous"], 0) + self.assertEqual(len(client.input_data), 0) + self.assertEqual(client.previous, "new") + + def test_deque_buffer_reset_clears_input_data(self): + client = self._make_client_without_thread(DequeSocketClient) + client.input_data = deque(["frame"]) + + client.reset() + + self.assertEqual(len(client.input_data), 0) + + def test_deque_buffer_get_input_data_returns_numpy_array(self): + client = self._make_client_without_thread(DequeSocketClient) + client.input_data = deque(["frame1", "frame2"]) + + out = client.get_input_data() + + self.assertIsInstance(out, np.ndarray) + self.assertEqual(out.shape[0], 2) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/unit/test_gui_transfer.py b/tests/unit/test_gui_transfer.py index 6ca35585f..c0227ac49 100644 --- a/tests/unit/test_gui_transfer.py +++ b/tests/unit/test_gui_transfer.py @@ -181,6 +181,81 @@ def test_get_type(transfer_module): assert get_type("Testmouse_2023-02-22_2.pickle") == "teensy_path" +def test_get_processed_files_includes_gui_output(transfer_module, tmp_path): + """ + Regression test: get_transfer_files(key="gui_output") used to fall back to + returning the *entire* transfer_file dict, since "gui_output" is not one of + the button-driven keys in self.keys. get_processed_files() then handed that + whole blob to move_files(), which crashed with KeyError('src') on submit. + """ + Transfer = transfer_module.Transfer + + class StubWidget: + main_layout = None + + transfer = Transfer(widget=StubWidget(), keys=["teensy_path"]) + + teensy_file = tmp_path / "Testmouse_2023-02-22_2.pickle" + teensy_file.write_text("p") + transfer._set_file("teensy_path", str(teensy_file)) + + npy_file = tmp_path / "Testmouse_2023-02-22_2.npy" + npy_file.write_text("n") + transfer.set_npy(str(npy_file)) + + processed = transfer.get_processed_files() + assert len(processed) == 2 + for info in processed: + assert isinstance(info, dict) + assert "src" in info + assert "filename" in info + + +def test_get_processed_files_moves_all_transferred_types_but_not_video( + transfer_module, tmp_path +): + """ + dlc_path/camera_path/proc_path get scp'd to the server just like teensy_path, + so their rig originals must also be queued for processed_path - previously + only teensy_path and gui_output were. video_path is the one exception: videos + stay on the rig and are never transferred, so it must NOT show up here. + """ + Transfer = transfer_module.Transfer + + class StubWidget: + main_layout = None + + keys = ["teensy_path", "dlc_path", "camera_path", "video_path", "proc_path"] + transfer = Transfer(widget=StubWidget(), keys=keys) + + for key, suffix in [ + ("teensy_path", ".pickle"), + ("dlc_path", "_DLC.hdf5"), + ("camera_path", "_TS.npy"), + ("video_path", "_VIDEO3.avi"), + ("proc_path", "_PROC"), + ]: + f = tmp_path / f"Testmouse_2023-02-22_2{suffix}" + f.write_text("x") + transfer._set_file(key, str(f)) + + npy_file = tmp_path / "Testmouse_2023-02-22_2.npy" + npy_file.write_text("n") + transfer.set_npy(str(npy_file)) + + processed = transfer.get_processed_files() + processed_filenames = {info["filename"] for info in processed} + + assert processed_filenames == { + "Testmouse_2023-02-22_2.pickle", + "Testmouse_2023-02-22_2_DLC.hdf5", + "Testmouse_2023-02-22_2_TS.npy", + "Testmouse_2023-02-22_2_PROC", + "Testmouse_2023-02-22_2.npy", + } + assert "Testmouse_2023-02-22_2_VIDEO3.avi" not in processed_filenames + + def test_transfer_file_localhost_copy(gui_modules, tmp_path): utils = gui_modules["utils"] src_dir = tmp_path / "src" @@ -265,6 +340,178 @@ def test_find_related_files(gui_modules, tmp_path): assert related["dlc_path"] == dlc +def test_camera_number_from_filename(): + from utils.session_files import camera_number_from_filename + + assert ( + camera_number_from_filename("TS_vr4mice_Yurumi_2026-07-23_1_CAMERA3.npy") == 3 + ) + assert camera_number_from_filename("vr4mice_Yurumi_2026-07-23_1_VIDEO3.avi") == 3 + assert ( + camera_number_from_filename("Imagingsource_Testmouse_2023-02-22_2_VIDEO.mp4") + is None + ) + assert camera_number_from_filename("Testmouse_2023-02-22_2.pickle") is None + + +def test_find_related_files_multi_camera_matches_selected_camera(gui_modules): + """ + On a multi-camera rig, picking the CAMERA3 timestamps file should find the + matching VIDEO3 file, not whichever camera number sorts first. + """ + from utils.session_files import find_related_files + + def get_type(filename): + if "VIDEO" in filename: + return "video_path" + if "CAMERA" in filename: + return "camera_path" + return "teensy_path" + + config_data = gui_modules["config_data"] + camera_dir = Path(config_data["camera_path"]) + video_dir = Path(config_data["video_path"]) + camera_dir.mkdir(parents=True, exist_ok=True) + video_dir.mkdir(parents=True, exist_ok=True) + + stem = "Yurumi_2026-07-23_1" + for n in (1, 2, 3): + (camera_dir / f"TS_vr4mice_{stem}_CAMERA{n}.npy").write_text("t") + (video_dir / f"vr4mice_{stem}_VIDEO{n}.avi").write_text("v") + + path_by_key = { + k: config_data[k] + for k in config_data + if k.endswith("_path") or k == "raw_data_src" + } + + related = find_related_files(stem, path_by_key, get_type, camera_number=3) + assert related["camera_path"] == camera_dir / f"TS_vr4mice_{stem}_CAMERA3.npy" + assert related["video_path"] == video_dir / f"vr4mice_{stem}_VIDEO3.avi" + + +def _setup_ambiguous_cameras(gui_modules, stem, camera_numbers, with_dlc=True): + def get_type(filename): + if "VIDEO" in filename: + return "video_path" + if "CAMERA" in filename: + return "camera_path" + if "DLC" in filename: + return "dlc_path" + return "teensy_path" + + config_data = gui_modules["config_data"] + camera_dir = Path(config_data["camera_path"]) + video_dir = Path(config_data["video_path"]) + dlc_dir = Path(config_data["dlc_path"]) + camera_dir.mkdir(parents=True, exist_ok=True) + video_dir.mkdir(parents=True, exist_ok=True) + dlc_dir.mkdir(parents=True, exist_ok=True) + + for n in camera_numbers: + (camera_dir / f"TS_vr4mice_{stem}_CAMERA{n}.npy").write_text("t") + (video_dir / f"vr4mice_{stem}_VIDEO{n}.avi").write_text("v") + if with_dlc: + (dlc_dir / f"vr4mice_{stem}_DLC.hdf5").write_text("d") + + path_by_key = { + k: config_data[k] + for k in config_data + if k.endswith("_path") or k == "raw_data_src" + } + return get_type, path_by_key, camera_dir, video_dir + + +def test_find_related_files_multi_camera_leaves_blank_when_default_absent( + gui_modules, +): + """ + Picking a file with no camera suffix (e.g. DLC output) gives no camera + number to match on. When DEFAULT_CAMERA_NUMBER (3) isn't among the + present cameras, there's no safe default to autocomplete to, so those + roles are left out of the result rather than falling back to e.g. the + highest camera number. + """ + from utils.session_files import find_related_files + + stem = "Yurumi_2026-07-23_1" + get_type, path_by_key, camera_dir, video_dir = _setup_ambiguous_cameras( + gui_modules, stem, (1, 2, 4) + ) + + related = find_related_files(stem, path_by_key, get_type, camera_number=None) + assert "camera_path" not in related + assert "video_path" not in related + + +def test_find_related_files_multi_camera_prefers_default_camera_when_ambiguous( + gui_modules, +): + """ + When DEFAULT_CAMERA_NUMBER (3) is among the present cameras, it should be + preferred over the max camera number. + """ + from utils.session_files import find_related_files + + stem = "Yurumi_2026-07-23_1" + get_type, path_by_key, camera_dir, video_dir = _setup_ambiguous_cameras( + gui_modules, stem, (1, 2, 3, 4) + ) + + related = find_related_files(stem, path_by_key, get_type, camera_number=None) + assert related["camera_path"] == camera_dir / f"TS_vr4mice_{stem}_CAMERA3.npy" + assert related["video_path"] == video_dir / f"vr4mice_{stem}_VIDEO3.avi" + + +def test_find_related_files_explicit_camera_absent_leaves_role_unfilled(gui_modules): + """ + If the requested camera_number has no matching file for a role, that + role is simply left out of the result rather than falling back. + """ + from utils.session_files import find_related_files + + stem = "Yurumi_2026-07-23_1" + get_type, path_by_key, camera_dir, video_dir = _setup_ambiguous_cameras( + gui_modules, stem, (1, 2), with_dlc=False + ) + + related = find_related_files(stem, path_by_key, get_type, camera_number=3) + assert "camera_path" not in related + assert "video_path" not in related + + +def test_find_related_files_explicit_camera_prefers_exact_over_unnumbered( + gui_modules, +): + """ + When camera_number is set and a role has both an exact-numbered match + and a legacy unnumbered file, the exact match must win. + """ + from utils.session_files import find_related_files + + def get_type(filename): + return "video_path" + + config_data = gui_modules["config_data"] + video_dir = Path(config_data["video_path"]) + video_dir.mkdir(parents=True, exist_ok=True) + + stem = "Yurumi_2026-07-23_1" + legacy = video_dir / f"vr4mice_{stem}_VIDEO.avi" + exact = video_dir / f"vr4mice_{stem}_VIDEO3.avi" + legacy.write_text("legacy") + exact.write_text("exact") + + path_by_key = { + k: config_data[k] + for k in config_data + if k.endswith("_path") or k == "raw_data_src" + } + + related = find_related_files(stem, path_by_key, get_type, camera_number=3) + assert related["video_path"] == exact + + def test_adjust_keys_uses_display_text(gui_modules): utils = gui_modules["utils"] info = {"Rig": "12 - AR"} diff --git a/tests/unit/test_teensy_behavior.py b/tests/unit/test_teensy_behavior.py new file mode 100644 index 000000000..806c0e246 --- /dev/null +++ b/tests/unit/test_teensy_behavior.py @@ -0,0 +1,181 @@ +import importlib +import importlib.util +import threading +import unittest +import sys +import types +from pathlib import Path +from unittest.mock import MagicMock, patch + +_MODULE_PATH = ( + Path(__file__).resolve().parents[2] + / "mouse_task" + / "latency_tests" + / "Teensy_latency" + / "TeensyLatency.py" +) +_SPEC = importlib.util.spec_from_file_location("TeensyLatency_module", _MODULE_PATH) +_MODULE = importlib.util.module_from_spec(_SPEC) + + +class _FakeSerialException(Exception): + pass + + +_SERIAL_STUB = types.SimpleNamespace(SerialException=_FakeSerialException) +with patch.dict(sys.modules, {"serial": _SERIAL_STUB}): + _SPEC.loader.exec_module(_MODULE) +TeensyLatency = _MODULE.TeensyLatency + + +def _load_teensy_experiment_gui(): + try: + from teensyexp.teensy_experiment import TeensyExperimentGUI + return TeensyExperimentGUI + except ModuleNotFoundError as err: + if err.name != "serial": + raise + with patch.dict(sys.modules, {"serial": MagicMock()}): + module = importlib.import_module("teensyexp.teensy_experiment") + return module.TeensyExperimentGUI + + +TeensyExperimentGUI = _load_teensy_experiment_gui() + + +class TestTeensyLatencyReadExceptions(unittest.TestCase): + """Regression tests for serial-read close-race exception handling.""" + + def _make_latency(self, readline_side_effect): + latency = TeensyLatency.__new__(TeensyLatency) + latency.reading_teensy = True + latency.stop_event = threading.Event() + latency.ser = MagicMock() + latency.ser.readline.side_effect = readline_side_effect + latency.input_data = [] + latency.input_data_time = [] + return latency + + def test_known_windows_byref_typeerror_is_swallowed(self): + latency = self._make_latency( + TypeError("byref() argument must be a ctypes instance, not 'NoneType'") + ) + + latency.read_on_thread() + + self.assertEqual(latency.ser.readline.call_count, 1) + + def test_unrelated_typeerror_is_raised(self): + latency = self._make_latency(TypeError("unexpected type problem")) + + with self.assertRaises(TypeError): + latency.read_on_thread() + + +class TestTeensyGuiCloseBehavior(unittest.TestCase): + """Regression tests for GUI close behavior and Ctrl+C handling.""" + + def test_init_task_unsaved_data_confirmation_cancel_reverts_ready(self): + gui = TeensyExperimentGUI.__new__(TeensyExperimentGUI) + gui.teensy = object() + gui.task_on_button = False + gui.task = object() + gui.saved_ok = False + gui.window = object() + gui.task_on = MagicMock() + + with patch("teensyexp.teensy_experiment.messagebox.askokcancel", return_value=False) as askokcancel: + gui.init_task() + + askokcancel.assert_called_once_with( + "Unsaved Data", + "The previous task's data has not been saved.\n" + "Click Cancel to save first, or OK to initialize a new task and discard access to the previous task data.", + parent=gui.window, + ) + gui.task_on.set.assert_called_once_with(0) + + def test_init_task_unsaved_data_confirmation_ok_keeps_ready_flow(self): + gui = TeensyExperimentGUI.__new__(TeensyExperimentGUI) + gui.teensy = object() + gui.task_on_button = False + gui.task = object() + gui.saved_ok = False + gui.window = object() + gui.task_on = MagicMock() + gui.task_name = MagicMock() + gui.task_name.get.return_value = "FakeTask" + gui.task_params = {"FakeTask": {}} + gui.task_module = types.SimpleNamespace(FakeTask=MagicMock(return_value=MagicMock())) + gui._reset_progress_labels = MagicMock() + gui.task_label = {} + + with patch("teensyexp.teensy_experiment.messagebox.askokcancel", return_value=True): + gui.init_task() + + gui.task_on.set.assert_any_call(-1) + self.assertNotIn(unittest.mock.call(0), gui.task_on.set.mock_calls) + + def test_check_close_unsaved_uses_parented_warning(self): + gui = TeensyExperimentGUI.__new__(TeensyExperimentGUI) + gui.task_on = MagicMock() + gui.task_on.get.return_value = 0 + gui.saved_ok = False + gui.gui_on = True + gui.window = object() + + with patch("teensyexp.teensy_experiment.messagebox.askokcancel", return_value=False) as askokcancel: + gui.check_close() + askokcancel.assert_called_once_with( + "Exit", + "ARE YOU SURE YOU SAVED YOUR Data?", + parent=gui.window, + ) + self.assertTrue(gui.gui_on) + + with patch("teensyexp.teensy_experiment.messagebox.askokcancel", return_value=True): + gui.check_close() + self.assertFalse(gui.gui_on) + + def test_run_experiment_repeated_keyboard_interrupt_still_closes(self): + gui = TeensyExperimentGUI.__new__(TeensyExperimentGUI) + gui.task_on_button = False + gui.task_on = MagicMock() + gui.task_on.get.return_value = 0 + gui.gui_task = None + gui.gui_on = True + + class _FakeWindow: + def __init__(self, owner): + self.owner = owner + self.calls = 0 + + def update(self): + self.calls += 1 + if self.calls == 1: + raise KeyboardInterrupt() + self.owner.gui_on = False + + gui.window = _FakeWindow(gui) + + close_calls = [] + + def _close_window(): + close_calls.append(1) + if len(close_calls) == 1: + raise KeyboardInterrupt() + + gui.close_window = _close_window + + with patch( + "teensyexp.teensy_experiment.messagebox.showwarning", + side_effect=KeyboardInterrupt, + ) as showwarning: + gui.run_experiment() + + self.assertEqual(showwarning.call_count, 1) + self.assertEqual(len(close_calls), 2) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/unit/test_unity_task_unit.py b/tests/unit/test_unity_task_unit.py new file mode 100644 index 000000000..044c3bf53 --- /dev/null +++ b/tests/unit/test_unity_task_unit.py @@ -0,0 +1,152 @@ +import importlib +import sys +import types +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +import numpy as np + + +def _load_unity_task_with_stubs(): + class FakeActionTuple: + def __init__(self): + self.continuous = None + + def add_continuous(self, arr): + self.continuous = arr + + class FakeChannel: + def __init__(self): + self._props = {} + + def list_properties(self): + return list(self._props.keys()) + + def get_property(self, name): + return self._props[name] + + class FakeFloatChannel: + pass + + env_module = types.ModuleType("mlagents_envs.environment") + env_module.ActionTuple = FakeActionTuple + env_module.UnityEnvironment = object + + params_module = types.ModuleType( + "mlagents_envs.side_channel.environment_parameters_channel" + ) + params_module.EnvironmentParametersChannel = FakeChannel + + float_module = types.ModuleType("mlagents_envs.side_channel.float_properties_channel") + float_module.FloatPropertiesChannel = FakeFloatChannel + + side_channel_module = types.ModuleType("mlagents_envs.side_channel") + mlagents_module = types.ModuleType("mlagents_envs") + + stubs = { + "cv2": MagicMock(), + "mlagents_envs": mlagents_module, + "mlagents_envs.environment": env_module, + "mlagents_envs.side_channel": side_channel_module, + "mlagents_envs.side_channel.environment_parameters_channel": params_module, + "mlagents_envs.side_channel.float_properties_channel": float_module, + } + + with patch.dict(sys.modules, stubs): + module = importlib.import_module("teensyexp.tasks_abc.unity_task") + module = importlib.reload(module) + + return module, FakeActionTuple + + +def _build_fake_env(terminal=False): + class FakeActionSpec: + continuous_size = 4 + discrete_size = [] + + @staticmethod + def is_continuous(): + return True + + @staticmethod + def is_discrete(): + return False + + class FakeStepResult: + def __init__(self): + self.obs = [np.array([[1.0, 2.0, 3.0]], dtype=np.float32)] + self.reward = 1.5 + + class FakeDecisionSteps: + def __init__(self): + self.obs = [np.array([[0.1, 0.2, 0.3]], dtype=np.float32)] + self._step = FakeStepResult() + + def __getitem__(self, _idx): + return self._step + + class FakeTerminalSteps: + def __init__(self, done): + self.agent_id = [0] if done else [] + self._step = FakeStepResult() + + def __getitem__(self, _idx): + return self._step + + class FakeEnv: + def __init__(self): + self.behavior_specs = { + "MockBehavior": SimpleNamespace( + observation_specs=[SimpleNamespace(shape=(3,))], + action_spec=FakeActionSpec(), + ) + } + self.reset_calls = 0 + self.step_calls = 0 + self.closed = False + self.last_action = None + + def reset(self): + self.reset_calls += 1 + + def get_steps(self, _agent): + return FakeDecisionSteps(), FakeTerminalSteps(terminal) + + def set_actions(self, _agent, action_tuple): + self.last_action = action_tuple + + def step(self): + self.step_calls += 1 + + def close(self): + self.closed = True + + return FakeEnv() + + +def test_unity_task_start_loop_stop_without_mlagents_runtime(): + module, fake_action_tuple_cls = _load_unity_task_with_stubs() + UnityTask = module.UnityTask + + teensy = MagicMock() + fake_env = _build_fake_env(terminal=False) + + with patch.object(module, "UnityEnvironment", return_value=fake_env): + task = UnityTask(teensy=teensy, env="fake_unity_build", epochs=[10]) + task.start() + + assert task.episode == 1 + assert task.agent == "MockBehavior" + teensy.write.assert_any_call("start") + + keep_running, info = task.loop() + + assert keep_running is True + assert "episode" in info + assert fake_env.step_calls == 1 + assert isinstance(fake_env.last_action, fake_action_tuple_cls) + assert fake_env.last_action.continuous.shape == (1, 4) + + task.stop() + teensy.write.assert_any_call("stop") + assert fake_env.closed is True