From 647a00d84db798fc6b139f4bb0fe9219490ce3e6 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?C=C3=A9lia=20Benquet?= Date: Fri, 3 Jul 2026 10:31:32 +0200 Subject: [PATCH 01/26] Add further reward ports tasks --- mouse_task/__init__.py | 3 + mouse_task/far_reward_mouse_detection_p2.py | 103 ++++++++++++++++++++ 2 files changed, 106 insertions(+) create mode 100644 mouse_task/far_reward_mouse_detection_p2.py diff --git a/mouse_task/__init__.py b/mouse_task/__init__.py index 9c038da5f..8274a7312 100644 --- a/mouse_task/__init__.py +++ b/mouse_task/__init__.py @@ -7,6 +7,9 @@ from .mouse_discrim_occluders import DiscriminationWithOccludersTask from .mouse_discrim_multioccluders import DiscriminationWithMultiOccludersTask +# White target contrast task with far reward ports +from .far_reward_mouse_detection_p2 import FarDetectionWithoutVelocityThresholdTask + # Black target contrast task from .inv_mouse_detection_p1 import InvDetectionNoVelThrTask from .inv_mouse_detection_p2 import InvDetectionVelThrTask diff --git a/mouse_task/far_reward_mouse_detection_p2.py b/mouse_task/far_reward_mouse_detection_p2.py new file mode 100644 index 000000000..de271059c --- /dev/null +++ b/mouse_task/far_reward_mouse_detection_p2.py @@ -0,0 +1,103 @@ +""" +Detection task with velocity threshold to initiate the trials. +""" + +import os + +os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID" +os.environ["CUDA_VISIBLE_DEVICES"] = "0,1" + +import pathlib +import time as time + +from mouse_task.task_active_sensing import ActiveSensingTask + + +config_name = pathlib.Path("task_config.json") +current_dir = pathlib.Path(__file__).parent +config_path = current_dir.joinpath(config_name) # default class constructor input + + + +class FarDetectionWithVelocityThresholdTask(ActiveSensingTask): + + """ + Detection task with velocity threshold to initiate the trials. + Reward ports are further away from the screen, by ~2 cm, compared to DetectionWithVelocityThresholdTask. + """ + + def __init__( + self, + teensy, + monitor=None, + write_video=False, + fps=60.0, + session_label=["ar_detection_velthr"], + epochs=[250], + epoch_labels=["single_teardrop"], + config_file_path=config_path, + reward_size=100, + cropped_image=[0, 530, 0, 510], + unity_arena_size=[-9, 9, -10, -2], + r_report_box=[5, 10, -6, -4], # changed from [5, 10, -4, -2] + l_report_box=[-10, -5, -6, -4], # changed from [-10, -5, -4, -2] + start_box=[-4, 4, -9, -5, 90], + rotate_camera=90.0, + prob_obj_on_left=0.5, + prob_block_coherence = 0.5, + mouse_report_delay=0.0, + slit_size=[19.0, 20.0, 2], + slit_depth=0.2, + target_selection=6.0, + distractor_selection=4.0, + occlusion_type=0.0, + camera_type=1.0, + target_spread=4.0, + target_rotation=0, + target_size=2.0, + target_height=3.0, + block_length=1.0, + start_box_delay=0.25, + velocity_threshold=10.0, + distractor=0.0, + grey_screen_active=0.0, + target_distance=3, + use_dlc=True, + ): + super().__init__( + teensy=teensy, + monitor=monitor, + write_video=write_video, + fps=fps, + session_label=session_label, + epochs=epochs, + epoch_labels=epoch_labels, + config_file_path=config_file_path, + reward_size=reward_size, + cropped_image=cropped_image, + unity_arena_size=unity_arena_size, + r_report_box=r_report_box, + l_report_box=l_report_box, + start_box=start_box, + rotate_camera=rotate_camera, + prob_obj_on_left=prob_obj_on_left, + prob_block_coherence=prob_block_coherence, + mouse_report_delay=mouse_report_delay, + slit_size=slit_size, + slit_depth=slit_depth, + target_selection=target_selection, + distractor_selection=distractor_selection, + occlusion_type=occlusion_type, + camera_type=camera_type, + target_spread=target_spread, + target_rotation=target_rotation, + target_size=target_size, + target_height=target_height, + block_length=block_length, + start_box_delay=start_box_delay, + velocity_threshold=velocity_threshold, + distractor=distractor, + grey_screen_active=grey_screen_active, + target_distance=target_distance, + use_dlc=use_dlc, + ) From 98050b7b7f9082447a8898e834ad6bc369bb79ae Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?L=C3=A9o=20Bruneau?= <94896339+leobruneau@users.noreply.github.com> Date: Fri, 3 Jul 2026 16:27:50 +0200 Subject: [PATCH 02/26] Modularize Active Sensing task configs into YAML + a class registry (#319) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * feat: add modular task config files and loader Add a common.yaml holding every task parameter's default, plus one tasks/.yaml per variant declaring only the entries that differ (along with its class_name and description). helpers.load_task_config() deep-merges the common defaults with a task's overrides, replacing the duplicated per-task parameter sets that previously lived in separate Python modules. Includes a config for far_reward_mouse_detection_p2, the one task variant not otherwise covered by an existing tasks/*.yaml. * feat: make ActiveSensingTask config-driven and generate task classes ActiveSensingTask now accepts a task_config argument selecting a task YAML; parameters resolve from the merged config, with any explicitly passed argument taking precedence. Missing parameters (no task_config and not passed) raise a clear error. _registry builds one ActiveSensingTask subclass per task YAML at import and exposes them as mouse_task.. The generated __init__ keeps a real, introspectable signature (every parameter as a keyword arg with its merged default) so the teensyexp GUI can still build its editable parameter form. * refactor: remove per-task subclass modules The 20 inv_*/mouse_*/shape_*/far_reward_* modules were near-identical copies differing only in parameter values, which now live in configs/tasks/*.yaml and are materialised as classes by the registry. * test: cover task-config merge, registry, and generated signatures Add tests verifying every task YAML yields an exposed ActiveSensingTask subclass, load_task_config merges common defaults with task overrides, the generated __init__ signatures expose every parameter with its merged default, and constructing a task resolves parameters correctly (explicit arguments overriding the config). * docs: describe the modular task-config registry Repoint links to the YAML configs and document how task variants are now defined (common.yaml + tasks/*.yaml) and materialised as classes by the registry, replacing the per-task module description. * Fix the modular task set up by testing on the rig (#320) * Fix task modularization * Refine further reward port distance based on rig --------- Co-authored-by: Célia Benquet <32598028+CeliaBenquet@users.noreply.github.com> --- docs/software_package/active_sensing_task.md | 2 +- docs/training_protocols/training_protocol.md | 18 +- mouse_task/__init__.py | 40 +--- mouse_task/_registry.py | 68 ++++++ mouse_task/configs/README.md | 57 +++++ mouse_task/configs/common.yaml | 57 +++++ .../tasks/far_reward_mouse_detection_p2.yaml | 28 +++ .../configs/tasks/inv_mouse_detection_p1.yaml | 21 ++ .../configs/tasks/inv_mouse_detection_p2.yaml | 19 ++ .../configs/tasks/inv_mouse_discrim.yaml | 11 + .../inv_mouse_discrim_multioccluders.yaml | 20 ++ .../tasks/inv_mouse_discrim_occluders.yaml | 17 ++ .../tasks/inv_shape_mouse_discrim.yaml | 19 ++ ...nv_shape_mouse_discrim_multioccluders.yaml | 22 ++ ..._shape_mouse_discrim_narrow_occluders.yaml | 20 ++ .../inv_shape_mouse_discrim_occluders.yaml | 20 ++ .../configs/tasks/mouse_detection_p1.yaml | 17 ++ .../configs/tasks/mouse_detection_p2.yaml | 15 ++ mouse_task/configs/tasks/mouse_discrim.yaml | 6 + .../tasks/mouse_discrim_multioccluders.yaml | 15 ++ .../tasks/mouse_discrim_occluders.yaml | 13 + .../tasks/shape_mouse_detection_p1.yaml | 22 ++ .../tasks/shape_mouse_detection_p2.yaml | 22 ++ .../configs/tasks/shape_mouse_discrim.yaml | 22 ++ .../shape_mouse_discrim_multioccluders.yaml | 25 ++ .../shape_mouse_discrim_narrow_occluders.yaml | 23 ++ .../tasks/shape_mouse_discrim_occluders.yaml | 23 ++ mouse_task/far_reward_mouse_detection_p2.py | 103 -------- mouse_task/helpers.py | 43 ++++ mouse_task/inv_mouse_detection_p1.py | 101 -------- mouse_task/inv_mouse_detection_p2.py | 103 -------- mouse_task/inv_mouse_discrim.py | 100 -------- .../inv_mouse_discrim_multioccluders.py | 104 -------- mouse_task/inv_mouse_discrim_occluders.py | 101 -------- mouse_task/inv_shape_mouse_discrim.py | 102 -------- .../inv_shape_mouse_discrim_multioccluders.py | 102 -------- ...nv_shape_mouse_discrim_narrow_occluders.py | 102 -------- .../inv_shape_mouse_discrim_occluders.py | 102 -------- mouse_task/mouse_detection_p1.py | 100 -------- mouse_task/mouse_detection_p2.py | 102 -------- mouse_task/mouse_discrim.py | 100 -------- mouse_task/mouse_discrim_multioccluders.py | 102 -------- mouse_task/mouse_discrim_occluders.py | 100 -------- mouse_task/shape_mouse_detection_p1.py | 100 -------- mouse_task/shape_mouse_detection_p2.py | 100 -------- mouse_task/shape_mouse_discrim.py | 102 -------- .../shape_mouse_discrim_multioccluders.py | 101 -------- .../shape_mouse_discrim_narrow_occluders.py | 101 -------- mouse_task/shape_mouse_discrim_occluders.py | 101 -------- mouse_task/task_active_sensing.py | 124 +++++++--- mouse_task/tests/test_task_registry.py | 222 ++++++++++++++++++ teensyexp/teensy_experiment.py | 28 ++- 52 files changed, 983 insertions(+), 2205 deletions(-) create mode 100644 mouse_task/_registry.py create mode 100644 mouse_task/configs/README.md create mode 100644 mouse_task/configs/common.yaml create mode 100644 mouse_task/configs/tasks/far_reward_mouse_detection_p2.yaml create mode 100644 mouse_task/configs/tasks/inv_mouse_detection_p1.yaml create mode 100644 mouse_task/configs/tasks/inv_mouse_detection_p2.yaml create mode 100644 mouse_task/configs/tasks/inv_mouse_discrim.yaml create mode 100644 mouse_task/configs/tasks/inv_mouse_discrim_multioccluders.yaml create mode 100644 mouse_task/configs/tasks/inv_mouse_discrim_occluders.yaml create mode 100644 mouse_task/configs/tasks/inv_shape_mouse_discrim.yaml create mode 100644 mouse_task/configs/tasks/inv_shape_mouse_discrim_multioccluders.yaml create mode 100644 mouse_task/configs/tasks/inv_shape_mouse_discrim_narrow_occluders.yaml create mode 100644 mouse_task/configs/tasks/inv_shape_mouse_discrim_occluders.yaml create mode 100644 mouse_task/configs/tasks/mouse_detection_p1.yaml create mode 100644 mouse_task/configs/tasks/mouse_detection_p2.yaml create mode 100644 mouse_task/configs/tasks/mouse_discrim.yaml create mode 100644 mouse_task/configs/tasks/mouse_discrim_multioccluders.yaml create mode 100644 mouse_task/configs/tasks/mouse_discrim_occluders.yaml create mode 100644 mouse_task/configs/tasks/shape_mouse_detection_p1.yaml create mode 100644 mouse_task/configs/tasks/shape_mouse_detection_p2.yaml create mode 100644 mouse_task/configs/tasks/shape_mouse_discrim.yaml create mode 100644 mouse_task/configs/tasks/shape_mouse_discrim_multioccluders.yaml create mode 100644 mouse_task/configs/tasks/shape_mouse_discrim_narrow_occluders.yaml create mode 100644 mouse_task/configs/tasks/shape_mouse_discrim_occluders.yaml delete mode 100644 mouse_task/far_reward_mouse_detection_p2.py delete mode 100644 mouse_task/inv_mouse_detection_p1.py delete mode 100644 mouse_task/inv_mouse_detection_p2.py delete mode 100644 mouse_task/inv_mouse_discrim.py delete mode 100644 mouse_task/inv_mouse_discrim_multioccluders.py delete mode 100644 mouse_task/inv_mouse_discrim_occluders.py delete mode 100644 mouse_task/inv_shape_mouse_discrim.py delete mode 100644 mouse_task/inv_shape_mouse_discrim_multioccluders.py delete mode 100644 mouse_task/inv_shape_mouse_discrim_narrow_occluders.py delete mode 100644 mouse_task/inv_shape_mouse_discrim_occluders.py delete mode 100644 mouse_task/mouse_detection_p1.py delete mode 100644 mouse_task/mouse_detection_p2.py delete mode 100644 mouse_task/mouse_discrim.py delete mode 100644 mouse_task/mouse_discrim_multioccluders.py delete mode 100644 mouse_task/mouse_discrim_occluders.py delete mode 100644 mouse_task/shape_mouse_detection_p1.py delete mode 100644 mouse_task/shape_mouse_detection_p2.py delete mode 100644 mouse_task/shape_mouse_discrim.py delete mode 100644 mouse_task/shape_mouse_discrim_multioccluders.py delete mode 100644 mouse_task/shape_mouse_discrim_narrow_occluders.py delete mode 100644 mouse_task/shape_mouse_discrim_occluders.py create mode 100644 mouse_task/tests/test_task_registry.py diff --git a/docs/software_package/active_sensing_task.md b/docs/software_package/active_sensing_task.md index b4b064064..eb9f71079 100644 --- a/docs/software_package/active_sensing_task.md +++ b/docs/software_package/active_sensing_task.md @@ -17,7 +17,7 @@ In this task, a mouse the mouse looks through a slit in the wall and has to repo ## Task active sensing - Outline -The python [class](https://github.com/MMathisLab/FreelyMovingVR4Mice/blob/main/mouse_task/task_active_sensing.py) acts as an interface between DLClive and the unity build and the teensy. This task can be imported in different task scripts such as [mouse_detection_p1.py](https://github.com/MMathisLab/FreelyMovingVR4Mice/blob/main/mouse_task/mouse_detection_p1.py) where the input parameters are changed for the different phases of training. These tasks then be selected from within the teensy experiments GUI using the task drop down menu. The parameters can also be manually edited by clicking clicking on the "edit" button. Here parameters such as reward size and the probability that the OOI will appear on the left can be set. This python script then logs all the data about the experiment such as the mouses position in the arena, when water was given and which side the mouse reported that the target was on. +The python [class](https://github.com/MMathisLab/FreelyMovingVR4Mice/blob/main/mouse_task/task_active_sensing.py) acts as an interface between DLClive and the unity build and the teensy. Rather than one near-identical script per task type, each task variant is now described by a small YAML file in [mouse_task/configs/tasks/](https://github.com/MMathisLab/FreelyMovingVR4Mice/tree/main/mouse_task/configs/tasks) that lists only the parameters which differ from the shared defaults in [configs/common.yaml](https://github.com/MMathisLab/FreelyMovingVR4Mice/blob/main/mouse_task/configs/common.yaml). At import time, [`mouse_task/_registry.py`](https://github.com/MMathisLab/FreelyMovingVR4Mice/blob/main/mouse_task/_registry.py) generates one `ActiveSensingTask` subclass per YAML (passing its name as the `task_config` argument), so each variant still appears as a named task in the teensy experiments GUI drop-down menu. The parameters can also be manually edited by clicking clicking on the "edit" button. Here parameters such as reward size and the probability that the OOI will appear on the left can be set. This python script then logs all the data about the experiment such as the mouses position in the arena, when water was given and which side the mouse reported that the target was on. In takes the form of a parent class over a base class (called `unity_task`) and receives inputs within the `__init__()` function. These inputs can easily be modified from within the teensy experiments GUI by first loading the task and clicking on the edit button. This will present you with a window where these inputs can be modified. When the task is run (by clicking ready, followed by start) these inputs are assigned to class variables so that they can be made available to all the methods of the class. diff --git a/docs/training_protocols/training_protocol.md b/docs/training_protocols/training_protocol.md index 38c78d1af..85286d4e7 100644 --- a/docs/training_protocols/training_protocol.md +++ b/docs/training_protocols/training_protocol.md @@ -202,8 +202,8 @@ May be combined in single arena session with {ref}`sec:arena-habituation`. - To encourage mice to form connection between position and visual stimuli, longer initial session (ex 75 minutes) may be necessary to allow enough serendipitous trial initiations to occur. - If mice are allowed to overtrain on this step their performance can reach up to 100%. - **Python Task** to be loaded in the GUI: - - [DetectionWithoutVelocityThresholdTask](../../mouse_task/mouse_detection_p1.py) - - [ShapeDetectionWithoutVelocityThresholdTask](../../mouse_task/mouse_shape_detection_p1.py) + - [DetectionWithoutVelocityThresholdTask](../../mouse_task/configs/tasks/mouse_detection_p1.yaml) + - [ShapeDetectionWithoutVelocityThresholdTask](../../mouse_task/configs/tasks/shape_mouse_detection_p1.yaml) - **Parameters**: - [**Contrast**](./contrast_discrimination_training_parameters.md#stage-1-p1) - [**Shape**](./shape_discrimination_training_parameters.md#stage-1-p1) @@ -213,8 +213,8 @@ May be combined in single arena session with {ref}`sec:arena-habituation`. - Standard arena preparation (see {ref}`sec:arena-habituation`) and control software preparation (see {ref}`sec:gui`). - Same as in **1.1** except we want the mice to learn to mark a stop at trial initiation so we lower `velocity_threshold` and increase `start_box_delay` from the task parameters compared to **1.1**. - **Python Task** to be loaded in the GUI: - - [DetectionWithVelocityThresholdTask](../../mouse_task/mouse_detection_p2.py) - - [ShapeDetectionWithVelocityThresholdTask](../../mouse_task/mouse_shape_detection_p2.py) + - [DetectionWithVelocityThresholdTask](../../mouse_task/configs/tasks/mouse_detection_p2.yaml) + - [ShapeDetectionWithVelocityThresholdTask](../../mouse_task/configs/tasks/shape_mouse_detection_p2.yaml) - **Parameters**: - [**Contrast**](./contrast_discrimination_training_parameters.md#stage-1-p2) - [**Shape**](./shape_discrimination_training_parameters.md#stage-1-p2) @@ -227,8 +227,8 @@ May be combined in single arena session with {ref}`sec:arena-habituation`. (see {ref}`sec:gui`). - At this stage, a distractor object is introduced so that the task becomes a discrimination task. The mouse gets no reward when going on the distractor object side. The distractor object can be selected by modifying the value of the `distractor_selection` parameter. Its presence (or absence) is controlled by the `distractor` parameter which can be either `1` (presence) or `0` (absence). - **Python Task** to be loaded in the GUI: - - [DiscriminationTask](../../mouse_task/mouse_discrim.py) - - [ShapeDiscrimination](../../mouse_task/mouse_shape_discrim.py) + - [DiscriminationTask](../../mouse_task/configs/tasks/mouse_discrim.yaml) + - [ShapeDiscrim](../../mouse_task/configs/tasks/shape_mouse_discrim.yaml) - **Parameters**: - [**Contrast**](./contrast_discrimination_training_parameters.md#stage-2) - [**Shape**](./shape_discrimination_training_parameters.md#stage-2) @@ -240,8 +240,8 @@ May be combined in single arena session with {ref}`sec:arena-habituation`. (see {ref}`sec:gui`). - Ooccluder walls are introduced to hide varying amounts of both the target and the distractor. - **Python Task** to be loaded in the GUI: - - [DiscriminationWithOccludersTask](../../mouse_task/mouse_discrim_occluders.py) - - [ShapeDiscriminationOccluders](../../mouse_task/mouse_shape_discrim_occluders.py) + - [DiscriminationWithOccludersTask](../../mouse_task/configs/tasks/mouse_discrim_occluders.yaml) + - [ShapeDiscrimOccluders](../../mouse_task/configs/tasks/shape_mouse_discrim_occluders.yaml) - **Parameters**: - [**Contrast**](./contrast_discrimination_training_parameters.md#stage-3) - [**Shape**](./shape_discrimination_training_parameters.md#stage-3) @@ -255,7 +255,7 @@ This is currently the **test stage**. Mice that have reached it should do **== 5 - Standard arena preparation (see {ref}`sec:arena-habituation`) and control software preparation (see {ref}`sec:gui`). - **Python Task** to be selected in the GUI: - - [DiscriminationWithMultiOccludersTask](../../mouse_task/mouse_discrim_multioccluders.py) + - [DiscriminationWithMultiOccludersTask](../../mouse_task/configs/tasks/mouse_discrim_multioccluders.yaml) - **Parameters**: - [**Contrast**](./contrast_discrimination_training_parameters.md#stage-4) - [**Shape**](./shape_discrimination_training_parameters.md#stage-4) diff --git a/mouse_task/__init__.py b/mouse_task/__init__.py index 8274a7312..2a42a88c3 100644 --- a/mouse_task/__init__.py +++ b/mouse_task/__init__.py @@ -1,32 +1,18 @@ -from .manual_water import ManualWater +"""Mouse task package. -# White target contrast task -from .mouse_detection_p1 import DetectionWithoutVelocityThresholdTask -from .mouse_detection_p2 import DetectionWithVelocityThresholdTask -from .mouse_discrim import DiscriminationTask -from .mouse_discrim_occluders import DiscriminationWithOccludersTask -from .mouse_discrim_multioccluders import DiscriminationWithMultiOccludersTask +Task variants are described by ``configs/tasks/.yaml`` (class name, +docstring, and the parameters that differ from ``configs/common.yaml``); the +corresponding ``ActiveSensingTask`` subclass is generated at import time by +:mod:`mouse_task._registry`. -# White target contrast task with far reward ports -from .far_reward_mouse_detection_p2 import FarDetectionWithoutVelocityThresholdTask +Generated classes are bound into this namespace, so ``mouse_task.`` +keeps working (e.g. the teensyexp GUI task list). +""" -# Black target contrast task -from .inv_mouse_detection_p1 import InvDetectionNoVelThrTask -from .inv_mouse_detection_p2 import InvDetectionVelThrTask -from .inv_mouse_discrim import InvDiscrimTask -from .inv_mouse_discrim_occluders import InvDiscrimOccludersTask -from .inv_mouse_discrim_multioccluders import InvDiscrimMultiOccludersTask +from .manual_water import ManualWater +from ._registry import build_task_classes -# White Pacman target shape task -from .shape_mouse_detection_p1 import ShapeDetectionWithoutVelocityThresholdTask -from .shape_mouse_detection_p2 import ShapeDetectionWithVelocityThresholdTask -from .shape_mouse_discrim import ShapeDiscrim -from .shape_mouse_discrim_occluders import ShapeDiscrimOccluders -from .shape_mouse_discrim_narrow_occluders import ShapeDiscrimNarrowOccluders -from .shape_mouse_discrim_multioccluders import ShapeDiscrimMultiOccluders +_TASK_CLASSES = build_task_classes() +globals().update(_TASK_CLASSES) -# Black teardrop target shape task -from .inv_shape_mouse_discrim import InvShapeDiscrim -from .inv_shape_mouse_discrim_occluders import InvShapeDiscrimOccluders -from .inv_shape_mouse_discrim_multioccluders import InvShapeDiscrimMultiOccluders -from .inv_shape_mouse_discrim_narrow_occluders import InvShapeDiscrimNarrowOccluders +__all__ = ["ManualWater", *sorted(_TASK_CLASSES)] diff --git a/mouse_task/_registry.py b/mouse_task/_registry.py new file mode 100644 index 000000000..ce5554999 --- /dev/null +++ b/mouse_task/_registry.py @@ -0,0 +1,68 @@ +"""Dynamic task-class registry. + +Each task variant is described by ``configs/tasks/.yaml`` (class name, +docstring, and only the parameters that differ from ``configs/common.yaml``). +This module generates one ``ActiveSensingTask`` subclass per YAML, with a +real, introspectable ``__init__`` signature (required by the teensyexp GUI, +which reads args/defaults via ``inspect.getargspec`` to build its editable +parameter form). + +Generated classes are re-exported from ``mouse_task/__init__.py`` as +``mouse_task.``. +""" + +import pathlib + +from mouse_task.helpers import CONFIG_DIR, _load_yaml, load_task_config +from mouse_task.task_active_sensing import _DEFAULT_CONFIG_FILE, ActiveSensingTask + +_TASKS_DIR = CONFIG_DIR / "tasks" + + +def _build_task_class(task_name: str) -> type: + """Generate an ``ActiveSensingTask`` subclass from ``tasks/.yaml``.""" + meta = _load_yaml(_TASKS_DIR / f"{task_name}.yaml") + class_name = meta["class_name"] + doc = meta.get("description") or "" + params = load_task_config(task_name) # common.yaml + task overrides + + # Build a real __init__ source so inspect.getargspec sees every param. + arg_lines = [f" {name}={value!r}," for name, value in params.items()] + fwd_lines = [f" {name}={name}," for name in params] + src = ( + "def __init__(\n" + " self,\n" + " teensy,\n" + + "\n".join(arg_lines) + "\n" + " config_file_path=_DEFAULT_CONFIG_FILE,\n" + " **kwargs,\n" + "):\n" + " ActiveSensingTask.__init__(\n" + " self,\n" + " teensy,\n" + f" task_config={task_name!r},\n" + " config_file_path=config_file_path,\n" + + "\n".join(fwd_lines) + "\n" + " **kwargs,\n" + " )\n" + ) + namespace = { + "ActiveSensingTask": ActiveSensingTask, + "_DEFAULT_CONFIG_FILE": _DEFAULT_CONFIG_FILE, + } + exec(src, namespace) # noqa: S102 - trusted, generated from in-repo configs + + return type( + class_name, + (ActiveSensingTask,), + {"__init__": namespace["__init__"], "__doc__": doc, "__module__": "mouse_task"}, + ) + + +def build_task_classes() -> dict[str, type]: + """Generate every task class, keyed by class name.""" + classes = {} + for path in sorted(_TASKS_DIR.glob("*.yaml")): + cls = _build_task_class(path.stem) + classes[cls.__name__] = cls + return classes diff --git a/mouse_task/configs/README.md b/mouse_task/configs/README.md new file mode 100644 index 000000000..90a20512e --- /dev/null +++ b/mouse_task/configs/README.md @@ -0,0 +1,57 @@ +# Task configs + +Modular configuration for the Active Sensing task variants, mirroring the +`rl_task/configs/` pattern. + +``` +configs/ +├── common.yaml # parameters shared by every task variant +└── tasks/ + ├── mouse_discrim.yaml + ├── shape_mouse_discrim.yaml + └── ... # one file per variant: only what differs +``` + +## How it works + +- `common.yaml` holds the default value for **every** task parameter. +- Each `tasks/.yaml` declares a `class_name`, a `description`, and a + `params:` block containing **only the entries that differ** from `common.yaml`. +- `mouse_task.helpers.load_task_config()` deep-merges the two. +- At import, `mouse_task._registry` generates one `ActiveSensingTask` subclass + per task YAML. The generated `__init__` keeps a real, introspectable signature + (every parameter as a keyword arg with its merged default) so the teensyexp + GUI can still build its editable parameter form, and it passes + `task_config=""` down to `ActiveSensingTask`. +- Generated classes are exposed as `mouse_task.` (used by the GUI + task list and the RL pipeline's `task_variant` strings). + +## Add a new task + +1. Create `tasks/my_task.yaml`: + + ```yaml + class_name: MyTask + description: | + One-line summary of the task. + + params: + session_label: ["ar_my_task"] + slit_size: [4.0, 4.0, 1] + # ...only the parameters that differ from common.yaml + ``` + +2. It is picked up automatically — `mouse_task.MyTask` becomes importable and + appears in the GUI task drop-down. No Python file to write. + +## Run a task directly + +```python +from mouse_task.task_active_sensing import ActiveSensingTask +task = ActiveSensingTask(teensy, task_config="shape_mouse_discrim") +``` + +Any parameter passed explicitly overrides the value from the merged config. +The machine-specific `config_file_path` (the `task_config.json` holding the +absolute Unity build path) is separate and defaults to +`mouse_task/task_config.json`. diff --git a/mouse_task/configs/common.yaml b/mouse_task/configs/common.yaml new file mode 100644 index 000000000..d39fbe42d --- /dev/null +++ b/mouse_task/configs/common.yaml @@ -0,0 +1,57 @@ +# Common task parameters. +# +# These are the parameter values shared across every task variant. A task +# config in ``configs/tasks/.yaml`` overrides only the entries that +# differ for that task; everything not listed there falls back to the values +# below (see ``mouse_task.helpers.load_task_config``). +# +# ``config_file_path`` (the machine-specific ``task_config.json`` holding the +# absolute Unity build path) is NOT a task parameter — it is handled in code +# and defaults to ``mouse_task/task_config.json``. + +# --- Operational / runtime (not task-specific behaviour) -------------------- +monitor: null # unused +write_video: false # record video output +fps: 60.0 # recorded-video frame rate (currently unused) +use_dlc: true # use the DLC tracking socket (set false only for debugging) + +# --- Session / trial structure --------------------------------------------- +session_label: ["ar_discrim"] +epochs: [250] +epoch_labels: ["dual_teardrop"] +reward_size: 100 # water-valve open time in ms (~3 µL) + +# --- Arena / camera geometry (rig-specific, constant across tasks) ---------- +cropped_image: [0, 530, 0, 510] # [left, right, top, bottom] +unity_arena_size: [-9, 9, -10, -2] # [left, right, top, bottom] +r_report_box: [5, 10, -4, -2] # right report box [left, right, top, bottom] +l_report_box: [-10, -5, -4, -2] # left report box [left, right, top, bottom] +start_box: [-4, 4, -9, -5, 90] # [left, right, top, bottom, angle] +rotate_camera: 90.0 # camera rotation, set once per rig + +# --- Trial-side / block logic ---------------------------------------------- +prob_obj_on_left: 0.5 +prob_block_coherence: 0.5 +block_length: 1.0 +mouse_report_delay: 0.0 + +# --- Occluder / slit geometry ---------------------------------------------- +slit_size: [4.0, 10.0, 2] # [min, max, n] or an explicit list of sizes (len > 3) +slit_depth: 0.2 # occluder wall thickness +occlusion_type: 0.0 # 0 = none, 1 = slit, 2 = central wall + +# --- Object selection / appearance ----------------------------------------- +target_selection: 6.0 # OOI object id (see ActiveSensingTask docstring) +distractor_selection: 4.0 # distractor object id +distractor: 1.0 # 0 = no distractor, 1 = distractor present +camera_type: 1.0 # 0 = on-axis, 1 = off-axis +target_spread: 4.0 # distance between targets +target_rotation: 0 # inward rotation of target tips +target_size: 2.0 +target_height: 3.0 +target_distance: 3 # target distance in y +grey_screen_active: 0.0 # 0 = no grey ITI screen, 1 = grey screen + +# --- Trial initiation ------------------------------------------------------ +start_box_delay: 0.25 # time below velocity threshold required in start box +velocity_threshold: 10.0 # max speed in start box to initiate a trial diff --git a/mouse_task/configs/tasks/far_reward_mouse_detection_p2.yaml b/mouse_task/configs/tasks/far_reward_mouse_detection_p2.yaml new file mode 100644 index 000000000..a0b2f2031 --- /dev/null +++ b/mouse_task/configs/tasks/far_reward_mouse_detection_p2.yaml @@ -0,0 +1,28 @@ +class_name: FarDetectionWithVelocityThresholdTask +description: | + Detection task with velocity threshold to initiate the trials. + + Reward ports are further away from the screen, by ~2 cm, compared to + DetectionWithVelocityThresholdTask. + +# Only the parameters that differ from configs/common.yaml. +params: + session_label: + - ar_detection_velthr + epoch_labels: + - single_teardrop + slit_size: + - 19.0 + - 20.0 + - 2 + r_report_box: + - 5 + - 10 + - -5 + - -3 + l_report_box: + - -10 + - -5 + - -5 + - -3 + distractor: 0.0 diff --git a/mouse_task/configs/tasks/inv_mouse_detection_p1.yaml b/mouse_task/configs/tasks/inv_mouse_detection_p1.yaml new file mode 100644 index 000000000..45fa1d036 --- /dev/null +++ b/mouse_task/configs/tasks/inv_mouse_detection_p1.yaml @@ -0,0 +1,21 @@ +class_name: InvDetectionNoVelThrTask +description: | + Inverse Detection task without velocity threshold to initiate the trials. + + Black teardrop as target, no velocity threshold. + +# Only the parameters that differ from configs/common.yaml. +params: + session_label: + - ar_det_no_velthr_inv + epoch_labels: + - single_teardrop + slit_size: + - 4.0 + - 4.0 + - 1 + target_selection: 4.0 + distractor_selection: 6.0 + start_box_delay: 0.1 + velocity_threshold: 20.0 + distractor: 0.0 diff --git a/mouse_task/configs/tasks/inv_mouse_detection_p2.yaml b/mouse_task/configs/tasks/inv_mouse_detection_p2.yaml new file mode 100644 index 000000000..4efda738c --- /dev/null +++ b/mouse_task/configs/tasks/inv_mouse_detection_p2.yaml @@ -0,0 +1,19 @@ +class_name: InvDetectionVelThrTask +description: | + Detection task with velocity threshold to initiate the trials. + + Black teardrop as target, with velocity threshold. + +# Only the parameters that differ from configs/common.yaml. +params: + session_label: + - ar_detection_velthr_inv + epoch_labels: + - single_teardrop + slit_size: + - 19.0 + - 20.0 + - 2 + target_selection: 4.0 + distractor_selection: 6.0 + distractor: 0.0 diff --git a/mouse_task/configs/tasks/inv_mouse_discrim.yaml b/mouse_task/configs/tasks/inv_mouse_discrim.yaml new file mode 100644 index 000000000..e2254caff --- /dev/null +++ b/mouse_task/configs/tasks/inv_mouse_discrim.yaml @@ -0,0 +1,11 @@ +class_name: InvDiscrimTask +description: | + Discrimination task (without occluders). + Black teardrop as target, white teardrop as distractor. + +# Only the parameters that differ from configs/common.yaml. +params: + session_label: + - ar_discrim_inv + target_selection: 4.0 + distractor_selection: 6.0 diff --git a/mouse_task/configs/tasks/inv_mouse_discrim_multioccluders.yaml b/mouse_task/configs/tasks/inv_mouse_discrim_multioccluders.yaml new file mode 100644 index 000000000..f041afe5f --- /dev/null +++ b/mouse_task/configs/tasks/inv_mouse_discrim_multioccluders.yaml @@ -0,0 +1,20 @@ +class_name: InvDiscrimMultiOccludersTask +description: | + Discrimination task with multiple occluders. + + Sampled from a log space these are [12.0, 8.48, 6.0, 4.2, 3]. + Black teardrop as target, white teardrop as distractor. + +# Only the parameters that differ from configs/common.yaml. +params: + session_label: + - ar_discrim_5_occluders_inv + slit_size: + - 12.0 + - 8.48 + - 6.0 + - 4.2 + - 3 + target_selection: 4.0 + distractor_selection: 6.0 + occlusion_type: 1.0 diff --git a/mouse_task/configs/tasks/inv_mouse_discrim_occluders.yaml b/mouse_task/configs/tasks/inv_mouse_discrim_occluders.yaml new file mode 100644 index 000000000..660aa979d --- /dev/null +++ b/mouse_task/configs/tasks/inv_mouse_discrim_occluders.yaml @@ -0,0 +1,17 @@ +class_name: InvDiscrimOccludersTask +description: | + Discrimination task with occluders on. + + Black teardrop as target, white teardrop as distractor. + +# Only the parameters that differ from configs/common.yaml. +params: + session_label: + - ar_discrim_occluders_inv + slit_size: + - 4.3 + - 12.0 + - 2 + target_selection: 4.0 + distractor_selection: 6.0 + occlusion_type: 1.0 diff --git a/mouse_task/configs/tasks/inv_shape_mouse_discrim.yaml b/mouse_task/configs/tasks/inv_shape_mouse_discrim.yaml new file mode 100644 index 000000000..f3d646d8e --- /dev/null +++ b/mouse_task/configs/tasks/inv_shape_mouse_discrim.yaml @@ -0,0 +1,19 @@ +class_name: InvShapeDiscrim +description: | + Discrimination for shape task,with velocity threshold for trial initiation, no occluders. + The mouse must report the location of the black teardrop object and ignore the wide Pacman. + +# Only the parameters that differ from configs/common.yaml. +params: + session_label: + - ar_shape_discrimination_inv + epoch_labels: + - pacman_vs_teardrop + slit_size: + - 4.0 + - 4.0 + - 1 + slit_depth: 0.02 + target_selection: 4.0 + distractor_selection: 9.0 + target_rotation: 15.0 diff --git a/mouse_task/configs/tasks/inv_shape_mouse_discrim_multioccluders.yaml b/mouse_task/configs/tasks/inv_shape_mouse_discrim_multioccluders.yaml new file mode 100644 index 000000000..c650a9765 --- /dev/null +++ b/mouse_task/configs/tasks/inv_shape_mouse_discrim_multioccluders.yaml @@ -0,0 +1,22 @@ +class_name: InvShapeDiscrimMultiOccluders +description: | + Shape discrimination task, with occluders of different sizes + The mouse must report the location of the black teardrop object and ignore the wide Pacman. + +# Only the parameters that differ from configs/common.yaml. +params: + session_label: + - ar_shape_discrim_multi_occluders_inv + epoch_labels: + - pacman_vs_teardrop + slit_size: + - 15.0 + - 10.78 + - 7.75 + - 5.57 + - 4.0 + slit_depth: 0.02 + target_selection: 4.0 + distractor_selection: 9.0 + occlusion_type: 1.0 + target_rotation: 15.0 diff --git a/mouse_task/configs/tasks/inv_shape_mouse_discrim_narrow_occluders.yaml b/mouse_task/configs/tasks/inv_shape_mouse_discrim_narrow_occluders.yaml new file mode 100644 index 000000000..93006a1fa --- /dev/null +++ b/mouse_task/configs/tasks/inv_shape_mouse_discrim_narrow_occluders.yaml @@ -0,0 +1,20 @@ +class_name: InvShapeDiscrimNarrowOccluders +description: | + Shape discrimination task, with occluders of different sizes + The mouse must report the location of the black teardrop object and ignore the wide Pacman. + +# Only the parameters that differ from configs/common.yaml. +params: + session_label: + - ar_shape_discrim_narrow_occluders_inv + epoch_labels: + - pacman_vs_teardrop + slit_size: + - 5.0 + - 15.0 + - 2 + slit_depth: 0.02 + target_selection: 4.0 + distractor_selection: 9.0 + occlusion_type: 1.0 + target_rotation: 15.0 diff --git a/mouse_task/configs/tasks/inv_shape_mouse_discrim_occluders.yaml b/mouse_task/configs/tasks/inv_shape_mouse_discrim_occluders.yaml new file mode 100644 index 000000000..33f8dafaa --- /dev/null +++ b/mouse_task/configs/tasks/inv_shape_mouse_discrim_occluders.yaml @@ -0,0 +1,20 @@ +class_name: InvShapeDiscrimOccluders +description: | + Shape discrimination task, with occluders of different sizes + The mouse must report the location of the black teardrop object and ignore the wide Pacman. + +# Only the parameters that differ from configs/common.yaml. +params: + session_label: + - ar_shape_discrim_occluders_inv + epoch_labels: + - pacman_vs_teardrop + slit_size: + - 8.0 + - 15.0 + - 2 + slit_depth: 0.02 + target_selection: 4.0 + distractor_selection: 9.0 + occlusion_type: 1.0 + target_rotation: 15.0 diff --git a/mouse_task/configs/tasks/mouse_detection_p1.yaml b/mouse_task/configs/tasks/mouse_detection_p1.yaml new file mode 100644 index 000000000..25be4f019 --- /dev/null +++ b/mouse_task/configs/tasks/mouse_detection_p1.yaml @@ -0,0 +1,17 @@ +class_name: DetectionWithoutVelocityThresholdTask +description: | + Detection task without velocity threshold to initiate the trials. + +# Only the parameters that differ from configs/common.yaml. +params: + session_label: + - ar_detection_no_velthr + epoch_labels: + - single_teardrop + slit_size: + - 4.0 + - 4.0 + - 1 + start_box_delay: 0.1 + velocity_threshold: 20.0 + distractor: 0.0 diff --git a/mouse_task/configs/tasks/mouse_detection_p2.yaml b/mouse_task/configs/tasks/mouse_detection_p2.yaml new file mode 100644 index 000000000..c8009fde8 --- /dev/null +++ b/mouse_task/configs/tasks/mouse_detection_p2.yaml @@ -0,0 +1,15 @@ +class_name: DetectionWithVelocityThresholdTask +description: | + Detection task with velocity threshold to initiate the trials. + +# Only the parameters that differ from configs/common.yaml. +params: + session_label: + - ar_detection_velthr + epoch_labels: + - single_teardrop + slit_size: + - 19.0 + - 20.0 + - 2 + distractor: 0.0 diff --git a/mouse_task/configs/tasks/mouse_discrim.yaml b/mouse_task/configs/tasks/mouse_discrim.yaml new file mode 100644 index 000000000..a7bc3f34f --- /dev/null +++ b/mouse_task/configs/tasks/mouse_discrim.yaml @@ -0,0 +1,6 @@ +class_name: DiscriminationTask +description: | + Discrimination task (without occluders). + +# No parameter overrides — uses configs/common.yaml verbatim. +params: {} diff --git a/mouse_task/configs/tasks/mouse_discrim_multioccluders.yaml b/mouse_task/configs/tasks/mouse_discrim_multioccluders.yaml new file mode 100644 index 000000000..2f9eda6fd --- /dev/null +++ b/mouse_task/configs/tasks/mouse_discrim_multioccluders.yaml @@ -0,0 +1,15 @@ +class_name: DiscriminationWithMultiOccludersTask +description: | + Discrimination task with multiple occluders, sampled from a log space these are [12.0, 8.48, 6.0, 4.24, 3]. + +# Only the parameters that differ from configs/common.yaml. +params: + session_label: + - ar_discrim_5_occluders + slit_size: + - 12.0 + - 8.48 + - 6.0 + - 4.2 + - 3 + occlusion_type: 1.0 diff --git a/mouse_task/configs/tasks/mouse_discrim_occluders.yaml b/mouse_task/configs/tasks/mouse_discrim_occluders.yaml new file mode 100644 index 000000000..56bc2c979 --- /dev/null +++ b/mouse_task/configs/tasks/mouse_discrim_occluders.yaml @@ -0,0 +1,13 @@ +class_name: DiscriminationWithOccludersTask +description: | + Discrimination task with occluders on. + +# Only the parameters that differ from configs/common.yaml. +params: + session_label: + - ar_discrim_occluders + slit_size: + - 4.3 + - 12.0 + - 2 + occlusion_type: 1.0 diff --git a/mouse_task/configs/tasks/shape_mouse_detection_p1.yaml b/mouse_task/configs/tasks/shape_mouse_detection_p1.yaml new file mode 100644 index 000000000..fee6aa8ad --- /dev/null +++ b/mouse_task/configs/tasks/shape_mouse_detection_p1.yaml @@ -0,0 +1,22 @@ +class_name: ShapeDetectionWithoutVelocityThresholdTask +description: | + Detection of shape task without velocity threshold to initiate the trials. + +# Only the parameters that differ from configs/common.yaml. +params: + session_label: + - ar_shape_detection_no_velthr + epoch_labels: + - single_wide_pacman + slit_size: + - 4.0 + - 4.0 + - 1 + slit_depth: 0.02 + target_selection: 13.0 + distractor_selection: 6.0 + target_spread: 3.0 + target_rotation: 15.0 + start_box_delay: 0.1 + distractor: 0.0 + target_distance: 4.0 diff --git a/mouse_task/configs/tasks/shape_mouse_detection_p2.yaml b/mouse_task/configs/tasks/shape_mouse_detection_p2.yaml new file mode 100644 index 000000000..ef74a3dd7 --- /dev/null +++ b/mouse_task/configs/tasks/shape_mouse_detection_p2.yaml @@ -0,0 +1,22 @@ +class_name: ShapeDetectionWithVelocityThresholdTask +description: | + Detection of shape task with a lower velocity threshold to initiate the trials. + +# Only the parameters that differ from configs/common.yaml. +params: + session_label: + - ar_shape_detection_velthr + epoch_labels: + - single_wide_pacman + slit_size: + - 4.0 + - 4.0 + - 1 + slit_depth: 0.02 + target_selection: 13.0 + distractor_selection: 6.0 + target_spread: 3.0 + target_rotation: 15.0 + velocity_threshold: 5.0 + distractor: 0.0 + target_distance: 4.0 diff --git a/mouse_task/configs/tasks/shape_mouse_discrim.yaml b/mouse_task/configs/tasks/shape_mouse_discrim.yaml new file mode 100644 index 000000000..5e833e2cd --- /dev/null +++ b/mouse_task/configs/tasks/shape_mouse_discrim.yaml @@ -0,0 +1,22 @@ +class_name: ShapeDiscrim +description: | + Discrimination for shape task,with velocity threshold for trial initiation, no occluders. + The mouse must report the location of the white Pacman object and ignore the teardrop. + +# Only the parameters that differ from configs/common.yaml. +params: + session_label: + - ar_shape_discrimination + epoch_labels: + - pacman_vs_teardrop + slit_size: + - 4.0 + - 4.0 + - 1 + slit_depth: 0.02 + target_selection: 13.0 + distractor_selection: 6.0 + target_spread: 3.0 + target_rotation: 15.0 + velocity_threshold: 5.0 + target_distance: 4.0 diff --git a/mouse_task/configs/tasks/shape_mouse_discrim_multioccluders.yaml b/mouse_task/configs/tasks/shape_mouse_discrim_multioccluders.yaml new file mode 100644 index 000000000..738e00527 --- /dev/null +++ b/mouse_task/configs/tasks/shape_mouse_discrim_multioccluders.yaml @@ -0,0 +1,25 @@ +class_name: ShapeDiscrimMultiOccluders +description: | + Shape discrimination task, with occluders of different sizes + Mouse must report the white pacman location and ignore the teardrop. + +# Only the parameters that differ from configs/common.yaml. +params: + session_label: + - ar_shape_discrim_multi_occluders + epoch_labels: + - pacman_vs_teardrop + slit_size: + - 15.0 + - 10.78 + - 7.75 + - 5.57 + - 4.0 + slit_depth: 0.02 + target_selection: 13.0 + distractor_selection: 6.0 + occlusion_type: 1.0 + target_spread: 3.0 + target_rotation: 15.0 + velocity_threshold: 5.0 + target_distance: 4.0 diff --git a/mouse_task/configs/tasks/shape_mouse_discrim_narrow_occluders.yaml b/mouse_task/configs/tasks/shape_mouse_discrim_narrow_occluders.yaml new file mode 100644 index 000000000..54290901c --- /dev/null +++ b/mouse_task/configs/tasks/shape_mouse_discrim_narrow_occluders.yaml @@ -0,0 +1,23 @@ +class_name: ShapeDiscrimNarrowOccluders +description: | + Shape discrimination task, with occluders of different sizes + Mouse must report the white pacman location and ignore the teardrop. + +# Only the parameters that differ from configs/common.yaml. +params: + session_label: + - ar_shape_discrim_narrow_occluders + epoch_labels: + - pacman_vs_teardrop + slit_size: + - 5.0 + - 15.0 + - 2 + slit_depth: 0.02 + target_selection: 13.0 + distractor_selection: 6.0 + occlusion_type: 1.0 + target_spread: 3.0 + target_rotation: 15.0 + velocity_threshold: 5.0 + target_distance: 4.0 diff --git a/mouse_task/configs/tasks/shape_mouse_discrim_occluders.yaml b/mouse_task/configs/tasks/shape_mouse_discrim_occluders.yaml new file mode 100644 index 000000000..e0a0d7812 --- /dev/null +++ b/mouse_task/configs/tasks/shape_mouse_discrim_occluders.yaml @@ -0,0 +1,23 @@ +class_name: ShapeDiscrimOccluders +description: | + Shape discrimination task, with occluders of different sizes + Mouse must report the white pacman location and ignore the teardrop. + +# Only the parameters that differ from configs/common.yaml. +params: + session_label: + - ar_shape_discrim_occluders + epoch_labels: + - pacman_vs_teardrop + slit_size: + - 8.0 + - 15.0 + - 2 + slit_depth: 0.02 + target_selection: 13.0 + distractor_selection: 6.0 + occlusion_type: 1.0 + target_spread: 3.0 + target_rotation: 15.0 + velocity_threshold: 5.0 + target_distance: 3.5 diff --git a/mouse_task/far_reward_mouse_detection_p2.py b/mouse_task/far_reward_mouse_detection_p2.py deleted file mode 100644 index de271059c..000000000 --- a/mouse_task/far_reward_mouse_detection_p2.py +++ /dev/null @@ -1,103 +0,0 @@ -""" -Detection task with velocity threshold to initiate the trials. -""" - -import os - -os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID" -os.environ["CUDA_VISIBLE_DEVICES"] = "0,1" - -import pathlib -import time as time - -from mouse_task.task_active_sensing import ActiveSensingTask - - -config_name = pathlib.Path("task_config.json") -current_dir = pathlib.Path(__file__).parent -config_path = current_dir.joinpath(config_name) # default class constructor input - - - -class FarDetectionWithVelocityThresholdTask(ActiveSensingTask): - - """ - Detection task with velocity threshold to initiate the trials. - Reward ports are further away from the screen, by ~2 cm, compared to DetectionWithVelocityThresholdTask. - """ - - def __init__( - self, - teensy, - monitor=None, - write_video=False, - fps=60.0, - session_label=["ar_detection_velthr"], - epochs=[250], - epoch_labels=["single_teardrop"], - config_file_path=config_path, - reward_size=100, - cropped_image=[0, 530, 0, 510], - unity_arena_size=[-9, 9, -10, -2], - r_report_box=[5, 10, -6, -4], # changed from [5, 10, -4, -2] - l_report_box=[-10, -5, -6, -4], # changed from [-10, -5, -4, -2] - start_box=[-4, 4, -9, -5, 90], - rotate_camera=90.0, - prob_obj_on_left=0.5, - prob_block_coherence = 0.5, - mouse_report_delay=0.0, - slit_size=[19.0, 20.0, 2], - slit_depth=0.2, - target_selection=6.0, - distractor_selection=4.0, - occlusion_type=0.0, - camera_type=1.0, - target_spread=4.0, - target_rotation=0, - target_size=2.0, - target_height=3.0, - block_length=1.0, - start_box_delay=0.25, - velocity_threshold=10.0, - distractor=0.0, - grey_screen_active=0.0, - target_distance=3, - use_dlc=True, - ): - super().__init__( - teensy=teensy, - monitor=monitor, - write_video=write_video, - fps=fps, - session_label=session_label, - epochs=epochs, - epoch_labels=epoch_labels, - config_file_path=config_file_path, - reward_size=reward_size, - cropped_image=cropped_image, - unity_arena_size=unity_arena_size, - r_report_box=r_report_box, - l_report_box=l_report_box, - start_box=start_box, - rotate_camera=rotate_camera, - prob_obj_on_left=prob_obj_on_left, - prob_block_coherence=prob_block_coherence, - mouse_report_delay=mouse_report_delay, - slit_size=slit_size, - slit_depth=slit_depth, - target_selection=target_selection, - distractor_selection=distractor_selection, - occlusion_type=occlusion_type, - camera_type=camera_type, - target_spread=target_spread, - target_rotation=target_rotation, - target_size=target_size, - target_height=target_height, - block_length=block_length, - start_box_delay=start_box_delay, - velocity_threshold=velocity_threshold, - distractor=distractor, - grey_screen_active=grey_screen_active, - target_distance=target_distance, - use_dlc=use_dlc, - ) diff --git a/mouse_task/helpers.py b/mouse_task/helpers.py index 0bd464739..b1cc4546d 100755 --- a/mouse_task/helpers.py +++ b/mouse_task/helpers.py @@ -5,6 +5,49 @@ import logging from pathlib import Path +import yaml + +# Directory holding the modular task configs (``common.yaml`` + ``tasks/*.yaml``). +CONFIG_DIR = Path(__file__).parent / "configs" + + +def _load_yaml(path: Path) -> dict: + with open(path) as f: + return yaml.safe_load(f) or {} + + +def _deep_merge(base: dict, override: dict) -> None: + """Recursively merge ``override`` into ``base`` in-place.""" + for key, val in override.items(): + if key in base and isinstance(base[key], dict) and isinstance(val, dict): + _deep_merge(base[key], val) + else: + base[key] = val + + +def load_task_config(task_config: str, config_dir: Path = CONFIG_DIR) -> dict: + """Load ``common.yaml`` and overlay the task-specific overrides. + + Args: + task_config: Stem of a file in ``configs/tasks/`` (e.g. ``"shape_mouse_discrim"``). + config_dir: Root of the config tree (defaults to ``mouse_task/configs``). + + Returns: + Flat dict of task parameters: every key from ``common.yaml`` with the + task's ``params`` block merged on top. + """ + params = _load_yaml(config_dir / "common.yaml") + + task_path = config_dir / "tasks" / f"{task_config}.yaml" + if not task_path.exists(): + available = sorted(p.stem for p in (config_dir / "tasks").glob("*.yaml")) + raise FileNotFoundError( + f"Task config not found: {task_path}. Available: {available}" + ) + _deep_merge(params, _load_yaml(task_path).get("params") or {}) + return params + + def process_config(config_file_path: Path) -> dict: """ Function that processes the task_config file and verifies its content diff --git a/mouse_task/inv_mouse_detection_p1.py b/mouse_task/inv_mouse_detection_p1.py deleted file mode 100644 index d866a9a25..000000000 --- a/mouse_task/inv_mouse_detection_p1.py +++ /dev/null @@ -1,101 +0,0 @@ -""" -Detection task without velocity threshold to initiate the trials. -""" - -import os - -os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID" -os.environ["CUDA_VISIBLE_DEVICES"] = "0,1" - -import pathlib - -from mouse_task.task_active_sensing import ActiveSensingTask - - -config_name = pathlib.Path("task_config.json") -current_dir = pathlib.Path(__file__).parent -config_path = current_dir.joinpath(config_name) # default class constructor input - - -class InvDetectionNoVelThrTask(ActiveSensingTask): - """ - Inverse Detection task without velocity threshold to initiate the trials. - - Black teardrop as target, no velocity threshold. - """ - - def __init__( - self, - teensy, - monitor=None, - write_video=False, - fps=60.0, - session_label=["ar_det_no_velthr_inv"], - epochs=[250], - epoch_labels=["single_teardrop"], - config_file_path=config_path, - reward_size=100, - cropped_image=[0, 530, 0, 510], - unity_arena_size=[-9, 9, -10, -2], - r_report_box=[5, 10, -4, -2], - l_report_box=[-10, -5, -4, -2], - start_box=[-4, 4, -9, -5, 90], - rotate_camera=90.0, - prob_obj_on_left=0.5, - prob_block_coherence=0.5, - mouse_report_delay=0.0, - slit_size=[4.0, 4.0, 1], - slit_depth=0.2, - target_selection=4.0, # changed to 4.0 from 6.0, black teardrop - distractor_selection=6.0, # changed to 6.0 from 4.0, white teardrop - occlusion_type=0.0, - camera_type=1.0, - target_spread=4.0, - target_rotation=0, - target_size=2.0, - target_height=3.0, - block_length=1.0, - start_box_delay=0.1, - velocity_threshold=20.0, - distractor=0.0, # no distractor - grey_screen_active=0.0, - target_distance=3, - use_dlc=True, - ): - super().__init__( - teensy=teensy, - monitor=monitor, - write_video=write_video, - fps=fps, - session_label=session_label, - epochs=epochs, - epoch_labels=epoch_labels, - config_file_path=config_file_path, - reward_size=reward_size, - cropped_image=cropped_image, - unity_arena_size=unity_arena_size, - r_report_box=r_report_box, - l_report_box=l_report_box, - start_box=start_box, - rotate_camera=rotate_camera, - prob_obj_on_left=prob_obj_on_left, - prob_block_coherence=prob_block_coherence, - mouse_report_delay=mouse_report_delay, - slit_size=slit_size, - slit_depth=slit_depth, - target_selection=target_selection, - distractor_selection=distractor_selection, - occlusion_type=occlusion_type, - camera_type=camera_type, - target_spread=target_spread, - target_rotation=target_rotation, - target_size=target_size, - target_height=target_height, - block_length=block_length, - start_box_delay=start_box_delay, - velocity_threshold=velocity_threshold, - distractor=distractor, - grey_screen_active=grey_screen_active, - target_distance=target_distance, - use_dlc=use_dlc, - ) diff --git a/mouse_task/inv_mouse_detection_p2.py b/mouse_task/inv_mouse_detection_p2.py deleted file mode 100644 index b30142bbb..000000000 --- a/mouse_task/inv_mouse_detection_p2.py +++ /dev/null @@ -1,103 +0,0 @@ -""" -Detection task with velocity threshold to initiate the trials. -""" - -import os - -os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID" -os.environ["CUDA_VISIBLE_DEVICES"] = "0,1" - -import pathlib - -from mouse_task.task_active_sensing import ActiveSensingTask - - -config_name = pathlib.Path("task_config.json") -current_dir = pathlib.Path(__file__).parent -config_path = current_dir.joinpath(config_name) # default class constructor input - - - -class InvDetectionVelThrTask(ActiveSensingTask): - - """ - Detection task with velocity threshold to initiate the trials. - - Black teardrop as target, with velocity threshold. - """ - - def __init__( - self, - teensy, - monitor=None, - write_video=False, - fps=60.0, - session_label=["ar_detection_velthr_inv"], - epochs=[250], - epoch_labels=["single_teardrop"], - config_file_path=config_path, - reward_size=100, - cropped_image=[0, 530, 0, 510], - unity_arena_size=[-9, 9, -10, -2], - r_report_box=[5, 10, -4, -2], - l_report_box=[-10, -5, -4, -2], - start_box=[-4, 4, -9, -5, 90], - rotate_camera=90.0, - prob_obj_on_left=0.5, - prob_block_coherence=0.5, - mouse_report_delay=0.0, - slit_size=[19.0, 20.0, 2], - slit_depth=0.2, - target_selection=4.0, # changed to 4.0 from 6.0, black teardrop - distractor_selection=6.0, # changed to 6.0 from 4.0, white teardrop - occlusion_type=0.0, - camera_type=1.0, - target_spread=4.0, - target_rotation=0, - target_size=2.0, - target_height=3.0, - block_length=1.0, - start_box_delay=0.25, - velocity_threshold=10.0, - distractor=0.0, - grey_screen_active=0.0, - target_distance=3, - use_dlc=True, - ): - super().__init__( - teensy=teensy, - monitor=monitor, - write_video=write_video, - fps=fps, - session_label=session_label, - epochs=epochs, - epoch_labels=epoch_labels, - config_file_path=config_file_path, - reward_size=reward_size, - cropped_image=cropped_image, - unity_arena_size=unity_arena_size, - r_report_box=r_report_box, - l_report_box=l_report_box, - start_box=start_box, - rotate_camera=rotate_camera, - prob_obj_on_left=prob_obj_on_left, - prob_block_coherence=prob_block_coherence, - mouse_report_delay=mouse_report_delay, - slit_size=slit_size, - slit_depth=slit_depth, - target_selection=target_selection, - distractor_selection=distractor_selection, - occlusion_type=occlusion_type, - camera_type=camera_type, - target_spread=target_spread, - target_rotation=target_rotation, - target_size=target_size, - target_height=target_height, - block_length=block_length, - start_box_delay=start_box_delay, - velocity_threshold=velocity_threshold, - distractor=distractor, - grey_screen_active=grey_screen_active, - target_distance=target_distance, - use_dlc=use_dlc, - ) diff --git a/mouse_task/inv_mouse_discrim.py b/mouse_task/inv_mouse_discrim.py deleted file mode 100644 index 1eb145e34..000000000 --- a/mouse_task/inv_mouse_discrim.py +++ /dev/null @@ -1,100 +0,0 @@ -""" -Discrimination task (without occluders). -""" - -import os - -os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID" -os.environ["CUDA_VISIBLE_DEVICES"] = "0,1" - -import pathlib - -from mouse_task.task_active_sensing import ActiveSensingTask - - -config_name = pathlib.Path("task_config.json") -current_dir = pathlib.Path(__file__).parent -config_path = current_dir.joinpath(config_name) # default class constructor input - - -class InvDiscrimTask(ActiveSensingTask): - """ - Discrimination task (without occluders). - Black teardrop as target, white teardrop as distractor. - """ - - def __init__( - self, - teensy, - monitor=None, - write_video=False, - fps=60.0, - session_label=["ar_discrim_inv"], - epochs=[250], - epoch_labels=["dual_teardrop"], - config_file_path=config_path, - reward_size=100, - cropped_image=[0, 530, 0, 510], - unity_arena_size=[-9, 9, -10, -2], - r_report_box=[5, 10, -4, -2], - l_report_box=[-10, -5, -4, -2], - start_box=[-4, 4, -9, -5, 90], - rotate_camera=90.0, - prob_obj_on_left=0.5, - prob_block_coherence=0.5, - mouse_report_delay=0.0, - slit_size=[4.0, 10.0, 2], - slit_depth=0.2, - target_selection=4.0, # changed to 4.0 from 6.0, black teardrop - distractor_selection=6.0, # changed to 6.0 from 4.0, white teardrop - occlusion_type=0.0, - camera_type=1.0, - target_spread=4.0, - target_rotation=0, - target_size=2.0, - target_height=3.0, - block_length=1.0, - start_box_delay=0.25, - velocity_threshold=10.0, - distractor=1.0, - grey_screen_active=0.0, - target_distance=3, - use_dlc=True, - ): - super().__init__( - teensy=teensy, - monitor=monitor, - write_video=write_video, - fps=fps, - session_label=session_label, - epochs=epochs, - epoch_labels=epoch_labels, - config_file_path=config_file_path, - reward_size=reward_size, - cropped_image=cropped_image, - unity_arena_size=unity_arena_size, - r_report_box=r_report_box, - l_report_box=l_report_box, - start_box=start_box, - rotate_camera=rotate_camera, - prob_obj_on_left=prob_obj_on_left, - prob_block_coherence=prob_block_coherence, - mouse_report_delay=mouse_report_delay, - slit_size=slit_size, - slit_depth=slit_depth, - target_selection=target_selection, - distractor_selection=distractor_selection, - occlusion_type=occlusion_type, - camera_type=camera_type, - target_spread=target_spread, - target_rotation=target_rotation, - target_size=target_size, - target_height=target_height, - block_length=block_length, - start_box_delay=start_box_delay, - velocity_threshold=velocity_threshold, - distractor=distractor, - grey_screen_active=grey_screen_active, - target_distance=target_distance, - use_dlc=use_dlc, - ) diff --git a/mouse_task/inv_mouse_discrim_multioccluders.py b/mouse_task/inv_mouse_discrim_multioccluders.py deleted file mode 100644 index d8caab0e0..000000000 --- a/mouse_task/inv_mouse_discrim_multioccluders.py +++ /dev/null @@ -1,104 +0,0 @@ -""" -Discrimination task with occluders on. -""" - -import os - -os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID" -os.environ["CUDA_VISIBLE_DEVICES"] = "0,1" - -import pathlib - -from mouse_task.task_active_sensing import ActiveSensingTask - - -config_name = pathlib.Path("task_config.json") -current_dir = pathlib.Path(__file__).parent -config_path = current_dir.joinpath(config_name) # default class constructor input - - - -class InvDiscrimMultiOccludersTask(ActiveSensingTask): - - """ - Discrimination task with multiple occluders. - - Sampled from a log space these are [12.0, 8.48, 6.0, 4.2, 3]. - Black teardrop as target, white teardrop as distractor. - """ - - def __init__( - self, - teensy, - monitor=None, - write_video=False, - fps=60.0, - session_label=["ar_discrim_5_occluders_inv"], - epochs=[250], - epoch_labels=["dual_teardrop"], - config_file_path=config_path, - reward_size=100, - cropped_image=[0, 530, 0, 510], - unity_arena_size=[-9, 9, -10, -2], - r_report_box=[5, 10, -4, -2], - l_report_box=[-10, -5, -4, -2], - start_box=[-4, 4, -9, -5, 90], - rotate_camera=90.0, - prob_obj_on_left=0.5, - prob_block_coherence=0.5, - mouse_report_delay=0.0, - slit_size=[12.0, 8.48, 6., 4.2, 3], - slit_depth=0.2, - target_selection=4.0, # changed to 4.0 from 6.0, black teardrop - distractor_selection=6.0, # changed to 6.0 from 4.0, white teardrop - occlusion_type=1.0, - camera_type=1.0, - target_spread=4.0, - target_rotation=0, - target_size=2.0, - target_height=3.0, - block_length=1.0, - start_box_delay=0.25, - velocity_threshold=10.0, - distractor=1.0, - grey_screen_active=0.0, - target_distance=3, - use_dlc=True, - ): - super().__init__( - teensy=teensy, - monitor=monitor, - write_video=write_video, - fps=fps, - session_label=session_label, - epochs=epochs, - epoch_labels=epoch_labels, - config_file_path=config_file_path, - reward_size=reward_size, - cropped_image=cropped_image, - unity_arena_size=unity_arena_size, - r_report_box=r_report_box, - l_report_box=l_report_box, - start_box=start_box, - rotate_camera=rotate_camera, - prob_obj_on_left=prob_obj_on_left, - prob_block_coherence=prob_block_coherence, - mouse_report_delay=mouse_report_delay, - slit_size=slit_size, - slit_depth=slit_depth, - target_selection=target_selection, - distractor_selection=distractor_selection, - occlusion_type=occlusion_type, - camera_type=camera_type, - target_spread=target_spread, - target_rotation=target_rotation, - target_size=target_size, - target_height=target_height, - block_length=block_length, - start_box_delay=start_box_delay, - velocity_threshold=velocity_threshold, - distractor=distractor, - grey_screen_active=grey_screen_active, - target_distance=target_distance, - use_dlc=use_dlc, - ) diff --git a/mouse_task/inv_mouse_discrim_occluders.py b/mouse_task/inv_mouse_discrim_occluders.py deleted file mode 100644 index 55423cbaa..000000000 --- a/mouse_task/inv_mouse_discrim_occluders.py +++ /dev/null @@ -1,101 +0,0 @@ -""" -Discrimination task with occluders on. -""" - -import os - -os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID" -os.environ["CUDA_VISIBLE_DEVICES"] = "0,1" - -import pathlib - -from mouse_task.task_active_sensing import ActiveSensingTask - - -config_name = pathlib.Path("task_config.json") -current_dir = pathlib.Path(__file__).parent -config_path = current_dir.joinpath(config_name) # default class constructor input - - -class InvDiscrimOccludersTask(ActiveSensingTask): - """ - Discrimination task with occluders on. - - Black teardrop as target, white teardrop as distractor. - """ - - def __init__( - self, - teensy, - monitor=None, - write_video=False, - fps=60.0, - session_label=["ar_discrim_occluders_inv"], - epochs=[250], - epoch_labels=["dual_teardrop"], - config_file_path=config_path, - reward_size=100, - cropped_image=[0, 530, 0, 510], - unity_arena_size=[-9, 9, -10, -2], - r_report_box=[5, 10, -4, -2], - l_report_box=[-10, -5, -4, -2], - start_box=[-4, 4, -9, -5, 90], - rotate_camera=90.0, - prob_obj_on_left=0.5, - prob_block_coherence=0.5, - mouse_report_delay=0.0, - slit_size=[4.3, 12.0, 2], - slit_depth=0.2, - target_selection=4.0, # changed to 4.0 from 6.0, black teardrop - distractor_selection=6.0, # changed to 6.0 from 4.0, white teardrop - occlusion_type=1.0, - camera_type=1.0, - target_spread=4.0, - target_rotation=0, - target_size=2.0, - target_height=3.0, - block_length=1.0, - start_box_delay=0.25, - velocity_threshold=10.0, - distractor=1.0, - grey_screen_active=0.0, - target_distance=3, - use_dlc=True, - ): - super().__init__( - teensy=teensy, - monitor=monitor, - write_video=write_video, - fps=fps, - session_label=session_label, - epochs=epochs, - epoch_labels=epoch_labels, - config_file_path=config_file_path, - reward_size=reward_size, - cropped_image=cropped_image, - unity_arena_size=unity_arena_size, - r_report_box=r_report_box, - l_report_box=l_report_box, - start_box=start_box, - rotate_camera=rotate_camera, - prob_obj_on_left=prob_obj_on_left, - prob_block_coherence=prob_block_coherence, - mouse_report_delay=mouse_report_delay, - slit_size=slit_size, - slit_depth=slit_depth, - target_selection=target_selection, - distractor_selection=distractor_selection, - occlusion_type=occlusion_type, - camera_type=camera_type, - target_spread=target_spread, - target_rotation=target_rotation, - target_size=target_size, - target_height=target_height, - block_length=block_length, - start_box_delay=start_box_delay, - velocity_threshold=velocity_threshold, - distractor=distractor, - grey_screen_active=grey_screen_active, - target_distance=target_distance, - use_dlc=use_dlc, - ) diff --git a/mouse_task/inv_shape_mouse_discrim.py b/mouse_task/inv_shape_mouse_discrim.py deleted file mode 100644 index c8b31fc92..000000000 --- a/mouse_task/inv_shape_mouse_discrim.py +++ /dev/null @@ -1,102 +0,0 @@ -""" -Discrimination task without occluders. -Mouse must report the location of the black teardrop and ignore the pacman. -""" - -import os - -os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID" -os.environ["CUDA_VISIBLE_DEVICES"] = "0,1" - -import pathlib -import time as time - -from mouse_task.task_active_sensing import ActiveSensingTask - - -config_name = pathlib.Path("task_config.json") -current_dir = pathlib.Path(__file__).parent -config_path = current_dir.joinpath(config_name) # default class constructor input - - -class InvShapeDiscrim(ActiveSensingTask): - """ - Discrimination for shape task,with velocity threshold for trial initiation, no occluders. - The mouse must report the location of the black teardrop object and ignore the wide Pacman. - """ - - def __init__( - self, - teensy, - monitor=None, - write_video=False, - fps=60.0, - session_label=["ar_shape_discrimination_inv"], - epochs=[250], - epoch_labels=["pacman_vs_teardrop"], - config_file_path=config_path, - reward_size=100, - cropped_image=[0, 530, 0, 510], - unity_arena_size=[-9, 9, -10, -2], - r_report_box=[5, 10, -4, -2], - l_report_box=[-10, -5, -4, -2], - start_box=[-4, 4, -9, -5, 90], - rotate_camera=90.0, - prob_obj_on_left=0.5, - prob_block_coherence = 0.5, - mouse_report_delay=0.0, - slit_size=[4.0, 4.0, 1], - slit_depth=0.02, - target_selection=4.0, # changed 13. to 4. --> black teardrop - distractor_selection=9.0, # changed 6. to 9. --> black pacman - occlusion_type=0.0, - camera_type=1.0, - target_spread=4.0, - target_rotation=15.0, - target_size=2.0, - target_height=3.0, - block_length=1.0, - start_box_delay=0.25, - velocity_threshold=10.0, - distractor=1.0, - grey_screen_active=0.0, - target_distance=3, - use_dlc=True, - ): - super().__init__( - teensy=teensy, - monitor=monitor, - write_video=write_video, - fps=fps, - session_label=session_label, - epochs=epochs, - epoch_labels=epoch_labels, - config_file_path=config_file_path, - reward_size=reward_size, - cropped_image=cropped_image, - unity_arena_size=unity_arena_size, - r_report_box=r_report_box, - l_report_box=l_report_box, - start_box=start_box, - rotate_camera=rotate_camera, - prob_obj_on_left=prob_obj_on_left, - prob_block_coherence=prob_block_coherence, - mouse_report_delay=mouse_report_delay, - slit_size=slit_size, - slit_depth=slit_depth, - target_selection=target_selection, - distractor_selection=distractor_selection, - occlusion_type=occlusion_type, - camera_type=camera_type, - target_spread=target_spread, - target_rotation=target_rotation, - target_size=target_size, - target_height=target_height, - block_length=block_length, - start_box_delay=start_box_delay, - velocity_threshold=velocity_threshold, - distractor=distractor, - grey_screen_active=grey_screen_active, - target_distance=target_distance, - use_dlc=use_dlc, - ) diff --git a/mouse_task/inv_shape_mouse_discrim_multioccluders.py b/mouse_task/inv_shape_mouse_discrim_multioccluders.py deleted file mode 100644 index 2fb8415b2..000000000 --- a/mouse_task/inv_shape_mouse_discrim_multioccluders.py +++ /dev/null @@ -1,102 +0,0 @@ -""" -Discrimination task with multiple occluder widths. -Mouse must report the location of the black teardrop and ignore the pacman. -""" - -import os - -os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID" -os.environ["CUDA_VISIBLE_DEVICES"] = "0,1" - -import pathlib -import time as time - -from mouse_task.task_active_sensing import ActiveSensingTask - - -config_name = pathlib.Path("task_config.json") -current_dir = pathlib.Path(__file__).parent -config_path = current_dir.joinpath(config_name) # default class constructor input - - -class InvShapeDiscrimMultiOccluders(ActiveSensingTask): - """ - Shape discrimination task, with occluders of different sizes - The mouse must report the location of the black teardrop object and ignore the wide Pacman. - """ - - def __init__( - self, - teensy, - monitor=None, - write_video=False, - fps=60.0, - session_label=["ar_shape_discrim_multi_occluders_inv"], - epochs=[250], - epoch_labels=["pacman_vs_teardrop"], - config_file_path=config_path, - reward_size=100, - cropped_image=[0, 530, 0, 510], - unity_arena_size=[-9, 9, -10, -2], - r_report_box=[5, 10, -4, -2], - l_report_box=[-10, -5, -4, -2], - start_box=[-4, 4, -9, -5, 90], - rotate_camera=90.0, - prob_obj_on_left=0.5, - prob_block_coherence=0.5, - mouse_report_delay=0.0, - slit_size=[15.0, 10.78, 7.75, 5.57, 4.0], - slit_depth=0.02, - target_selection=4.0, # changed 13. to 4. --> black teardrop - distractor_selection=9.0, # changed 6. to 9. --> black pacman - occlusion_type=1.0, - camera_type=1.0, - target_spread=4.0, - target_rotation=15.0, - target_size=2.0, - target_height=3.0, - block_length=1.0, - start_box_delay=0.25, - velocity_threshold=10.0, - distractor=1.0, - grey_screen_active=0.0, - target_distance=3, - use_dlc=True, - ): - super().__init__( - teensy=teensy, - monitor=monitor, - write_video=write_video, - fps=fps, - session_label=session_label, - epochs=epochs, - epoch_labels=epoch_labels, - config_file_path=config_file_path, - reward_size=reward_size, - cropped_image=cropped_image, - unity_arena_size=unity_arena_size, - r_report_box=r_report_box, - l_report_box=l_report_box, - start_box=start_box, - rotate_camera=rotate_camera, - prob_obj_on_left=prob_obj_on_left, - prob_block_coherence=prob_block_coherence, - mouse_report_delay=mouse_report_delay, - slit_size=slit_size, - slit_depth=slit_depth, - target_selection=target_selection, - distractor_selection=distractor_selection, - occlusion_type=occlusion_type, - camera_type=camera_type, - target_spread=target_spread, - target_rotation=target_rotation, - target_size=target_size, - target_height=target_height, - block_length=block_length, - start_box_delay=start_box_delay, - velocity_threshold=velocity_threshold, - distractor=distractor, - grey_screen_active=grey_screen_active, - target_distance=target_distance, - use_dlc=use_dlc, - ) diff --git a/mouse_task/inv_shape_mouse_discrim_narrow_occluders.py b/mouse_task/inv_shape_mouse_discrim_narrow_occluders.py deleted file mode 100644 index b244e8067..000000000 --- a/mouse_task/inv_shape_mouse_discrim_narrow_occluders.py +++ /dev/null @@ -1,102 +0,0 @@ -""" -Discrimination task with 2 occluders of different sizes, narrower occluders. -Mouse must report the location of the black teardrop and ignore the pacman. -""" - -import os - -os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID" -os.environ["CUDA_VISIBLE_DEVICES"] = "0,1" - -import pathlib -import time as time - -from mouse_task.task_active_sensing import ActiveSensingTask - - -config_name = pathlib.Path("task_config.json") -current_dir = pathlib.Path(__file__).parent -config_path = current_dir.joinpath(config_name) # default class constructor input - - -class InvShapeDiscrimNarrowOccluders(ActiveSensingTask): - """ - Shape discrimination task, with occluders of different sizes - The mouse must report the location of the black teardrop object and ignore the wide Pacman. - """ - - def __init__( - self, - teensy, - monitor=None, - write_video=False, - fps=60.0, - session_label=["ar_shape_discrim_narrow_occluders_inv"], - epochs=[250], - epoch_labels=["pacman_vs_teardrop"], - config_file_path=config_path, - reward_size=100, - cropped_image=[0, 530, 0, 510], - unity_arena_size=[-9, 9, -10, -2], - r_report_box=[5, 10, -4, -2], - l_report_box=[-10, -5, -4, -2], - start_box=[-4, 4, -9, -5, 90], - rotate_camera=90.0, - prob_obj_on_left=0.5, - prob_block_coherence=0.5, - mouse_report_delay=0.0, - slit_size=[5.0, 15.0, 2], - slit_depth=0.02, - target_selection=4.0, # changed 13. to 4. --> black teardrop - distractor_selection=9.0, # changed 6. to 9. --> black pacman - occlusion_type=1.0, - camera_type=1.0, - target_spread=4.0, - target_rotation=15.0, - target_size=2.0, - target_height=3.0, - block_length=1.0, - start_box_delay=0.25, - velocity_threshold=10.0, - distractor=1.0, - grey_screen_active=0.0, - target_distance=3, - use_dlc=True, - ): - super().__init__( - teensy=teensy, - monitor=monitor, - write_video=write_video, - fps=fps, - session_label=session_label, - epochs=epochs, - epoch_labels=epoch_labels, - config_file_path=config_file_path, - reward_size=reward_size, - cropped_image=cropped_image, - unity_arena_size=unity_arena_size, - r_report_box=r_report_box, - l_report_box=l_report_box, - start_box=start_box, - rotate_camera=rotate_camera, - prob_obj_on_left=prob_obj_on_left, - prob_block_coherence=prob_block_coherence, - mouse_report_delay=mouse_report_delay, - slit_size=slit_size, - slit_depth=slit_depth, - target_selection=target_selection, - distractor_selection=distractor_selection, - occlusion_type=occlusion_type, - camera_type=camera_type, - target_spread=target_spread, - target_rotation=target_rotation, - target_size=target_size, - target_height=target_height, - block_length=block_length, - start_box_delay=start_box_delay, - velocity_threshold=velocity_threshold, - distractor=distractor, - grey_screen_active=grey_screen_active, - target_distance=target_distance, - use_dlc=use_dlc, - ) diff --git a/mouse_task/inv_shape_mouse_discrim_occluders.py b/mouse_task/inv_shape_mouse_discrim_occluders.py deleted file mode 100644 index 29f18757e..000000000 --- a/mouse_task/inv_shape_mouse_discrim_occluders.py +++ /dev/null @@ -1,102 +0,0 @@ -""" -Discrimination task with 2 occluders of different sizes, wider occluders. -Mouse must report the location of the black teardrop and ignore the pacman -""" - -import os - -os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID" -os.environ["CUDA_VISIBLE_DEVICES"] = "0,1" - -import pathlib -import time as time - -from mouse_task.task_active_sensing import ActiveSensingTask - - -config_name = pathlib.Path("task_config.json") -current_dir = pathlib.Path(__file__).parent -config_path = current_dir.joinpath(config_name) # default class constructor input - - -class InvShapeDiscrimOccluders(ActiveSensingTask): - """ - Shape discrimination task, with occluders of different sizes - The mouse must report the location of the black teardrop object and ignore the wide Pacman. - """ - - def __init__( - self, - teensy, - monitor=None, - write_video=False, - fps=60.0, - session_label=["ar_shape_discrim_occluders_inv"], - epochs=[250], - epoch_labels=["pacman_vs_teardrop"], - config_file_path=config_path, - reward_size=100, - cropped_image=[0, 530, 0, 510], - unity_arena_size=[-9, 9, -10, -2], - r_report_box=[5, 10, -4, -2], - l_report_box=[-10, -5, -4, -2], - start_box=[-4, 4, -9, -5, 90], - rotate_camera=90.0, - prob_obj_on_left=0.5, - prob_block_coherence=0.5, - mouse_report_delay=0.0, - slit_size=[8.0, 15.0, 2], - slit_depth=0.02, - target_selection=4.0, # changed 13. to 4. --> black teardrop - distractor_selection=9.0, # changed 6. to 9. --> black pacman - occlusion_type=1.0, - camera_type=1.0, - target_spread=4.0, - target_rotation=15.0, - target_size=2.0, - target_height=3.0, - block_length=1.0, - start_box_delay=0.25, - velocity_threshold=10.0, - distractor=1.0, - grey_screen_active=0.0, - target_distance=3, - use_dlc=True, - ): - super().__init__( - teensy=teensy, - monitor=monitor, - write_video=write_video, - fps=fps, - session_label=session_label, - epochs=epochs, - epoch_labels=epoch_labels, - config_file_path=config_file_path, - reward_size=reward_size, - cropped_image=cropped_image, - unity_arena_size=unity_arena_size, - r_report_box=r_report_box, - l_report_box=l_report_box, - start_box=start_box, - rotate_camera=rotate_camera, - prob_obj_on_left=prob_obj_on_left, - prob_block_coherence=prob_block_coherence, - mouse_report_delay=mouse_report_delay, - slit_size=slit_size, - slit_depth=slit_depth, - target_selection=target_selection, - distractor_selection=distractor_selection, - occlusion_type=occlusion_type, - camera_type=camera_type, - target_spread=target_spread, - target_rotation=target_rotation, - target_size=target_size, - target_height=target_height, - block_length=block_length, - start_box_delay=start_box_delay, - velocity_threshold=velocity_threshold, - distractor=distractor, - grey_screen_active=grey_screen_active, - target_distance=target_distance, - use_dlc=use_dlc, - ) diff --git a/mouse_task/mouse_detection_p1.py b/mouse_task/mouse_detection_p1.py deleted file mode 100644 index 94fa6c961..000000000 --- a/mouse_task/mouse_detection_p1.py +++ /dev/null @@ -1,100 +0,0 @@ -""" -Detection task without velocity threshold to initiate the trials. -""" - -import os - -os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID" -os.environ["CUDA_VISIBLE_DEVICES"] = "0,1" - -import pathlib -import time as time - -from mouse_task.task_active_sensing import ActiveSensingTask - - -config_name = pathlib.Path("task_config.json") -current_dir = pathlib.Path(__file__).parent -config_path = current_dir.joinpath(config_name) # default class constructor input - - -class DetectionWithoutVelocityThresholdTask(ActiveSensingTask): - """ - Detection task without velocity threshold to initiate the trials. - """ - - def __init__( - self, - teensy, - monitor=None, - write_video=False, - fps=60.0, - session_label=["ar_detection_no_velthr"], - epochs=[250], - epoch_labels=["single_teardrop"], - config_file_path=config_path, - reward_size=100, - cropped_image=[0, 530, 0, 510], - unity_arena_size=[-9, 9, -10, -2], - r_report_box=[5, 10, -4, -2], - l_report_box=[-10, -5, -4, -2], - start_box=[-4, 4, -9, -5, 90], - rotate_camera=90.0, - prob_obj_on_left=0.5, - prob_block_coherence = 0.5, - mouse_report_delay=0.0, - slit_size=[4.0, 4.0, 1], - slit_depth=0.2, - target_selection=6.0, - distractor_selection=4.0, - occlusion_type=0.0, - camera_type=1.0, - target_spread=4.0, - target_rotation=0, - target_size=2.0, - target_height=3.0, - block_length=1.0, - start_box_delay=0.1, - velocity_threshold=20.0, - distractor=0.0, - grey_screen_active=0.0, - target_distance=3, - use_dlc=True, - ): - super().__init__( - teensy=teensy, - monitor=monitor, - write_video=write_video, - fps=fps, - session_label=session_label, - epochs=epochs, - epoch_labels=epoch_labels, - config_file_path=config_file_path, - reward_size=reward_size, - cropped_image=cropped_image, - unity_arena_size=unity_arena_size, - r_report_box=r_report_box, - l_report_box=l_report_box, - start_box=start_box, - rotate_camera=rotate_camera, - prob_obj_on_left=prob_obj_on_left, - prob_block_coherence=prob_block_coherence, - mouse_report_delay=mouse_report_delay, - slit_size=slit_size, - slit_depth=slit_depth, - target_selection=target_selection, - distractor_selection=distractor_selection, - occlusion_type=occlusion_type, - camera_type=camera_type, - target_spread=target_spread, - target_rotation=target_rotation, - target_size=target_size, - target_height=target_height, - block_length=block_length, - start_box_delay=start_box_delay, - velocity_threshold=velocity_threshold, - distractor=distractor, - grey_screen_active=grey_screen_active, - target_distance=target_distance, - use_dlc=use_dlc, - ) diff --git a/mouse_task/mouse_detection_p2.py b/mouse_task/mouse_detection_p2.py deleted file mode 100644 index 30cd5fc48..000000000 --- a/mouse_task/mouse_detection_p2.py +++ /dev/null @@ -1,102 +0,0 @@ -""" -Detection task with velocity threshold to initiate the trials. -""" - -import os - -os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID" -os.environ["CUDA_VISIBLE_DEVICES"] = "0,1" - -import pathlib -import time as time - -from mouse_task.task_active_sensing import ActiveSensingTask - - -config_name = pathlib.Path("task_config.json") -current_dir = pathlib.Path(__file__).parent -config_path = current_dir.joinpath(config_name) # default class constructor input - - - -class DetectionWithVelocityThresholdTask(ActiveSensingTask): - - """ - Detection task with velocity threshold to initiate the trials. - """ - - def __init__( - self, - teensy, - monitor=None, - write_video=False, - fps=60.0, - session_label=["ar_detection_velthr"], - epochs=[250], - epoch_labels=["single_teardrop"], - config_file_path=config_path, - reward_size=100, - cropped_image=[0, 530, 0, 510], - unity_arena_size=[-9, 9, -10, -2], - r_report_box=[5, 10, -4, -2], - l_report_box=[-10, -5, -4, -2], - start_box=[-4, 4, -9, -5, 90], - rotate_camera=90.0, - prob_obj_on_left=0.5, - prob_block_coherence = 0.5, - mouse_report_delay=0.0, - slit_size=[19.0, 20.0, 2], - slit_depth=0.2, - target_selection=6.0, - distractor_selection=4.0, - occlusion_type=0.0, - camera_type=1.0, - target_spread=4.0, - target_rotation=0, - target_size=2.0, - target_height=3.0, - block_length=1.0, - start_box_delay=0.25, - velocity_threshold=10.0, - distractor=0.0, - grey_screen_active=0.0, - target_distance=3, - use_dlc=True, - ): - super().__init__( - teensy=teensy, - monitor=monitor, - write_video=write_video, - fps=fps, - session_label=session_label, - epochs=epochs, - epoch_labels=epoch_labels, - config_file_path=config_file_path, - reward_size=reward_size, - cropped_image=cropped_image, - unity_arena_size=unity_arena_size, - r_report_box=r_report_box, - l_report_box=l_report_box, - start_box=start_box, - rotate_camera=rotate_camera, - prob_obj_on_left=prob_obj_on_left, - prob_block_coherence=prob_block_coherence, - mouse_report_delay=mouse_report_delay, - slit_size=slit_size, - slit_depth=slit_depth, - target_selection=target_selection, - distractor_selection=distractor_selection, - occlusion_type=occlusion_type, - camera_type=camera_type, - target_spread=target_spread, - target_rotation=target_rotation, - target_size=target_size, - target_height=target_height, - block_length=block_length, - start_box_delay=start_box_delay, - velocity_threshold=velocity_threshold, - distractor=distractor, - grey_screen_active=grey_screen_active, - target_distance=target_distance, - use_dlc=use_dlc, - ) diff --git a/mouse_task/mouse_discrim.py b/mouse_task/mouse_discrim.py deleted file mode 100644 index 074c2f1f9..000000000 --- a/mouse_task/mouse_discrim.py +++ /dev/null @@ -1,100 +0,0 @@ -""" -Discrimination task (without occluders). -""" - -import os - -os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID" -os.environ["CUDA_VISIBLE_DEVICES"] = "0,1" - -import pathlib -import time as time - -from mouse_task.task_active_sensing import ActiveSensingTask - - -config_name = pathlib.Path("task_config.json") -current_dir = pathlib.Path(__file__).parent -config_path = current_dir.joinpath(config_name) # default class constructor input - - -class DiscriminationTask(ActiveSensingTask): - """ - Discrimination task (without occluders). - """ - - def __init__( - self, - teensy, - monitor=None, - write_video=False, - fps=60.0, - session_label=["ar_discrim"], - epochs=[250], - epoch_labels=["dual_teardrop"], - config_file_path=config_path, - reward_size=100, - cropped_image=[0, 530, 0, 510], - unity_arena_size=[-9, 9, -10, -2], - r_report_box=[5, 10, -4, -2], - l_report_box=[-10, -5, -4, -2], - start_box=[-4, 4, -9, -5, 90], - rotate_camera=90.0, - prob_obj_on_left=0.5, - prob_block_coherence = 0.5, - mouse_report_delay=0.0, - slit_size=[4.0, 10.0, 2], - slit_depth=0.2, - target_selection=6.0, - distractor_selection=4.0, - occlusion_type=0.0, - camera_type=1.0, - target_spread=4.0, - target_rotation=0, - target_size=2.0, - target_height=3.0, - block_length=1.0, - start_box_delay=0.25, - velocity_threshold=10.0, - distractor=1.0, - grey_screen_active=0.0, - target_distance=3, - use_dlc=True, - ): - super().__init__( - teensy=teensy, - monitor=monitor, - write_video=write_video, - fps=fps, - session_label=session_label, - epochs=epochs, - epoch_labels=epoch_labels, - config_file_path=config_file_path, - reward_size=reward_size, - cropped_image=cropped_image, - unity_arena_size=unity_arena_size, - r_report_box=r_report_box, - l_report_box=l_report_box, - start_box=start_box, - rotate_camera=rotate_camera, - prob_obj_on_left=prob_obj_on_left, - prob_block_coherence=prob_block_coherence, - mouse_report_delay=mouse_report_delay, - slit_size=slit_size, - slit_depth=slit_depth, - target_selection=target_selection, - distractor_selection=distractor_selection, - occlusion_type=occlusion_type, - camera_type=camera_type, - target_spread=target_spread, - target_rotation=target_rotation, - target_size=target_size, - target_height=target_height, - block_length=block_length, - start_box_delay=start_box_delay, - velocity_threshold=velocity_threshold, - distractor=distractor, - grey_screen_active=grey_screen_active, - target_distance=target_distance, - use_dlc=use_dlc, - ) diff --git a/mouse_task/mouse_discrim_multioccluders.py b/mouse_task/mouse_discrim_multioccluders.py deleted file mode 100644 index 99a8c2099..000000000 --- a/mouse_task/mouse_discrim_multioccluders.py +++ /dev/null @@ -1,102 +0,0 @@ -""" -Discrimination task with occluders on. -""" - -import os - -os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID" -os.environ["CUDA_VISIBLE_DEVICES"] = "0,1" - -import pathlib -import time as time - -from mouse_task.task_active_sensing import ActiveSensingTask - - -config_name = pathlib.Path("task_config.json") -current_dir = pathlib.Path(__file__).parent -config_path = current_dir.joinpath(config_name) # default class constructor input - - - -class DiscriminationWithMultiOccludersTask(ActiveSensingTask): - - """ - Discrimination task with multiple occluders, sampled from a log space these are [12.0, 8.48, 6.0, 4.24, 3]. - """ - - def __init__( - self, - teensy, - monitor=None, - write_video=False, - fps=60.0, - session_label=["ar_discrim_5_occluders"], - epochs=[250], - epoch_labels=["dual_teardrop"], - config_file_path=config_path, - reward_size=100, - cropped_image=[0, 530, 0, 510], - unity_arena_size=[-9, 9, -10, -2], - r_report_box=[5, 10, -4, -2], - l_report_box=[-10, -5, -4, -2], - start_box=[-4, 4, -9, -5, 90], - rotate_camera=90.0, - prob_obj_on_left=0.5, - prob_block_coherence =0.5, - mouse_report_delay=0.0, - slit_size=[12.0, 8.48, 6., 4.2, 3], - slit_depth=0.2, - target_selection=6.0, - distractor_selection=4.0, - occlusion_type=1.0, - camera_type=1.0, - target_spread=4.0, - target_rotation=0, - target_size=2.0, - target_height=3.0, - block_length=1.0, - start_box_delay=0.25, - velocity_threshold=10.0, - distractor=1.0, - grey_screen_active=0.0, - target_distance=3, - use_dlc=True, - ): - super().__init__( - teensy=teensy, - monitor=monitor, - write_video=write_video, - fps=fps, - session_label=session_label, - epochs=epochs, - epoch_labels=epoch_labels, - config_file_path=config_file_path, - reward_size=reward_size, - cropped_image=cropped_image, - unity_arena_size=unity_arena_size, - r_report_box=r_report_box, - l_report_box=l_report_box, - start_box=start_box, - rotate_camera=rotate_camera, - prob_obj_on_left=prob_obj_on_left, - prob_block_coherence=prob_block_coherence, - mouse_report_delay=mouse_report_delay, - slit_size=slit_size, - slit_depth=slit_depth, - target_selection=target_selection, - distractor_selection=distractor_selection, - occlusion_type=occlusion_type, - camera_type=camera_type, - target_spread=target_spread, - target_rotation=target_rotation, - target_size=target_size, - target_height=target_height, - block_length=block_length, - start_box_delay=start_box_delay, - velocity_threshold=velocity_threshold, - distractor=distractor, - grey_screen_active=grey_screen_active, - target_distance=target_distance, - use_dlc=use_dlc, - ) diff --git a/mouse_task/mouse_discrim_occluders.py b/mouse_task/mouse_discrim_occluders.py deleted file mode 100644 index 8a917a7b8..000000000 --- a/mouse_task/mouse_discrim_occluders.py +++ /dev/null @@ -1,100 +0,0 @@ -""" -Discrimination task with occluders on. -""" - -import os - -os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID" -os.environ["CUDA_VISIBLE_DEVICES"] = "0,1" - -import pathlib -import time as time - -from mouse_task.task_active_sensing import ActiveSensingTask - - -config_name = pathlib.Path("task_config.json") -current_dir = pathlib.Path(__file__).parent -config_path = current_dir.joinpath(config_name) # default class constructor input - - -class DiscriminationWithOccludersTask(ActiveSensingTask): - """ - Discrimination task with occluders on. - """ - - def __init__( - self, - teensy, - monitor=None, - write_video=False, - fps=60.0, - session_label=["ar_discrim_occluders"], - epochs=[250], - epoch_labels=["dual_teardrop"], - config_file_path=config_path, - reward_size=100, - cropped_image=[0, 530, 0, 510], - unity_arena_size=[-9, 9, -10, -2], - r_report_box=[5, 10, -4, -2], - l_report_box=[-10, -5, -4, -2], - start_box=[-4, 4, -9, -5, 90], - rotate_camera=90.0, - prob_obj_on_left=0.5, - prob_block_coherence=0.5, - mouse_report_delay=0.0, - slit_size=[4.3, 12.0, 2], - slit_depth=0.2, - target_selection=6.0, - distractor_selection=4.0, - occlusion_type=1.0, - camera_type=1.0, - target_spread=4.0, - target_rotation=0, - target_size=2.0, - target_height=3.0, - block_length=1.0, - start_box_delay=0.25, - velocity_threshold=10.0, - distractor=1.0, - grey_screen_active=0.0, - target_distance=3, - use_dlc=True, - ): - super().__init__( - teensy=teensy, - monitor=monitor, - write_video=write_video, - fps=fps, - session_label=session_label, - epochs=epochs, - epoch_labels=epoch_labels, - config_file_path=config_file_path, - reward_size=reward_size, - cropped_image=cropped_image, - unity_arena_size=unity_arena_size, - r_report_box=r_report_box, - l_report_box=l_report_box, - start_box=start_box, - rotate_camera=rotate_camera, - prob_obj_on_left=prob_obj_on_left, - prob_block_coherence=prob_block_coherence, - mouse_report_delay=mouse_report_delay, - slit_size=slit_size, - slit_depth=slit_depth, - target_selection=target_selection, - distractor_selection=distractor_selection, - occlusion_type=occlusion_type, - camera_type=camera_type, - target_spread=target_spread, - target_rotation=target_rotation, - target_size=target_size, - target_height=target_height, - block_length=block_length, - start_box_delay=start_box_delay, - velocity_threshold=velocity_threshold, - distractor=distractor, - grey_screen_active=grey_screen_active, - target_distance=target_distance, - use_dlc=use_dlc, - ) diff --git a/mouse_task/shape_mouse_detection_p1.py b/mouse_task/shape_mouse_detection_p1.py deleted file mode 100644 index 5a69471e3..000000000 --- a/mouse_task/shape_mouse_detection_p1.py +++ /dev/null @@ -1,100 +0,0 @@ -""" -Detection task without velocity threshold to initiate the trials. -""" - -import os - -os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID" -os.environ["CUDA_VISIBLE_DEVICES"] = "0,1" - -import pathlib -import time as time - -from mouse_task.task_active_sensing import ActiveSensingTask - - -config_name = pathlib.Path("task_config.json") -current_dir = pathlib.Path(__file__).parent -config_path = current_dir.joinpath(config_name) # default class constructor input - - -class ShapeDetectionWithoutVelocityThresholdTask(ActiveSensingTask): - """ - Detection of shape task without velocity threshold to initiate the trials. - """ - - def __init__( - self, - teensy, - monitor=None, - write_video=False, - fps=60.0, - session_label=["ar_shape_detection_no_velthr"], - epochs=[250], - epoch_labels=["single_wide_pacman"], - config_file_path=config_path, - reward_size=100, - cropped_image=[0, 530, 0, 510], - unity_arena_size=[-9, 9, -10, -2], - r_report_box=[5, 10, -4, -2], - l_report_box=[-10, -5, -4, -2], - start_box=[-4, 4, -9, -5, 90], - rotate_camera=90.0, - prob_obj_on_left=0.5, - prob_block_coherence = 0.5, - mouse_report_delay=0.0, - slit_size=[4.0, 4.0, 1], - slit_depth=0.02, - target_selection=13.0, - distractor_selection=6.0, - occlusion_type=0.0, - camera_type=1.0, - target_spread=3.0, - target_rotation=15.0, - target_size=2.0, - target_height=3.0, - block_length=1.0, - start_box_delay=0.1, - velocity_threshold=10.0, - distractor=0.0, - grey_screen_active=0.0, - target_distance=4.0, - use_dlc=True, - ): - super().__init__( - teensy=teensy, - monitor=monitor, - write_video=write_video, - fps=fps, - session_label=session_label, - epochs=epochs, - epoch_labels=epoch_labels, - config_file_path=config_file_path, - reward_size=reward_size, - cropped_image=cropped_image, - unity_arena_size=unity_arena_size, - r_report_box=r_report_box, - l_report_box=l_report_box, - start_box=start_box, - rotate_camera=rotate_camera, - prob_obj_on_left=prob_obj_on_left, - prob_block_coherence=prob_block_coherence, - mouse_report_delay=mouse_report_delay, - slit_size=slit_size, - slit_depth=slit_depth, - target_selection=target_selection, - distractor_selection=distractor_selection, - occlusion_type=occlusion_type, - camera_type=camera_type, - target_spread=target_spread, - target_rotation=target_rotation, - target_size=target_size, - target_height=target_height, - block_length=block_length, - start_box_delay=start_box_delay, - velocity_threshold=velocity_threshold, - distractor=distractor, - grey_screen_active=grey_screen_active, - target_distance=target_distance, - use_dlc=use_dlc, - ) diff --git a/mouse_task/shape_mouse_detection_p2.py b/mouse_task/shape_mouse_detection_p2.py deleted file mode 100644 index 00b5636e3..000000000 --- a/mouse_task/shape_mouse_detection_p2.py +++ /dev/null @@ -1,100 +0,0 @@ -""" -Detection task without velocity threshold to initiate the trials. -""" - -import os - -os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID" -os.environ["CUDA_VISIBLE_DEVICES"] = "0,1" - -import pathlib -import time as time - -from mouse_task.task_active_sensing import ActiveSensingTask - - -config_name = pathlib.Path("task_config.json") -current_dir = pathlib.Path(__file__).parent -config_path = current_dir.joinpath(config_name) # default class constructor input - - -class ShapeDetectionWithVelocityThresholdTask(ActiveSensingTask): - """ - Detection of shape task with a lower velocity threshold to initiate the trials. - """ - - def __init__( - self, - teensy, - monitor=None, - write_video=False, - fps=60.0, - session_label=["ar_shape_detection_velthr"], - epochs=[250], - epoch_labels=["single_wide_pacman"], - config_file_path=config_path, - reward_size=100, - cropped_image=[0, 530, 0, 510], - unity_arena_size=[-9, 9, -10, -2], - r_report_box=[5, 10, -4, -2], - l_report_box=[-10, -5, -4, -2], - start_box=[-4, 4, -9, -5, 90], - rotate_camera=90.0, - prob_obj_on_left=0.5, - prob_block_coherence = 0.5, - mouse_report_delay=0.0, - slit_size=[4.0, 4.0, 1], - slit_depth=0.02, - target_selection=13.0, - distractor_selection=6.0, - occlusion_type=0.0, - camera_type=1.0, - target_spread=3.0, - target_rotation=15.0, - target_size=2.0, - target_height=3.0, - block_length=1.0, - start_box_delay=0.25, - velocity_threshold=5.0, - distractor=0.0, - grey_screen_active=0.0, - target_distance=4.0, - use_dlc=True, - ): - super().__init__( - teensy=teensy, - monitor=monitor, - write_video=write_video, - fps=fps, - session_label=session_label, - epochs=epochs, - epoch_labels=epoch_labels, - config_file_path=config_file_path, - reward_size=reward_size, - cropped_image=cropped_image, - unity_arena_size=unity_arena_size, - r_report_box=r_report_box, - l_report_box=l_report_box, - start_box=start_box, - rotate_camera=rotate_camera, - prob_obj_on_left=prob_obj_on_left, - prob_block_coherence=prob_block_coherence, - mouse_report_delay=mouse_report_delay, - slit_size=slit_size, - slit_depth=slit_depth, - target_selection=target_selection, - distractor_selection=distractor_selection, - occlusion_type=occlusion_type, - camera_type=camera_type, - target_spread=target_spread, - target_rotation=target_rotation, - target_size=target_size, - target_height=target_height, - block_length=block_length, - start_box_delay=start_box_delay, - velocity_threshold=velocity_threshold, - distractor=distractor, - grey_screen_active=grey_screen_active, - target_distance=target_distance, - use_dlc=use_dlc, - ) diff --git a/mouse_task/shape_mouse_discrim.py b/mouse_task/shape_mouse_discrim.py deleted file mode 100644 index ff205a61f..000000000 --- a/mouse_task/shape_mouse_discrim.py +++ /dev/null @@ -1,102 +0,0 @@ -""" -Discrimination task without occluders. -Mouse must report the location of the white pacman and ignore the teardrop -""" - -import os - -os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID" -os.environ["CUDA_VISIBLE_DEVICES"] = "0,1" - -import pathlib -import time as time - -from mouse_task.task_active_sensing import ActiveSensingTask - - -config_name = pathlib.Path("task_config.json") -current_dir = pathlib.Path(__file__).parent -config_path = current_dir.joinpath(config_name) # default class constructor input - - -class ShapeDiscrim(ActiveSensingTask): - """ - Discrimination for shape task,with velocity threshold for trial initiation, no occluders. - The mouse must report the location of the white Pacman object and ignore the teardrop. - """ - - def __init__( - self, - teensy, - monitor=None, - write_video=False, - fps=60.0, - session_label=["ar_shape_discrimination"], - epochs=[250], - epoch_labels=["pacman_vs_teardrop"], - config_file_path=config_path, - reward_size=100, - cropped_image=[0, 530, 0, 510], - unity_arena_size=[-9, 9, -10, -2], - r_report_box=[5, 10, -4, -2], - l_report_box=[-10, -5, -4, -2], - start_box=[-4, 4, -9, -5, 90], - rotate_camera=90.0, - prob_obj_on_left=0.5, - prob_block_coherence = 0.5, - mouse_report_delay=0.0, - slit_size=[4.0, 4.0, 1], - slit_depth=0.02, - target_selection=13.0, - distractor_selection=6.0, - occlusion_type=0.0, - camera_type=1.0, - target_spread=3.0, - target_rotation=15.0, - target_size=2.0, - target_height=3.0, - block_length=1.0, - start_box_delay=0.25, - velocity_threshold=5.0, - distractor=1.0, - grey_screen_active=0.0, - target_distance=4.0, - use_dlc=True, - ): - super().__init__( - teensy=teensy, - monitor=monitor, - write_video=write_video, - fps=fps, - session_label=session_label, - epochs=epochs, - epoch_labels=epoch_labels, - config_file_path=config_file_path, - reward_size=reward_size, - cropped_image=cropped_image, - unity_arena_size=unity_arena_size, - r_report_box=r_report_box, - l_report_box=l_report_box, - start_box=start_box, - rotate_camera=rotate_camera, - prob_obj_on_left=prob_obj_on_left, - prob_block_coherence=prob_block_coherence, - mouse_report_delay=mouse_report_delay, - slit_size=slit_size, - slit_depth=slit_depth, - target_selection=target_selection, - distractor_selection=distractor_selection, - occlusion_type=occlusion_type, - camera_type=camera_type, - target_spread=target_spread, - target_rotation=target_rotation, - target_size=target_size, - target_height=target_height, - block_length=block_length, - start_box_delay=start_box_delay, - velocity_threshold=velocity_threshold, - distractor=distractor, - grey_screen_active=grey_screen_active, - target_distance=target_distance, - use_dlc=use_dlc, - ) diff --git a/mouse_task/shape_mouse_discrim_multioccluders.py b/mouse_task/shape_mouse_discrim_multioccluders.py deleted file mode 100644 index 7dc8a5841..000000000 --- a/mouse_task/shape_mouse_discrim_multioccluders.py +++ /dev/null @@ -1,101 +0,0 @@ -""" -Detection task without velocity threshold to initiate the trials. -""" - -import os - -os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID" -os.environ["CUDA_VISIBLE_DEVICES"] = "0,1" - -import pathlib -import time as time - -from mouse_task.task_active_sensing import ActiveSensingTask - - -config_name = pathlib.Path("task_config.json") -current_dir = pathlib.Path(__file__).parent -config_path = current_dir.joinpath(config_name) # default class constructor input - - -class ShapeDiscrimMultiOccluders(ActiveSensingTask): - """ - Shape discrimination task, with occluders of different sizes - Mouse must report the white pacman location and ignore the teardrop. - """ - - def __init__( - self, - teensy, - monitor=None, - write_video=False, - fps=60.0, - session_label=["ar_shape_discrim_multi_occluders"], - epochs=[250], - epoch_labels=["pacman_vs_teardrop"], - config_file_path=config_path, - reward_size=100, - cropped_image=[0, 530, 0, 510], - unity_arena_size=[-9, 9, -10, -2], - r_report_box=[5, 10, -4, -2], - l_report_box=[-10, -5, -4, -2], - start_box=[-4, 4, -9, -5, 90], - rotate_camera=90.0, - prob_obj_on_left=0.5, - prob_block_coherence=0.5, - mouse_report_delay=0.0, - slit_size=[15.0, 10.78, 7.75, 5.57, 4.0], - slit_depth=0.02, - target_selection=13.0, - distractor_selection=6.0, - occlusion_type=1.0, - camera_type=1.0, - target_spread=3.0, - target_rotation=15.0, - target_size=2.0, - target_height=3.0, - block_length=1.0, - start_box_delay=0.25, - velocity_threshold=5.0, - distractor=1.0, - grey_screen_active=0.0, - target_distance=4.0, - use_dlc=True, - ): - super().__init__( - teensy=teensy, - monitor=monitor, - write_video=write_video, - fps=fps, - session_label=session_label, - epochs=epochs, - epoch_labels=epoch_labels, - config_file_path=config_file_path, - reward_size=reward_size, - cropped_image=cropped_image, - unity_arena_size=unity_arena_size, - r_report_box=r_report_box, - l_report_box=l_report_box, - start_box=start_box, - rotate_camera=rotate_camera, - prob_obj_on_left=prob_obj_on_left, - prob_block_coherence=prob_block_coherence, - mouse_report_delay=mouse_report_delay, - slit_size=slit_size, - slit_depth=slit_depth, - target_selection=target_selection, - distractor_selection=distractor_selection, - occlusion_type=occlusion_type, - camera_type=camera_type, - target_spread=target_spread, - target_rotation=target_rotation, - target_size=target_size, - target_height=target_height, - block_length=block_length, - start_box_delay=start_box_delay, - velocity_threshold=velocity_threshold, - distractor=distractor, - grey_screen_active=grey_screen_active, - target_distance=target_distance, - use_dlc=use_dlc, - ) diff --git a/mouse_task/shape_mouse_discrim_narrow_occluders.py b/mouse_task/shape_mouse_discrim_narrow_occluders.py deleted file mode 100644 index c8a155cbf..000000000 --- a/mouse_task/shape_mouse_discrim_narrow_occluders.py +++ /dev/null @@ -1,101 +0,0 @@ -""" -Detection task without velocity threshold to initiate the trials. -""" - -import os - -os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID" -os.environ["CUDA_VISIBLE_DEVICES"] = "0,1" - -import pathlib -import time as time - -from mouse_task.task_active_sensing import ActiveSensingTask - - -config_name = pathlib.Path("task_config.json") -current_dir = pathlib.Path(__file__).parent -config_path = current_dir.joinpath(config_name) # default class constructor input - - -class ShapeDiscrimNarrowOccluders(ActiveSensingTask): - """ - Shape discrimination task, with occluders of different sizes - Mouse must report the white pacman location and ignore the teardrop. - """ - - def __init__( - self, - teensy, - monitor=None, - write_video=False, - fps=60.0, - session_label=["ar_shape_discrim_narrow_occluders"], - epochs=[250], - epoch_labels=["pacman_vs_teardrop"], - config_file_path=config_path, - reward_size=100, - cropped_image=[0, 530, 0, 510], - unity_arena_size=[-9, 9, -10, -2], - r_report_box=[5, 10, -4, -2], - l_report_box=[-10, -5, -4, -2], - start_box=[-4, 4, -9, -5, 90], - rotate_camera=90.0, - prob_obj_on_left=0.5, - prob_block_coherence=0.5, - mouse_report_delay=0.0, - slit_size=[5.0, 15.0, 2], - slit_depth=0.02, - target_selection=13.0, - distractor_selection=6.0, - occlusion_type=1.0, - camera_type=1.0, - target_spread=3.0, - target_rotation=15.0, - target_size=2.0, - target_height=3.0, - block_length=1.0, - start_box_delay=0.25, - velocity_threshold=5.0, - distractor=1.0, - grey_screen_active=0.0, - target_distance=4.0, - use_dlc=True, - ): - super().__init__( - teensy=teensy, - monitor=monitor, - write_video=write_video, - fps=fps, - session_label=session_label, - epochs=epochs, - epoch_labels=epoch_labels, - config_file_path=config_file_path, - reward_size=reward_size, - cropped_image=cropped_image, - unity_arena_size=unity_arena_size, - r_report_box=r_report_box, - l_report_box=l_report_box, - start_box=start_box, - rotate_camera=rotate_camera, - prob_obj_on_left=prob_obj_on_left, - prob_block_coherence=prob_block_coherence, - mouse_report_delay=mouse_report_delay, - slit_size=slit_size, - slit_depth=slit_depth, - target_selection=target_selection, - distractor_selection=distractor_selection, - occlusion_type=occlusion_type, - camera_type=camera_type, - target_spread=target_spread, - target_rotation=target_rotation, - target_size=target_size, - target_height=target_height, - block_length=block_length, - start_box_delay=start_box_delay, - velocity_threshold=velocity_threshold, - distractor=distractor, - grey_screen_active=grey_screen_active, - target_distance=target_distance, - use_dlc=use_dlc, - ) diff --git a/mouse_task/shape_mouse_discrim_occluders.py b/mouse_task/shape_mouse_discrim_occluders.py deleted file mode 100644 index 3e69db701..000000000 --- a/mouse_task/shape_mouse_discrim_occluders.py +++ /dev/null @@ -1,101 +0,0 @@ -""" -Detection task without velocity threshold to initiate the trials. -""" - -import os - -os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID" -os.environ["CUDA_VISIBLE_DEVICES"] = "0,1" - -import pathlib -import time as time - -from mouse_task.task_active_sensing import ActiveSensingTask - - -config_name = pathlib.Path("task_config.json") -current_dir = pathlib.Path(__file__).parent -config_path = current_dir.joinpath(config_name) # default class constructor input - - -class ShapeDiscrimOccluders(ActiveSensingTask): - """ - Shape discrimination task, with occluders of different sizes - Mouse must report the white pacman location and ignore the teardrop. - """ - - def __init__( - self, - teensy, - monitor=None, - write_video=False, - fps=60.0, - session_label=["ar_shape_discrim_occluders"], - epochs=[250], - epoch_labels=["pacman_vs_teardrop"], - config_file_path=config_path, - reward_size=100, - cropped_image=[0, 530, 0, 510], - unity_arena_size=[-9, 9, -10, -2], - r_report_box=[5, 10, -4, -2], - l_report_box=[-10, -5, -4, -2], - start_box=[-4, 4, -9, -5, 90], - rotate_camera=90.0, - prob_obj_on_left=0.5, - prob_block_coherence=0.5, - mouse_report_delay=0.0, - slit_size=[8.0, 15.0, 2], - slit_depth=0.02, - target_selection=13.0, - distractor_selection=6.0, - occlusion_type=1.0, - camera_type=1.0, - target_spread=3.0, - target_rotation=15.0, - target_size=2.0, - target_height=3.0, - block_length=1.0, - start_box_delay=0.25, - velocity_threshold=5.0, - distractor=1.0, - grey_screen_active=0.0, - target_distance=3.5, - use_dlc=True, - ): - super().__init__( - teensy=teensy, - monitor=monitor, - write_video=write_video, - fps=fps, - session_label=session_label, - epochs=epochs, - epoch_labels=epoch_labels, - config_file_path=config_file_path, - reward_size=reward_size, - cropped_image=cropped_image, - unity_arena_size=unity_arena_size, - r_report_box=r_report_box, - l_report_box=l_report_box, - start_box=start_box, - rotate_camera=rotate_camera, - prob_obj_on_left=prob_obj_on_left, - prob_block_coherence=prob_block_coherence, - mouse_report_delay=mouse_report_delay, - slit_size=slit_size, - slit_depth=slit_depth, - target_selection=target_selection, - distractor_selection=distractor_selection, - occlusion_type=occlusion_type, - camera_type=camera_type, - target_spread=target_spread, - target_rotation=target_rotation, - target_size=target_size, - target_height=target_height, - block_length=block_length, - start_box_delay=start_box_delay, - velocity_threshold=velocity_threshold, - distractor=distractor, - grey_screen_active=grey_screen_active, - target_distance=target_distance, - use_dlc=use_dlc, - ) diff --git a/mouse_task/task_active_sensing.py b/mouse_task/task_active_sensing.py index c703733c3..0e2103d78 100644 --- a/mouse_task/task_active_sensing.py +++ b/mouse_task/task_active_sensing.py @@ -15,7 +15,7 @@ import numpy as np import time -from mouse_task.helpers import process_config +from mouse_task.helpers import load_task_config, process_config from teensyexp.tasks_abc.unity_task import UnityTask from teensyexp.tasks_abc.dlc_deque_socket import DLCClient @@ -24,6 +24,12 @@ from typing import List, Optional +# Machine-specific config holding the absolute Unity build path. +_DEFAULT_CONFIG_FILE = pathlib.Path(__file__).parent / "task_config.json" + +# Marks a param not explicitly passed, so it is resolved from task_config. +_UNSET = object() + class ActiveSensingTask(UnityTask): """ @@ -73,42 +79,90 @@ class ActiveSensingTask(UnityTask): def __init__( self, teensy: Teensy, - session_label: List[str], - config_file_path: pathlib.Path, - monitor: Optional[bool], - write_video: bool, - fps: float, - epochs: int, - epoch_labels: List[str], - reward_size: int, - cropped_image: List[int], - unity_arena_size: List[int], - r_report_box: List[int], - l_report_box: List[int], - start_box: List[int], - rotate_camera: float, - prob_obj_on_left: float, - prob_block_coherence: float, - mouse_report_delay: float, - slit_size: List[int], - slit_depth: float, - target_selection: float, - distractor_selection: float, - occlusion_type: float, - camera_type: float, - target_spread: float, - target_rotation: float, - target_size: float, - target_height: float, - block_length: float, - start_box_delay: float, - velocity_threshold: float, - distractor: float, - grey_screen_active: float, - target_distance: float, - use_dlc: bool, + task_config: Optional[str] = None, + config_file_path: pathlib.Path = _DEFAULT_CONFIG_FILE, + session_label: List[str] = _UNSET, + monitor: Optional[bool] = _UNSET, + write_video: bool = _UNSET, + fps: float = _UNSET, + epochs: int = _UNSET, + epoch_labels: List[str] = _UNSET, + reward_size: int = _UNSET, + cropped_image: List[int] = _UNSET, + unity_arena_size: List[int] = _UNSET, + r_report_box: List[int] = _UNSET, + l_report_box: List[int] = _UNSET, + start_box: List[int] = _UNSET, + rotate_camera: float = _UNSET, + prob_obj_on_left: float = _UNSET, + prob_block_coherence: float = _UNSET, + mouse_report_delay: float = _UNSET, + slit_size: List[int] = _UNSET, + slit_depth: float = _UNSET, + target_selection: float = _UNSET, + distractor_selection: float = _UNSET, + occlusion_type: float = _UNSET, + camera_type: float = _UNSET, + target_spread: float = _UNSET, + target_rotation: float = _UNSET, + target_size: float = _UNSET, + target_height: float = _UNSET, + block_length: float = _UNSET, + start_box_delay: float = _UNSET, + velocity_threshold: float = _UNSET, + distractor: float = _UNSET, + grey_screen_active: float = _UNSET, + target_distance: float = _UNSET, + use_dlc: bool = _UNSET, ): + # Resolve params: explicit args win, else fall back to task_config. + _explicit = {k: v for k, v in locals().items() if v is not _UNSET} + if task_config is not None: + merged = load_task_config(task_config) + merged.update({k: v for k, v in _explicit.items() if k in merged}) + else: + merged = _explicit + _missing = [k for k, v in locals().items() if v is _UNSET and k not in merged] + if _missing: + raise ValueError( + "Missing task parameters (pass them explicitly or via task_config): " + + ", ".join(sorted(_missing)) + ) + session_label = merged["session_label"] + monitor = merged["monitor"] + write_video = merged["write_video"] + fps = merged["fps"] + epochs = merged["epochs"] + epoch_labels = merged["epoch_labels"] + reward_size = merged["reward_size"] + cropped_image = merged["cropped_image"] + unity_arena_size = merged["unity_arena_size"] + r_report_box = merged["r_report_box"] + l_report_box = merged["l_report_box"] + start_box = merged["start_box"] + rotate_camera = merged["rotate_camera"] + prob_obj_on_left = merged["prob_obj_on_left"] + prob_block_coherence = merged["prob_block_coherence"] + mouse_report_delay = merged["mouse_report_delay"] + slit_size = merged["slit_size"] + slit_depth = merged["slit_depth"] + target_selection = merged["target_selection"] + distractor_selection = merged["distractor_selection"] + occlusion_type = merged["occlusion_type"] + camera_type = merged["camera_type"] + target_spread = merged["target_spread"] + target_rotation = merged["target_rotation"] + target_size = merged["target_size"] + target_height = merged["target_height"] + block_length = merged["block_length"] + start_box_delay = merged["start_box_delay"] + velocity_threshold = merged["velocity_threshold"] + distractor = merged["distractor"] + grey_screen_active = merged["grey_screen_active"] + target_distance = merged["target_distance"] + use_dlc = merged["use_dlc"] + # Initialized in init_DLC_live() self.t_count = None self.filt = None diff --git a/mouse_task/tests/test_task_registry.py b/mouse_task/tests/test_task_registry.py new file mode 100644 index 000000000..00149629c --- /dev/null +++ b/mouse_task/tests/test_task_registry.py @@ -0,0 +1,222 @@ +# Description: Unit tests for the modular task-config / registry setup. +# +# These verify that the per-task YAML configs (configs/common.yaml + +# configs/tasks/.yaml) and the dynamically generated ActiveSensingTask +# subclasses (mouse_task._registry) behave as expected: +# * every task YAML produces a correctly named, importable subclass; +# * load_task_config deep-merges common defaults with task overrides; +# * the generated __init__ keeps a real, introspectable signature (required +# by the teensyexp GUI, which reads args/defaults via inspect.getargspec); +# * constructing a task resolves parameters from the merged config, with +# explicitly-passed arguments taking precedence. + +import inspect +import types +import unittest +from pathlib import Path +from unittest.mock import MagicMock, patch + +import yaml + +import mouse_task +from mouse_task.helpers import CONFIG_DIR, load_task_config +from mouse_task.task_active_sensing import ActiveSensingTask + +_TASKS_DIR = CONFIG_DIR / "tasks" +_MOCK_CONFIG = {"ar_env_unity_absolute_path": "mock_path"} + + +def _task_names(): + """All task-config stems (one per configs/tasks/*.yaml).""" + return sorted(p.stem for p in _TASKS_DIR.glob("*.yaml")) + + +def _yaml(path): + with open(path) as f: + return yaml.safe_load(f) or {} + + +def _signature_defaults(cls): + """Map of {arg_name: default} for a generated class __init__.""" + spec = inspect.getfullargspec(cls.__init__) + return dict(zip(spec.args[-len(spec.defaults):], spec.defaults)) + + +class TestTaskRegistry(unittest.TestCase): + """The registry exposes one class per task YAML.""" + + def test_at_least_expected_tasks(self): + names = _task_names() + self.assertGreaterEqual(len(names), 20) + # A few well-known variants must be present. + for stem in ["mouse_discrim", "shape_mouse_discrim", "shape_mouse_discrim_occluders"]: + self.assertIn(stem, names) + + def test_every_yaml_yields_an_exposed_subclass(self): + for stem in _task_names(): + with self.subTest(task=stem): + class_name = _yaml(_TASKS_DIR / f"{stem}.yaml")["class_name"] + self.assertTrue( + hasattr(mouse_task, class_name), + f"{class_name} not exposed on mouse_task", + ) + cls = getattr(mouse_task, class_name) + self.assertTrue(issubclass(cls, ActiveSensingTask)) + self.assertEqual(cls.__name__, class_name) + self.assertIn(class_name, mouse_task.__all__) + + def test_class_names_are_unique(self): + names = [_yaml(_TASKS_DIR / f"{s}.yaml")["class_name"] for s in _task_names()] + self.assertEqual(len(names), len(set(names)), "duplicate class_name across task YAMLs") + + def test_docstring_comes_from_yaml(self): + cls = mouse_task.ShapeDiscrim + expected = _yaml(_TASKS_DIR / "shape_mouse_discrim.yaml")["description"].strip() + self.assertIn(expected.splitlines()[0], (cls.__doc__ or "")) + + +class TestLoadTaskConfig(unittest.TestCase): + """common.yaml + tasks/.yaml merge behaviour.""" + + def test_task_without_overrides_equals_common(self): + common = _yaml(CONFIG_DIR / "common.yaml") + merged = load_task_config("mouse_discrim") # has an empty params block + self.assertEqual(merged, common) + + def test_overrides_win_over_common(self): + merged = load_task_config("shape_mouse_discrim") + # overridden in the task YAML + self.assertEqual(merged["target_selection"], 13.0) + self.assertEqual(merged["slit_depth"], 0.02) + self.assertEqual(merged["velocity_threshold"], 5.0) + # NOT overridden -> falls back to common.yaml + self.assertEqual(merged["reward_size"], 100) + self.assertEqual(merged["camera_type"], 1.0) + + def test_every_merged_config_has_all_common_keys(self): + common_keys = set(_yaml(CONFIG_DIR / "common.yaml")) + for stem in _task_names(): + with self.subTest(task=stem): + merged = load_task_config(stem) + self.assertEqual(set(merged), common_keys) + + def test_unknown_task_raises(self): + with self.assertRaises(FileNotFoundError): + load_task_config("does_not_exist") + + +class TestGeneratedSignatures(unittest.TestCase): + """The teensyexp GUI introspects __init__ via inspect — keep it real.""" + + def test_signature_lists_every_param_with_merged_default(self): + for stem in _task_names(): + cls = getattr(mouse_task, _yaml(_TASKS_DIR / f"{stem}.yaml")["class_name"]) + merged = load_task_config(stem) + defaults = _signature_defaults(cls) + with self.subTest(task=stem): + for key, value in merged.items(): + self.assertIn(key, defaults) + self.assertEqual( + defaults[key], value, f"{stem}.{key} default mismatch" + ) + + def test_signature_shape(self): + spec = inspect.getfullargspec(mouse_task.ShapeDiscrim.__init__) + self.assertEqual(spec.args[:2], ["self", "teensy"]) + self.assertEqual(spec.varkw, "kwargs") # forwards extra kwargs + self.assertIn("config_file_path", spec.args) # machine-specific config kept separate + + def test_all_args_after_teensy_have_defaults(self): + # The GUI (teensyexp) reads defaults positionally from the introspected + # signature, so every arg except self/teensy must carry a default. + spec = inspect.getfullargspec(mouse_task.ShapeDiscrim.__init__) + self.assertEqual(len(spec.defaults), len(spec.args) - 2) # all but self, teensy default + + +class TestTaskConstruction(unittest.TestCase): + """End-to-end parameter resolution onto a real instance (no Unity launch).""" + + def _build(self, factory, **kwargs): + # use_dlc=False avoids opening the DLC socket; process_config is mocked + # so no real Unity build path is required. + kwargs.setdefault("use_dlc", False) + with patch( + "mouse_task.task_active_sensing.process_config", return_value=_MOCK_CONFIG + ): + return factory(teensy=MagicMock(), **kwargs) + + def test_generated_class_resolves_merged_params(self): + task = self._build(mouse_task.ShapeDiscrim) + self.assertEqual(task.session_label, ["ar_shape_discrimination"]) + self.assertEqual(task.velocity_threshold, 5.0) + self.assertEqual(task.target_selection_param, 13.0) + self.assertEqual(task.slit_depth_param, 0.02) + self.assertEqual(task.target_distance_param, 4.0) + self.assertEqual(task.reward_size, [100]) # common default, wrapped by as_list + self.assertFalse(task.use_dlc) + + def test_base_class_task_config_argument(self): + # The main class can be driven directly by task_config alone. + task = self._build(ActiveSensingTask, task_config="shape_mouse_discrim") + self.assertEqual(task.velocity_threshold, 5.0) + self.assertEqual(task.target_selection_param, 13.0) + + def test_explicit_argument_overrides_config(self): + task = self._build( + ActiveSensingTask, task_config="shape_mouse_discrim", velocity_threshold=99.0 + ) + self.assertEqual(task.velocity_threshold, 99.0) # explicit wins + self.assertEqual(task.target_selection_param, 13.0) # rest from config + + def test_generated_class_forwards_gui_edit(self): + # The GUI passes every signature param as a kwarg; an edited value must win. + task = self._build(mouse_task.ShapeDiscrim, velocity_threshold=42.0) + self.assertEqual(task.velocity_threshold, 42.0) + + def test_missing_params_without_task_config_raises(self): + with patch( + "mouse_task.task_active_sensing.process_config", return_value=_MOCK_CONFIG + ): + with self.assertRaises(ValueError): + ActiveSensingTask(teensy=MagicMock(), use_dlc=False) # no task_config, no params + + +class TestTeensyGuiTaskLoading(unittest.TestCase): + """Regression tests for teensyexp task discovery and parameter introspection.""" + + def test_update_tasks_ignores_non_task_symbols(self): + from teensyexp.tasks_abc.task import Task + from teensyexp.teensy_experiment import TeensyExperimentGUI + + class GoodTask(Task): + def __init__(self, teensy, velocity_threshold=7.5, optional=None): + super().__init__(teensy) + + class NotATask: + def __init__(self, teensy, ignored=1): + self.teensy = teensy + + class _Loader: + def exec_module(self, module): + module.GoodTask = GoodTask + module.NotATask = NotATask + module._TASK_CLASSES = {"GoodTask": GoodTask} + module.some_constant = 123 + + fake_spec = types.SimpleNamespace(loader=_Loader()) + fake_module = types.ModuleType("fake_tasks_pkg") + + gui = TeensyExperimentGUI.__new__(TeensyExperimentGUI) + + with patch("teensyexp.teensy_experiment.importlib.util.find_spec", return_value=fake_spec), \ + patch("teensyexp.teensy_experiment.importlib.util.module_from_spec", return_value=fake_module): + gui.update_tasks("C:/tmp/fake_tasks_pkg") + + self.assertIn("GoodTask", gui.task_params) + self.assertNotIn("NotATask", gui.task_params) + self.assertEqual(gui.task_params["GoodTask"]["velocity_threshold"], 7.5) + self.assertIsNone(gui.task_params["GoodTask"]["optional"]) + + +if __name__ == "__main__": + unittest.main() diff --git a/teensyexp/teensy_experiment.py b/teensyexp/teensy_experiment.py index 02739d6ae..d1a678fa4 100644 --- a/teensyexp/teensy_experiment.py +++ b/teensyexp/teensy_experiment.py @@ -262,7 +262,15 @@ def update_tasks(self, task_dir): try: self.task_module = importlib.util.module_from_spec(task_spec) task_spec.loader.exec_module(self.task_module) - task_list = [t for t in dir(self.task_module) if '__' not in t] + from teensyexp.tasks_abc.task import Task + task_list = [] + # Only load real Task subclasses; task packages may export helper symbols too. + for t in dir(self.task_module): + if t.startswith('_'): + continue + obj = getattr(self.task_module, t) + if inspect.isclass(obj) and issubclass(obj, Task): + task_list.append(t) except AttributeError: if hasattr(self, "window"): messagebox.showerror("Failed to load tasks!", @@ -276,10 +284,22 @@ def update_tasks(self, task_dir): self.task_params = {} for t in task_list: obj = getattr(self.task_module, t) # TODO no getattr - args = inspect.getargspec(obj) self.task_params[t] = {} - for i in range(2, len(args[0])): - self.task_params[t][args[0][i]] = args[3][i - 2] + try: + # signature() is safer than getargspec for modern/dynamic callables. + signature = inspect.signature(obj.__init__) + except (TypeError, ValueError): + continue + + for name, param in signature.parameters.items(): + if name in ('self', 'teensy'): + continue + if param.kind in (inspect.Parameter.VAR_POSITIONAL, inspect.Parameter.VAR_KEYWORD): + continue + if param.default is inspect.Parameter.empty: + self.task_params[t][name] = None + else: + self.task_params[t][name] = param.default if hasattr(self, 'task_entry'): # TODO no hasattr self.task_name.set("") # gui interaction : update From 4a8f6a88febbb9ec6d71859660cdc1835805f0a0 Mon Sep 17 00:00:00 2001 From: CeliaBenquet Date: Wed, 8 Jul 2026 15:18:17 +0200 Subject: [PATCH 03/26] Change reward and start zone more to back --- .../tasks/far_reward_mouse_detection_p2.yaml | 14 ++++++++++---- 1 file changed, 10 insertions(+), 4 deletions(-) diff --git a/mouse_task/configs/tasks/far_reward_mouse_detection_p2.yaml b/mouse_task/configs/tasks/far_reward_mouse_detection_p2.yaml index a0b2f2031..70685da1e 100644 --- a/mouse_task/configs/tasks/far_reward_mouse_detection_p2.yaml +++ b/mouse_task/configs/tasks/far_reward_mouse_detection_p2.yaml @@ -18,11 +18,17 @@ params: r_report_box: - 5 - 10 - - -5 - - -3 + - -6 + - -4 l_report_box: - -10 - -5 - - -5 - - -3 + - -6 + - -4 + start_box: + - -4 + - 4 + - -9 + - -7 + - 90 distractor: 0.0 From c10ad7135f2be1e0ea68dca581e882ea082aac29 Mon Sep 17 00:00:00 2001 From: CeliaBenquet Date: Wed, 1 Jul 2026 16:51:16 +0200 Subject: [PATCH 04/26] Make processors compatible for dlclive gui update --- mouse_task/dlc_utils/dlcProcessor_dlconly.py | 29 +++++ mouse_task/dlc_utils/dlc_processor_socket.py | 104 ++++++++++++++++-- .../dlc_utils/dlc_processor_socket_pd.py | 79 ++++++++++++- .../dlc_utils/dlc_processor_socket_pd_sync.py | 79 ++++++++++++- mouse_task/dlc_utils/simple_processor.py | 34 +++++- 5 files changed, 308 insertions(+), 17 deletions(-) diff --git a/mouse_task/dlc_utils/dlcProcessor_dlconly.py b/mouse_task/dlc_utils/dlcProcessor_dlconly.py index 012b3d3c1..58f21da57 100644 --- a/mouse_task/dlc_utils/dlcProcessor_dlconly.py +++ b/mouse_task/dlc_utils/dlcProcessor_dlconly.py @@ -1,9 +1,27 @@ 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 + +@register_processor class dlc_only(Processor): + 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 = [] @@ -45,6 +63,17 @@ def save(self, filename): save_code = False return save_code + + +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..6ec17b744 100644 --- a/mouse_task/dlc_utils/dlc_processor_socket.py +++ b/mouse_task/dlc_utils/dlc_processor_socket.py @@ -1,27 +1,86 @@ 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): + 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,6 +94,15 @@ 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 + def process(self, pose: NDArray[np.float64], **kwargs: Any) -> NDArray[np.float64]: xy = pose[:, :2] @@ -80,7 +148,12 @@ 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 @@ -114,3 +187,14 @@ def save_latency_data(self) -> Dict[str, Any]: save_dict["head_angle"] = np.array(self.head_angle) return save_dict + + +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..5e7e357ac 100644 --- a/mouse_task/dlc_utils/dlc_processor_socket_pd.py +++ b/mouse_task/dlc_utils/dlc_processor_socket_pd.py @@ -1,13 +1,77 @@ import numpy as np +import importlib.util +import sys +from pathlib import Path from typing import Any, Dict +from dlclivegui.processors import PROCESSOR_REGISTRY, register_processor -from latency_tests.Teensy_latency.TeensyLatency import TeensyLatency -from dlc_utils.dlc_processor_socket import MyProcessor_socket +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): + 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): + 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 +109,14 @@ def save_latency_data(self) -> Dict[str, Any]: save_dict["photodiode_time"] = np.array(self.teensy.input_data_time) return save_dict + + +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..0fb950b36 100644 --- a/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py +++ b/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py @@ -1,14 +1,73 @@ """Sync photodiode processor that reuses the shared DLC/socket behavior.""" +import importlib.util +import sys +from pathlib import Path from typing import Any, Dict import numpy as np +from dlclivegui.processors import PROCESSOR_REGISTRY, register_processor # type: ignore[import-not-found] -from dlc_utils.dlc_processor_socket_pd import dlc_inference_w_pd -from latency_tests.Teensy_latency.TeensyLatencySync import TeensyLatencySync +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 +try: + from latency_tests.Teensy_latency.TeensyLatencySync import TeensyLatencySync +except ModuleNotFoundError: + TeensyLatencySync = None +PROCESSOR_REGISTRY.pop("dlc_inference_w_pd_sync", None) + + +@register_processor class dlc_inference_w_pd_sync(dlc_inference_w_pd): + 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", @@ -28,6 +87,11 @@ def __init__( ) def _create_teensy(self, com, baudrate): + if TeensyLatencySync is None: + raise ImportError( + "TeensyLatencySync dependency is unavailable. Ensure mouse_task is on PYTHONPATH " + "and Teensy latency modules are installed." + ) return TeensyLatencySync(com, baudrate=baudrate) def save_latency_data(self) -> Dict[str, Any]: @@ -40,3 +104,14 @@ def save_latency_data(self) -> Dict[str, Any]: ) return save_dict + + +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..8bf197862 100644 --- a/mouse_task/dlc_utils/simple_processor.py +++ b/mouse_task/dlc_utils/simple_processor.py @@ -1,11 +1,28 @@ 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): + 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): @@ -34,4 +51,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", {}), + } + } From 88a0c9e2e289393ecfb6d8748714fca56e2ed8e7 Mon Sep 17 00:00:00 2001 From: CeliaBenquet Date: Wed, 1 Jul 2026 17:38:49 +0200 Subject: [PATCH 05/26] Update dlcliveonly --- mouse_task/dlc_utils/dlcProcessor_dlconly.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/mouse_task/dlc_utils/dlcProcessor_dlconly.py b/mouse_task/dlc_utils/dlcProcessor_dlconly.py index 58f21da57..4c5d8f4ab 100644 --- a/mouse_task/dlc_utils/dlcProcessor_dlconly.py +++ b/mouse_task/dlc_utils/dlcProcessor_dlconly.py @@ -4,6 +4,8 @@ from math import sqrt, acos, atan2, copysign, degrees import pickle +PROCESSOR_REGISTRY.pop("dlc_only", None) + @register_processor class dlc_only(Processor): From ceabff961307814db19b028421935e5801cd1c88 Mon Sep 17 00:00:00 2001 From: CeliaBenquet Date: Thu, 9 Jul 2026 16:24:24 +0200 Subject: [PATCH 06/26] Add HEAD_CONF_THRESHOLD to MyProcessor_socket and dlc_inference_w_pd_sync classes --- mouse_task/dlc_utils/dlc_processor_socket.py | 3 ++- .../dlc_utils/dlc_processor_socket_pd_sync.py | 14 +++++++++++--- 2 files changed, 13 insertions(+), 4 deletions(-) diff --git a/mouse_task/dlc_utils/dlc_processor_socket.py b/mouse_task/dlc_utils/dlc_processor_socket.py index 6ec17b744..399650953 100644 --- a/mouse_task/dlc_utils/dlc_processor_socket.py +++ b/mouse_task/dlc_utils/dlc_processor_socket.py @@ -33,6 +33,7 @@ @register_processor class MyProcessor_socket(ProcessorWithSignal): + HEAD_CONF_THRESHOLD = 0.6 PROCESSOR_NAME = "SocketProcessor" PROCESSOR_DESCRIPTION = "Sends DLC-derived kinematics over a local socket." PROCESSOR_PARAMS = { @@ -111,7 +112,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) 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 0fb950b36..df4d0450a 100644 --- a/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py +++ b/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py @@ -4,6 +4,7 @@ import sys from pathlib import Path from typing import Any, Dict +import logging import numpy as np from dlclivegui.processors import PROCESSOR_REGISTRY, register_processor # type: ignore[import-not-found] @@ -23,18 +24,22 @@ _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: +except ModuleNotFoundError as e: TeensyLatencySync = None + import_issues = e PROCESSOR_REGISTRY.pop("dlc_inference_w_pd_sync", None) - +logger = logging.getLogger(__name__) @register_processor class dlc_inference_w_pd_sync(dlc_inference_w_pd): PROCESSOR_NAME = "SocketProcessorWithPDSync" + HEAD_CONF_THRESHOLD = 0.01 PROCESSOR_DESCRIPTION = "Photodiode processor with Teensy sync timing capture." + PROCESSOR_BUILD_IN_WORKER = True # legacy initialization ensures compatibility with old dlclivegui processors PROCESSOR_PARAMS = { "com": { "type": "str", @@ -85,13 +90,16 @@ def __init__( freq=freq, use_teensy=use_teensy, ) + logger.info( + f"Listener status: {self.listener._listener._socket.getsockname()} with authkey: {self.authkey}" + ) def _create_teensy(self, com, baudrate): 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 save_latency_data(self) -> Dict[str, Any]: From 8a8e0d32ac9812ebbe5728d22d0ddd3f70ebbbe1 Mon Sep 17 00:00:00 2001 From: Cyril Achard Date: Thu, 9 Jul 2026 16:28:59 +0200 Subject: [PATCH 07/26] Harden PD sync processor initialization cleanup Wrap processor initialization in a guarded try/except so partial startup failures trigger cleanup and re-raise as a clear RuntimeError. Add a dedicated stop() method (plus close() alias) to optionally save output and reliably release Teensy, socket connection, and listener resources, reducing leaked serial handles/ports after errors or shutdown. --- .../dlc_utils/dlc_processor_socket_pd_sync.py | 100 ++++++++++++++++-- 1 file changed, 89 insertions(+), 11 deletions(-) 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 df4d0450a..161283770 100644 --- a/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py +++ b/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py @@ -82,17 +82,23 @@ def __init__( 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, - ) - logger.info( - f"Listener status: {self.listener._listener._socket.getsockname()} with authkey: {self.authkey}" - ) + try: + super().__init__( + com=com, + baudrate=baudrate, + signal_delay=signal_delay, + signal_type=signal_type, + freq=freq, + use_teensy=use_teensy, + ) + logger.info( + f"Listener status: {self.listener._listener._socket.getsockname()} with authkey: {self.authkey}" + ) + except Exception as e: + self.stop(save=False) + raise RuntimeError( + f"Failed to initialize dlc_inference_w_pd_sync: {e}." + ) from e def _create_teensy(self, com, baudrate): if TeensyLatencySync is None: @@ -112,6 +118,78 @@ def save_latency_data(self) -> Dict[str, Any]: ) return save_dict + + def stop(self, save: bool = False, file: str | None = None) -> None: + """Cleanly stop processor resources. + + Args: + save: + If True, call self.save(file) before closing resources. + file: + Output path for processor data. If None, no save is attempted unless + the parent class has a meaningful default filename. + """ + + # 1. Optional save first, while all buffers/objects still exist. + if save: + try: + if file is not None: + self.save(file) + else: + # Avoid saving to an unknown location unless explicitly requested. + print("Processor stop(save=True) called without file; skipping save.") + except Exception as exc: + print(f"Processor save during stop failed: {exc}") + + # 2. Close Teensy serial if available. + 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 exc: + print(f"Failed to close Teensy cleanly: {exc}") + finally: + try: + self.teensy = None + except Exception: + pass + + # 3. Close accepted socket connection. + try: + conn = getattr(self, "conn", None) + if conn is not None: + conn.close() + except Exception as exc: + print(f"Failed to close processor socket connection: {exc}") + finally: + try: + self.conn = None + except Exception: + pass + + # 4. Close listener on port 6000. + try: + listener = getattr(self, "listener", None) + if listener is not None: + listener.close() + except Exception as exc: + print(f"Failed to close processor listener: {exc}") + finally: + try: + self.listener = None + except Exception: + pass + + + def close(self) -> None: + """Alias for generic cleanup.""" + self.stop(save=False) def get_available_processors() -> Dict[str, Dict[str, Any]]: From c86f22de66e916d9ce03179e34d9dd0165e40fd8 Mon Sep 17 00:00:00 2001 From: Cyril Achard Date: Thu, 9 Jul 2026 16:32:59 +0200 Subject: [PATCH 08/26] Add default-path save for PD sync processor Implement a dedicated `save()` method in `dlc_inference_w_pd_sync` that persists latency data to a pickle file, supports an optional explicit path, and falls back to `self.save_path` when no file is passed. The change also initializes `save_path` on construction, ensures parent directories are created, and adds warning-based error handling for missing paths or save failures. --- .../dlc_utils/dlc_processor_socket_pd_sync.py | 35 ++++++++++++++++++- 1 file changed, 34 insertions(+), 1 deletion(-) 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 161283770..34176b0da 100644 --- a/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py +++ b/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py @@ -3,8 +3,10 @@ import importlib.util import sys from pathlib import Path -from typing import Any, Dict +from typing import Any, Dict, Optional import logging +import warnings +import pickle import numpy as np from dlclivegui.processors import PROCESSOR_REGISTRY, register_processor # type: ignore[import-not-found] @@ -82,6 +84,7 @@ def __init__( freq=5, use_teensy=1, ): + self.save_path: Optional[Path] = None try: super().__init__( com=com, @@ -118,6 +121,36 @@ def save_latency_data(self) -> Dict[str, Any]: ) return save_dict + + def save(self, file: Optional[str] = None) -> int: + """Save processor data. + + If `file` is not provided, uses `self.save_path`. + """ + target = file + + if target is None: + target = 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) + + print(f"Processor data saved to: {target}") + return 1 + + except Exception as e: + warnings.warn(f"Proc file was not saved, an exception occurred: {e}") + return -1 def stop(self, save: bool = False, file: str | None = None) -> None: """Cleanly stop processor resources. From d8d6bbae533dae0c523f6b7aa60062cb1baefaa0 Mon Sep 17 00:00:00 2001 From: C-Achard Date: Fri, 10 Jul 2026 11:19:13 +0200 Subject: [PATCH 09/26] Enhance dlc_inference_w_pd_sync with legacy recording support and timestamp handling --- .../dlc_utils/dlc_processor_socket_pd_sync.py | 283 +++++++++++++++++- 1 file changed, 277 insertions(+), 6 deletions(-) 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 34176b0da..1c4739972 100644 --- a/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py +++ b/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py @@ -3,12 +3,17 @@ import importlib.util import sys from pathlib import Path +import time from typing import Any, Dict, Optional import logging import warnings +import json import pickle +import pandas as pd import numpy as np +from numpy.typing import NDArray + from dlclivegui.processors import PROCESSOR_REGISTRY, register_processor # type: ignore[import-not-found] try: @@ -85,6 +90,13 @@ def __init__( use_teensy=1, ): self.save_path: Optional[Path] = None + ### + self._legacy_recording_active = False + self._legacy_poses = [] + self._legacy_pose_times = [] + self._legacy_frame_times = [] + + try: super().__init__( com=com, @@ -102,7 +114,22 @@ def __init__( raise RuntimeError( f"Failed to initialize dlc_inference_w_pd_sync: {e}." ) from e - + + def process(self, pose: NDArray[np.float64], **kwargs: Any) -> NDArray[np.float64]: + processed_pose = super().process(pose, **kwargs) + if getattr(self, "_legacy_recording_active", False): + try: + # NOTE from @C-Achard to @CeliaBenquet: + # I am casting pose as float32 to mitigate + # potential issues with RAM usage over time. + # If this is a problem, change it back to float64. + self._legacy_poses.append(np.asarray(pose, dtype=np.float32).copy()) + self._legacy_pose_times.append(time.time()) + self._legacy_frame_times.append(kwargs.get("frame_time", time.time())) + except Exception: + logger.exception("Failed to buffer pose for legacy DLC save") + return processed_pose + def _create_teensy(self, com, baudrate): if TeensyLatencySync is None: raise ImportError( @@ -122,6 +149,51 @@ def save_latency_data(self) -> Dict[str, Any]: return save_dict + def set_dlc_cfg(self, dlc_cfg): + self.dlc_cfg = dlc_cfg + + def on_recording_started(self, context: dict) -> None: + """Receive recording context from DLCLiveGUI.""" + self.recording_context = dict(context or {}) + + base_path = self.recording_context.get("processor_base_path") + if base_path is None: + self.save_path = None + logger.warning("Processor recording started without processor_base_path") + return + + base_path = Path(base_path) + + # Old-style processor output. + self.save_path = base_path.parent / f"{base_path.name}_PROC" + + # Optional future outputs. + self.dlc_h5_path = base_path.parent / f"{base_path.name}_DLC.hdf5" + self.legacy_timestamp_path = base_path.parent / f"TS_{base_path.name}.npy" + + logger.info("Processor save path set to %s", self.save_path) + + def on_recording_stopped(self, context: dict) -> None: + """Save custom processor outputs after GUI recording stops.""" + self.recording_context = dict(context or getattr(self, "recording_context", {})) + + # Save old PROC pickle. + result = self.save() + logger.info("Processor legacy PROC save result: %r", result) + + dlc_h5_result = self.save_legacy_dlc_h5() + logger.info("Processor legacy DLC h5 save result: %r", dlc_h5_result) + + npy_ts_result = self.save_legacy_timestamp_npy() + logger.info("Processor legacy timestamp npy save result: %r", npy_ts_result) + + 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 save(self, file: Optional[str] = None) -> int: """Save processor data. @@ -151,7 +223,210 @@ def save(self, file: Optional[str] = None) -> int: except Exception as e: warnings.warn(f"Proc file was not saved, an exception occurred: {e}") return -1 + + 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) + + # Expected shape: frames x keypoints x 3 + poses = np.asarray(poses) + if poses.ndim == 2: + # Single frame case. + poses = poses[None, :, :] + + flat = poses.reshape((poses.shape[0], poses.shape[1] * poses.shape[2])) + + bodyparts = None + dlc_cfg = getattr(self, "dlc_cfg", 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.shape[1]: + pdindex = pd.MultiIndex.from_product( + [bodyparts, ["x", "y", "likelihood"]], + names=["bodyparts", "coords"], + ) + pose_df = pd.DataFrame(flat, columns=pdindex) + else: + pose_df = pd.DataFrame(flat) + + pose_df["frame_time"] = list(getattr(self, "_legacy_frame_times", [])) + pose_df["pose_time"] = list(getattr(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 _extract_timestamps_from_json(self, json_path: Path) -> np.ndarray: + """Extract software timestamps from new GUI timestamp JSON.""" + json_path = Path(json_path) + + with json_path.open("r", encoding="utf-8") as f: + data = json.load(f) + + # Current VideoRecorder schema. + if isinstance(data, dict) and isinstance(data.get("frame_timestamps"), list): + timestamps = [] + for rec in data["frame_timestamps"]: + if isinstance(rec, dict) and "software_timestamp" in rec: + timestamps.append(float(rec["software_timestamp"])) + return np.asarray(timestamps, dtype=float) + + # Tolerant fallback schemas. + 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): + for key in ("software_timestamp", "timestamp", "frame_time", "time"): + if key in item: + values.append(float(item[key])) + break + return np.asarray(values, dtype=float) + + return np.asarray([], dtype=float) + + def _legacy_timestamp_output_path( + self, + json_path: Path, + *, + base_path: Path | None, + index: int, + total: int, + ) -> Path: + """Derive old-style TS_*.npy path from timestamp JSON path/context.""" + json_path = Path(json_path) + + # Single DLC/camera recording: use processor base name. + if total == 1 and base_path is not None: + return base_path.parent / f"TS_{base_path.name}.npy" + + # Multi-camera case: derive from video filename. + name = json_path.name + + for suffix in ( + ".avi_timestamps.json", + ".mp4_timestamps.json", + "_timestamps.json", + ): + if name.endswith(suffix): + video_stem = name[: -len(suffix)] + break + else: + video_stem = json_path.stem + + return json_path.parent / f"TS_{video_stem}.npy" + + def _find_timestamp_json_files(self) -> list[Path]: + """Resolve timestamp JSON files from recording context.""" + context = getattr(self, "recording_context", {}) or {} + + value = ( + context.get("timestamp_json_files") + or context.get("timestamp_files") + or {} + ) + + paths: list[Path] = [] + + if isinstance(value, (str, Path)): + paths.append(Path(value)) + + elif isinstance(value, dict): + for item in value.values(): + if isinstance(item, (str, Path)): + paths.append(Path(item)) + + elif isinstance(value, (list, tuple, set)): + for item in value: + if isinstance(item, (str, Path)): + paths.append(Path(item)) + + paths = [p for p in paths if p.exists()] + + if paths: + return sorted(paths) + + # Fallback: scan run_dir. + run_dir = context.get("run_dir") + if run_dir is not None: + return sorted(Path(run_dir).glob("*_timestamps.json")) + + return [] + def save_legacy_timestamp_npy(self) -> int: + """Convert new GUI timestamp JSON files to old-style TS_*.npy files.""" + 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 + + context = getattr(self, "recording_context", {}) or {} + base_path = context.get("processor_base_path") + base_path = Path(base_path) if base_path is not None else None + + saved = 0 + + 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 + + out_path = self._legacy_timestamp_output_path( + json_path, + base_path=base_path, + index=index, + total=len(json_paths), + ) + + out_path.parent.mkdir(parents=True, exist_ok=True) + np.save(out_path, timestamps) + + logger.info( + "Legacy 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 stop(self, save: bool = False, file: str | None = None) -> None: """Cleanly stop processor resources. @@ -166,11 +441,7 @@ def stop(self, save: bool = False, file: str | None = None) -> None: # 1. Optional save first, while all buffers/objects still exist. if save: try: - if file is not None: - self.save(file) - else: - # Avoid saving to an unknown location unless explicitly requested. - print("Processor stop(save=True) called without file; skipping save.") + self.save(file) except Exception as exc: print(f"Processor save during stop failed: {exc}") From 89487a47b675ebf0488ddd61b405b30f42d95374 Mon Sep 17 00:00:00 2001 From: C-Achard Date: Fri, 10 Jul 2026 11:20:51 +0200 Subject: [PATCH 10/26] Refactor dlc_inference_w_pd_sync for improved legacy support and enhanced logging --- .../dlc_utils/dlc_processor_socket_pd_sync.py | 648 ++++++++++++------ 1 file changed, 422 insertions(+), 226 deletions(-) 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 1c4739972..aa8bf685e 100644 --- a/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py +++ b/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py @@ -1,21 +1,25 @@ """Sync photodiode processor that reuses the shared DLC/socket behavior.""" +from __future__ import annotations + import importlib.util +import json +import logging +import pickle +import shutil import sys -from pathlib import Path import time -from typing import Any, Dict, Optional -import logging import warnings -import json -import pickle +from pathlib import Path +from typing import Any, Dict, Optional -import pandas as pd 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: @@ -24,13 +28,16 @@ _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 @@ -38,15 +45,22 @@ TeensyLatencySync = None import_issues = e + PROCESSOR_REGISTRY.pop("dlc_inference_w_pd_sync", None) logger = logging.getLogger(__name__) + @register_processor class dlc_inference_w_pd_sync(dlc_inference_w_pd): PROCESSOR_NAME = "SocketProcessorWithPDSync" - HEAD_CONF_THRESHOLD = 0.01 PROCESSOR_DESCRIPTION = "Photodiode processor with Teensy sync timing capture." - PROCESSOR_BUILD_IN_WORKER = True # legacy initialization ensures compatibility with old dlclivegui processors + + # Legacy initialization ensures compatibility with old DLCLiveGUI processors: + # sockets / serial / side-effect-heavy resources are created inside DLCLiveWorker. + PROCESSOR_BUILD_IN_WORKER = True + + HEAD_CONF_THRESHOLD = 0.01 + PROCESSOR_PARAMS = { "com": { "type": "str", @@ -82,21 +96,26 @@ class dlc_inference_w_pd_sync(dlc_inference_w_pd): def __init__( self, - com="COM3", - baudrate=9600, - signal_delay=10, - signal_type="pulse_geo", - freq=5, - use_teensy=1, - ): + com: str = "COM3", + baudrate: int = 9600, + signal_delay: float = 10, + signal_type: str = "pulse_geo", + freq: float = 5, + use_teensy: int | bool = 1, + ) -> None: + 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 = [] - self._legacy_pose_times = [] - self._legacy_frame_times = [] + 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, @@ -106,103 +125,132 @@ def __init__( freq=freq, use_teensy=use_teensy, ) - logger.info( - f"Listener status: {self.listener._listener._socket.getsockname()} with authkey: {self.authkey}" - ) + + 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 - + 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) - if getattr(self, "_legacy_recording_active", False): + + # 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: - # NOTE from @C-Achard to @CeliaBenquet: - # I am casting pose as float32 to mitigate - # potential issues with RAM usage over time. - # If this is a problem, change it back to float64. - self._legacy_poses.append(np.asarray(pose, dtype=np.float32).copy()) + # 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(kwargs.get("frame_time", 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 processed_pose - + + return pose_to_buffer + def _create_teensy(self, com, baudrate): 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 save_latency_data(self) -> Dict[str, Any]: - save_dict = super().save_latency_data() - if self.use_teensy == 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 TeensyLatencySync(com, baudrate=baudrate) - return save_dict - def set_dlc_cfg(self, dlc_cfg): 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.recording_context.get("processor_base_path") + 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 - base_path = Path(base_path) - - # Old-style processor output. + # Primary/native processor outputs. self.save_path = base_path.parent / f"{base_path.name}_PROC" - - # Optional future outputs. self.dlc_h5_path = base_path.parent / f"{base_path.name}_DLC.hdf5" - self.legacy_timestamp_path = base_path.parent / f"TS_{base_path.name}.npy" + + # 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 on_recording_stopped(self, context: dict) -> None: - """Save custom processor outputs after GUI recording stops.""" - self.recording_context = dict(context or getattr(self, "recording_context", {})) + """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 - # Save old PROC pickle. - result = self.save() - logger.info("Processor legacy PROC save result: %r", result) + proc_result = self.save() + logger.info("Processor legacy PROC save result: %r", proc_result) dlc_h5_result = self.save_legacy_dlc_h5() logger.info("Processor legacy DLC h5 save result: %r", dlc_h5_result) - + npy_ts_result = self.save_legacy_timestamp_npy() logger.info("Processor legacy timestamp npy save result: %r", npy_ts_result) - - 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 save(self, file: Optional[str] = None) -> int: - """Save processor data. + + video_copy_result = self.copy_legacy_video_files() + logger.info("Processor legacy video copy result: %r", video_copy_result) + + align_result = self.copy_processor_outputs_to_primary_legacy_base() + logger.info("Processor legacy output alignment result: %r", align_result) + + self._clear_legacy_pose_buffers() + + # ------------------------------------------------------------------ + # Primary PROC save + # ------------------------------------------------------------------ + + def save_latency_data(self) -> Dict[str, Any]: + save_dict = super().save_latency_data() + + 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 = file - - if target is None: - target = getattr(self, "save_path", None) + 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.") @@ -212,18 +260,21 @@ def save(self, file: Optional[str] = None) -> int: 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) + pickle.dump(self.save_latency_data(), f) - print(f"Processor data saved to: {target}") + 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) @@ -240,34 +291,28 @@ def save_legacy_dlc_h5(self) -> int: target = Path(target) target.parent.mkdir(parents=True, exist_ok=True) - # Expected shape: frames x keypoints x 3 - poses = np.asarray(poses) if poses.ndim == 2: - # Single frame case. poses = poses[None, :, :] - flat = poses.reshape((poses.shape[0], poses.shape[1] * poses.shape[2])) - - bodyparts = None - dlc_cfg = getattr(self, "dlc_cfg", None) + if poses.ndim != 3 or poses.shape[-1] != 3: + logger.warning("Skipping DLC h5 save: unexpected pose shape %s", poses.shape) + return 0 - if isinstance(dlc_cfg, dict): - bodyparts = ( - dlc_cfg.get("all_joints_names") - or dlc_cfg.get("metadata", {}).get("bodyparts") - ) + flat = poses.reshape((poses.shape[0], poses.shape[1] * poses.shape[2])) - if bodyparts and len(bodyparts) * 3 == flat.shape[1]: + 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(getattr(self, "_legacy_frame_times", [])) - pose_df["pose_time"] = list(getattr(self, "_legacy_pose_times", [])) + 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") @@ -277,23 +322,82 @@ def save_legacy_dlc_h5(self) -> int: 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: + 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: + """Convert new GUI timestamp JSON files to pipeline-compatible .npy files. + + Saves both: + - _TS.npy old GUI style + - TS_.npy pipeline/database matcher style + """ + 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 + + for json_path in 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 + + legacy_base = self._legacy_base_for_timestamp_json(json_path) + for out_path in self._timestamp_output_paths(legacy_base): + self._save_npy(out_path, timestamps) + logger.info( + "Legacy 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 software timestamps from new GUI timestamp JSON.""" + """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) - # Current VideoRecorder schema. if isinstance(data, dict) and isinstance(data.get("frame_timestamps"), list): - timestamps = [] - for rec in data["frame_timestamps"]: - if isinstance(rec, dict) and "software_timestamp" in rec: - timestamps.append(float(rec["software_timestamp"])) + 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) - # Tolerant fallback schemas. if isinstance(data, dict): for key in ("timestamps", "frame_times", "times"): values = data.get(key) @@ -306,146 +410,243 @@ def _extract_timestamps_from_json(self, json_path: Path) -> np.ndarray: if isinstance(item, (int, float)): values.append(float(item)) elif isinstance(item, dict): - for key in ("software_timestamp", "timestamp", "frame_time", "time"): - if key in item: - values.append(float(item[key])) - break + 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 _legacy_timestamp_output_path( - self, - json_path: Path, - *, - base_path: Path | None, - index: int, - total: int, - ) -> Path: - """Derive old-style TS_*.npy path from timestamp JSON path/context.""" - json_path = Path(json_path) - - # Single DLC/camera recording: use processor base name. - if total == 1 and base_path is not None: - return base_path.parent / f"TS_{base_path.name}.npy" - - # Multi-camera case: derive from video filename. - name = json_path.name - for suffix in ( - ".avi_timestamps.json", - ".mp4_timestamps.json", - "_timestamps.json", - ): - if name.endswith(suffix): - video_stem = name[: -len(suffix)] - break - else: - video_stem = json_path.stem - - return json_path.parent / f"TS_{video_stem}.npy" - def _find_timestamp_json_files(self) -> list[Path]: """Resolve timestamp JSON files from recording context.""" - context = getattr(self, "recording_context", {}) or {} - - value = ( - context.get("timestamp_json_files") - or context.get("timestamp_files") - or {} + paths = self._paths_from_context_value( + self.recording_context.get("timestamp_json_files") + or self.recording_context.get("timestamp_files") ) - paths: list[Path] = [] - - if isinstance(value, (str, Path)): - paths.append(Path(value)) - - elif isinstance(value, dict): - for item in value.values(): - if isinstance(item, (str, Path)): - paths.append(Path(item)) - - elif isinstance(value, (list, tuple, set)): - for item in value: - if isinstance(item, (str, Path)): - paths.append(Path(item)) - paths = [p for p in paths if p.exists()] - if paths: return sorted(paths) - # Fallback: scan run_dir. - run_dir = context.get("run_dir") + run_dir = self.recording_context.get("run_dir") if run_dir is not None: return sorted(Path(run_dir).glob("*_timestamps.json")) return [] - - def save_legacy_timestamp_npy(self) -> int: - """Convert new GUI timestamp JSON files to old-style TS_*.npy files.""" - 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 + def _timestamp_output_paths(self, legacy_base: Path) -> list[Path]: + """Return both old-GUI and pipeline-compatible timestamp names.""" + suffix_style = legacy_base.parent / f"{legacy_base.name}_TS.npy" + prefix_style = legacy_base.parent / f"TS_{legacy_base.name}.npy" + return self._unique_paths([suffix_style, prefix_style]) - context = getattr(self, "recording_context", {}) or {} - base_path = context.get("processor_base_path") - base_path = Path(base_path) if base_path is not None else None + # ------------------------------------------------------------------ + # Legacy video/sidecar compatibility copies + # ------------------------------------------------------------------ - saved = 0 + def copy_legacy_video_files(self) -> int: + """Copy new GUI video files to old-style _VIDEO. files.""" + copied = 0 - for index, json_path in enumerate(json_paths): + for video_path in self._find_video_files(): try: - timestamps = self._extract_timestamps_from_json(json_path) + legacy_base = self._legacy_base_for_video(video_path) + out_path = legacy_base.parent / f"{legacy_base.name}_VIDEO{video_path.suffix}" - if timestamps.size == 0: - logger.warning("No timestamps extracted from %s", json_path) - continue + if self._copy_file_if_needed(video_path, out_path): + copied += 1 - out_path = self._legacy_timestamp_output_path( - json_path, - base_path=base_path, - index=index, - total=len(json_paths), - ) - - out_path.parent.mkdir(parents=True, exist_ok=True) - np.save(out_path, timestamps) + except Exception: + logger.exception("Failed to copy legacy video file %s", video_path) - logger.info( - "Legacy timestamp npy saved to %s with %d timestamps", - out_path, - len(timestamps), - ) - saved += 1 + return 1 if copied else 0 - except Exception: - logger.exception("Failed to convert timestamp JSON to npy: %s", json_path) + def copy_processor_outputs_to_primary_legacy_base(self) -> int: + """Copy PROC and DLC H5 to match the primary video/timestamp legacy base. - return 1 if saved else 0 - - - def stop(self, save: bool = False, file: str | None = None) -> None: - """Cleanly stop processor resources. - - Args: - save: - If True, call self.save(file) before closing resources. - file: - Output path for processor data. If None, no save is attempted unless - the parent class has a meaningful default filename. + This is useful when processor_base_path differs from the actual recorder + video stem. """ + legacy_base = self._primary_legacy_base() + if legacy_base is None: + return 0 + + copied = 0 + + src_proc = getattr(self, "save_path", None) + if src_proc is not None: + dst_proc = legacy_base.parent / f"{legacy_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 = legacy_base.parent / f"{legacy_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 _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: + timestamp_jsons = self._find_timestamp_json_files() + if timestamp_jsons: + return self._legacy_base_for_timestamp_json(timestamp_jsons[0]) + + video_files = self._find_video_files() + if video_files: + return self._legacy_base_for_video(video_files[0]) - # 1. Optional save first, while all buffers/objects still exist. + return self._context_processor_base_path() + + 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") + + # ------------------------------------------------------------------ + # Cleanup + # ------------------------------------------------------------------ + + def stop(self, save: bool = False, file: str | Path | None = None) -> None: + """Cleanly stop processor resources.""" if save: try: self.save(file) - except Exception as exc: - print(f"Processor save during stop failed: {exc}") + except Exception: + logger.exception("Processor save during stop failed") - # 2. Close Teensy serial if available. + self._close_teensy() + self._close_socket_connection() + self._close_listener() + + def close(self) -> None: + """Alias for generic cleanup.""" + self.stop(save=False) + + def _close_teensy(self) -> None: try: teensy = getattr(self, "teensy", None) if teensy is not None: @@ -456,34 +657,34 @@ def stop(self, save: bool = False, file: str | None = None) -> None: close = getattr(teensy, "close", None) if callable(close): close() - except Exception as exc: - print(f"Failed to close Teensy cleanly: {exc}") + except Exception: + logger.exception("Failed to close Teensy cleanly") finally: try: self.teensy = None except Exception: pass - # 3. Close accepted socket connection. + def _close_socket_connection(self) -> None: try: conn = getattr(self, "conn", None) if conn is not None: conn.close() - except Exception as exc: - print(f"Failed to close processor socket connection: {exc}") + except Exception: + logger.exception("Failed to close processor socket connection") finally: try: self.conn = None except Exception: pass - # 4. Close listener on port 6000. + def _close_listener(self) -> None: try: listener = getattr(self, "listener", None) if listener is not None: listener.close() - except Exception as exc: - print(f"Failed to close processor listener: {exc}") + except Exception: + logger.exception("Failed to close processor listener") finally: try: self.listener = None @@ -491,11 +692,6 @@ def stop(self, save: bool = False, file: str | None = None) -> None: pass - def close(self) -> None: - """Alias for generic cleanup.""" - self.stop(save=False) - - def get_available_processors() -> Dict[str, Dict[str, Any]]: return { "dlc_inference_w_pd_sync": { @@ -504,4 +700,4 @@ def get_available_processors() -> Dict[str, Dict[str, Any]]: "description": getattr(dlc_inference_w_pd_sync, "PROCESSOR_DESCRIPTION", ""), "params": getattr(dlc_inference_w_pd_sync, "PROCESSOR_PARAMS", {}), } - } + } \ No newline at end of file From 73e12ae6a9f9c2169bef35d42899bebf53306235 Mon Sep 17 00:00:00 2001 From: C-Achard Date: Fri, 10 Jul 2026 11:51:34 +0200 Subject: [PATCH 11/26] Refactor dlc_inference_w_pd_sync for DB compatibility and improved timestamp handling --- .../dlc_utils/dlc_processor_socket_pd_sync.py | 201 ++++++++++++++---- 1 file changed, 159 insertions(+), 42 deletions(-) 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 aa8bf685e..6dd7b2cc9 100644 --- a/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py +++ b/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py @@ -2,10 +2,12 @@ from __future__ import annotations +import datetime import importlib.util import json import logging import pickle +import re import shutil import sys import time @@ -343,12 +345,6 @@ def _get_bodyparts_for_pose_width(self, flat_width: int) -> list[str] | None: # ------------------------------------------------------------------ def save_legacy_timestamp_npy(self) -> int: - """Convert new GUI timestamp JSON files to pipeline-compatible .npy files. - - Saves both: - - _TS.npy old GUI style - - TS_.npy pipeline/database matcher style - """ json_paths = self._find_timestamp_json_files() if not json_paths: @@ -356,19 +352,24 @@ def save_legacy_timestamp_npy(self) -> int: return 0 saved = 0 + total = len(json_paths) + compat_base = self._db_compat_base() - for json_path in json_paths: + 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 - legacy_base = self._legacy_base_for_timestamp_json(json_path) - for out_path in self._timestamp_output_paths(legacy_base): + for out_path in self._timestamp_output_paths( + compat_base, + index=index, + total=total, + ): self._save_npy(out_path, timestamps) logger.info( - "Legacy timestamp npy saved to %s with %d timestamps", + "DB-compatible timestamp npy saved to %s with %d timestamps", out_path, len(timestamps), ) @@ -378,7 +379,7 @@ def save_legacy_timestamp_npy(self) -> int: 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. @@ -434,54 +435,57 @@ def _find_timestamp_json_files(self) -> list[Path]: return [] - def _timestamp_output_paths(self, legacy_base: Path) -> list[Path]: - """Return both old-GUI and pipeline-compatible timestamp names.""" - suffix_style = legacy_base.parent / f"{legacy_base.name}_TS.npy" - prefix_style = legacy_base.parent / f"TS_{legacy_base.name}.npy" - return self._unique_paths([suffix_style, prefix_style]) + 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: - """Copy new GUI video files to old-style _VIDEO. files.""" copied = 0 + video_files = self._find_video_files() + compat_base = self._db_compat_base() + total = len(video_files) - for video_path in self._find_video_files(): + for index, video_path in enumerate(video_files): try: - legacy_base = self._legacy_base_for_video(video_path) - out_path = legacy_base.parent / f"{legacy_base.name}_VIDEO{video_path.suffix}" + 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 legacy video file %s", video_path) + 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: - """Copy PROC and DLC H5 to match the primary video/timestamp legacy base. - - This is useful when processor_base_path differs from the actual recorder - video stem. - """ - legacy_base = self._primary_legacy_base() - if legacy_base is None: - return 0 - + compat_base = self._db_compat_base() copied = 0 src_proc = getattr(self, "save_path", None) if src_proc is not None: - dst_proc = legacy_base.parent / f"{legacy_base.name}_PROC" + 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 = legacy_base.parent / f"{legacy_base.name}_DLC.hdf5" + dst_h5 = compat_base.parent / f"{compat_base.name}_DLC.hdf5" if self._copy_file_if_needed(Path(src_h5), dst_h5): copied += 1 @@ -490,21 +494,134 @@ def copy_processor_outputs_to_primary_legacy_base(self) -> int: # ------------------------------------------------------------------ # Legacy base / path helpers # ------------------------------------------------------------------ + 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() + + 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 / f"{mouse}_{date}_{attempt}" + + + 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: - timestamp_jsons = self._find_timestamp_json_files() - if timestamp_jsons: - return self._legacy_base_for_timestamp_json(timestamp_jsons[0]) - - video_files = self._find_video_files() - if video_files: - return self._legacy_base_for_video(video_files[0]) - - return self._context_processor_base_path() + 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.""" From 1f66996d139545e9919b1b5937bc890a9096e1ab Mon Sep 17 00:00:00 2001 From: C-Achard Date: Fri, 10 Jul 2026 17:45:55 +0200 Subject: [PATCH 12/26] Use direct datetime import in DLC sync Replaced `import datetime` with `from datetime import datetime` in `dlc_processor_socket_pd_sync.py` to align the import with direct `datetime` usage and avoid module/class ambiguity. --- mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) 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 6dd7b2cc9..bae6ddc2c 100644 --- a/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py +++ b/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py @@ -2,7 +2,7 @@ from __future__ import annotations -import datetime +from datetime import datetime import importlib.util import json import logging From b57ad2825db09624313f0393b5fe1edd6189afc9 Mon Sep 17 00:00:00 2001 From: C-Achard Date: Fri, 10 Jul 2026 18:20:44 +0200 Subject: [PATCH 13/26] Use video prefix in DB compat base path Add a `_video_prefix()` helper that derives the base name from discovered video files, falling back to `filename_stem` and then `recording`. Update `_db_compat_base()` to use this prefix instead of the mouse identifier when building the output path, improving alignment with recording/video naming. --- .../dlc_utils/dlc_processor_socket_pd_sync.py | 17 +++++++++++++++-- 1 file changed, 15 insertions(+), 2 deletions(-) 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 bae6ddc2c..dcb7af52e 100644 --- a/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py +++ b/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py @@ -494,6 +494,18 @@ def copy_processor_outputs_to_primary_legacy_base(self) -> int: # ------------------------------------------------------------------ # Legacy base / path helpers # ------------------------------------------------------------------ + def _video_prefix(self) -> str: + videos = self._find_video_files() + + if videos: + return videos[0].stem + + filename_stem = self.recording_context.get("filename_stem") + if filename_stem: + return str(filename_stem) + + return "recording" + def _db_compat_base(self) -> Path: """Return DB-GUI-compatible base path. @@ -515,11 +527,12 @@ def _db_compat_base(self) -> Path: run_dir = context.get("run_dir") run_dir = Path(run_dir) if run_dir is not None else self._fallback_output_dir() - mouse = self._mouse_from_context_or_run_dir(run_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 / f"{mouse}_{date}_{attempt}" + return run_dir / f"{prefix}_{date}_{attempt}" def _fallback_output_dir(self) -> Path: From 2c2f5fc530082ddefb0409fda091274c8931bd70 Mon Sep 17 00:00:00 2001 From: C-Achard Date: Fri, 10 Jul 2026 18:22:06 +0200 Subject: [PATCH 14/26] Update dlc_processor_socket_pd_sync.py --- mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) 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 dcb7af52e..b9b44b417 100644 --- a/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py +++ b/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py @@ -498,11 +498,11 @@ def _video_prefix(self) -> str: videos = self._find_video_files() if videos: - return videos[0].stem + return videos[0].stem.split("_", 1)[0] filename_stem = self.recording_context.get("filename_stem") if filename_stem: - return str(filename_stem) + return str(filename_stem).split("_", 1)[0] return "recording" From 3b64ba8cc683fb0d2b8423528d4ab19a98979797 Mon Sep 17 00:00:00 2001 From: C-Achard Date: Fri, 10 Jul 2026 18:27:46 +0200 Subject: [PATCH 15/26] Update dlc_processor_socket_pd_sync.py --- mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) 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 b9b44b417..294b4e9cd 100644 --- a/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py +++ b/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py @@ -532,7 +532,7 @@ def _db_compat_base(self) -> Path: date = self._date_from_context_or_run_dir(run_dir) attempt = self._attempt_from_context(default="1") - return run_dir / f"{prefix}_{date}_{attempt}" + return run_dir / f"vr4mice_{prefix}_{date}_{attempt}" def _fallback_output_dir(self) -> Path: From 1585c93e13c542bcff1496e41b3ac4859a938593 Mon Sep 17 00:00:00 2001 From: Cyril Achard Date: Mon, 13 Jul 2026 10:47:25 +0200 Subject: [PATCH 16/26] Only save copy to one timestamp file --- mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) 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 294b4e9cd..419e54fb2 100644 --- a/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py +++ b/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py @@ -447,7 +447,7 @@ def _timestamp_output_paths( 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", + # compat_base.parent / f"TIMESTAMP_{compat_base.name}_{camera_token}.npy", ] ) From 8db9d0b756e1bd6e2d9262e42485171b8b86a09c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?C=C3=A9lia=20Benquet?= Date: Mon, 13 Jul 2026 19:30:09 +0200 Subject: [PATCH 17/26] Set dlc threshold back to 0.6 --- mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) 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 419e54fb2..ace3932fb 100644 --- a/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py +++ b/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py @@ -61,7 +61,7 @@ class dlc_inference_w_pd_sync(dlc_inference_w_pd): # sockets / serial / side-effect-heavy resources are created inside DLCLiveWorker. PROCESSOR_BUILD_IN_WORKER = True - HEAD_CONF_THRESHOLD = 0.01 + HEAD_CONF_THRESHOLD = 0.6 # NOTE: this might need to be adjusted based on the model PROCESSOR_PARAMS = { "com": { From f620e2a196b76f9cef50f7ff06a61737a68fee9d Mon Sep 17 00:00:00 2001 From: Cyril Achard Date: Fri, 17 Jul 2026 13:48:05 +0200 Subject: [PATCH 18/26] Handle multi-detection poses in PD sync Add `_select_single_pose` to normalize incoming pose arrays before processing. The method now accepts both `(K, 3)` and `(N, K, 3)` inputs and, when multiple detections are present, picks the one with the highest mean keypoint likelihood (with a warning fallback when scores are non-finite). `process()` now feeds this selected single pose into the parent processor to prevent shape mismatches while keeping legacy buffering behavior. --- .../dlc_utils/dlc_processor_socket_pd_sync.py | 50 ++++++++++++++++++- 1 file changed, 49 insertions(+), 1 deletion(-) 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 ace3932fb..fb57e4e28 100644 --- a/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py +++ b/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py @@ -144,10 +144,58 @@ def __init__( # ------------------------------------------------------------------ # DLCLive processor API # ------------------------------------------------------------------ + @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(): + logger.warning( + "No detection has a finite confidence score; " + "selecting detection 0" + ) + return poses[0] + + index = int(np.nanargmax(scores)) + + + return poses[index] + + raise ValueError( + "Expected pose shape (K, 3) or (N, K, 3), " + f"got {poses.shape}" + ) 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) + most_likely_detected_pose = self._select_single_pose(pose) # IMPORTANT: not multi-animal friendly + processed_pose = super().process(most_likely_detected_pose, **kwargs) # Parent should return pose, but guard just in case. pose_to_buffer = processed_pose if processed_pose is not None else pose From 32909c2b65c10efceeab531e52747db9394cc845 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?C=C3=A9lia=20Benquet?= Date: Thu, 23 Jul 2026 14:37:08 +0200 Subject: [PATCH 19/26] Refactor processor classes --- mouse_task/dlc_utils/dlc_processor_socket.py | 120 +++++++++- .../dlc_utils/dlc_processor_socket_pd.py | 41 ++++ .../dlc_utils/dlc_processor_socket_pd_sync.py | 219 ++++++------------ 3 files changed, 233 insertions(+), 147 deletions(-) diff --git a/mouse_task/dlc_utils/dlc_processor_socket.py b/mouse_task/dlc_utils/dlc_processor_socket.py index 399650953..7285ea6d5 100644 --- a/mouse_task/dlc_utils/dlc_processor_socket.py +++ b/mouse_task/dlc_utils/dlc_processor_socket.py @@ -33,7 +33,23 @@ @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 = { @@ -104,7 +120,54 @@ def _ensure_connection(self) -> None: 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] @@ -152,13 +215,19 @@ def process(self, pose: NDArray[np.float64], **kwargs: Any) -> NDArray[np.float6 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]]) + 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: @@ -176,6 +245,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) @@ -189,6 +259,54 @@ def save_latency_data(self) -> Dict[str, Any]: 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 { diff --git a/mouse_task/dlc_utils/dlc_processor_socket_pd.py b/mouse_task/dlc_utils/dlc_processor_socket_pd.py index 5e7e357ac..273ed9d0a 100644 --- a/mouse_task/dlc_utils/dlc_processor_socket_pd.py +++ b/mouse_task/dlc_utils/dlc_processor_socket_pd.py @@ -1,6 +1,7 @@ import numpy as np import importlib.util import sys +import warnings from pathlib import Path from typing import Any, Dict @@ -31,6 +32,20 @@ @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 = { @@ -67,6 +82,7 @@ class dlc_inference_w_pd(MyProcessor_socket): } 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 " @@ -110,6 +126,31 @@ def save_latency_data(self) -> Dict[str, Any]: 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 { 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 fb57e4e28..325fb588a 100644 --- a/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py +++ b/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py @@ -54,14 +54,28 @@ @register_processor class dlc_inference_w_pd_sync(dlc_inference_w_pd): - PROCESSOR_NAME = "SocketProcessorWithPDSync" - PROCESSOR_DESCRIPTION = "Photodiode processor with Teensy sync timing capture." + """`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 - HEAD_CONF_THRESHOLD = 0.6 # NOTE: this might need to be adjusted based on the model + PROCESSOR_NAME = "SocketProcessorWithPDSync" + PROCESSOR_DESCRIPTION = "Photodiode processor with Teensy sync timing capture." PROCESSOR_PARAMS = { "com": { @@ -105,6 +119,7 @@ def __init__( 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 @@ -135,67 +150,23 @@ def __init__( self.authkey, ) except Exception: - logger.info("Listener initialized with authkey: %r", getattr(self, "authkey", None)) + 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 + raise RuntimeError( + f"Failed to initialize dlc_inference_w_pd_sync: {e}." + ) from e # ------------------------------------------------------------------ # DLCLive processor API # ------------------------------------------------------------------ - @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(): - logger.warning( - "No detection has a finite confidence score; " - "selecting detection 0" - ) - return poses[0] - - index = int(np.nanargmax(scores)) - - - return poses[index] - - raise ValueError( - "Expected pose shape (K, 3) or (N, K, 3), " - f"got {poses.shape}" - ) - 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.""" - most_likely_detected_pose = self._select_single_pose(pose) # IMPORTANT: not multi-animal friendly - processed_pose = super().process(most_likely_detected_pose, **kwargs) + 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 @@ -203,15 +174,20 @@ def process(self, pose: NDArray[np.float64], **kwargs: Any) -> NDArray[np.float6 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_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()))) + 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 " @@ -221,6 +197,7 @@ def _create_teensy(self, com, baudrate): 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 # ------------------------------------------------------------------ @@ -285,6 +262,7 @@ def on_recording_stopped(self, context: dict) -> None: # ------------------------------------------------------------------ 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 getattr(self, "use_teensy", 0) == 1: @@ -345,7 +323,9 @@ def save_legacy_dlc_h5(self) -> int: poses = poses[None, :, :] if poses.ndim != 3 or poses.shape[-1] != 3: - logger.warning("Skipping DLC h5 save: unexpected pose shape %s", poses.shape) + 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])) @@ -358,7 +338,9 @@ def save_legacy_dlc_h5(self) -> int: ) pose_df = pd.DataFrame(flat, columns=pdindex) else: - logger.warning("Bodyparts information not found or mismatched; saving DLC h5 without labels.") + 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) @@ -374,14 +356,14 @@ def save_legacy_dlc_h5(self) -> int: 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") - ) + 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) @@ -396,7 +378,9 @@ def save_legacy_timestamp_npy(self) -> int: json_paths = self._find_timestamp_json_files() if not json_paths: - logger.warning("Skipping legacy timestamp npy save: no timestamp JSON files found") + logger.warning( + "Skipping legacy timestamp npy save: no timestamp JSON files found" + ) return 0 saved = 0 @@ -424,10 +408,12 @@ def save_legacy_timestamp_npy(self) -> int: saved += 1 except Exception: - logger.exception("Failed to convert timestamp JSON to npy: %s", json_path) + 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. @@ -459,7 +445,9 @@ def _extract_timestamps_from_json(self, json_path: Path) -> np.ndarray: if isinstance(item, (int, float)): values.append(float(item)) elif isinstance(item, dict): - value = self._first_present(item, ("software_timestamp", "timestamp", "frame_time", "time")) + 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) @@ -511,13 +499,18 @@ def copy_legacy_video_files(self) -> int: 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}" + 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) + logger.exception( + "Failed to copy DB-compatible video file %s", video_path + ) return 1 if copied else 0 @@ -553,7 +546,7 @@ def _video_prefix(self) -> str: return str(filename_stem).split("_", 1)[0] return "recording" - + def _db_compat_base(self) -> Path: """Return DB-GUI-compatible base path. @@ -582,7 +575,6 @@ def _db_compat_base(self) -> Path: return run_dir / f"vr4mice_{prefix}_{date}_{attempt}" - def _fallback_output_dir(self) -> Path: base_path = self._context_processor_base_path() if base_path is not None: @@ -594,7 +586,6 @@ def _fallback_output_dir(self) -> Path: return Path.cwd() - def _mouse_from_context_or_run_dir(self, run_dir: Path) -> str: context = getattr(self, "recording_context", {}) or {} @@ -610,7 +601,6 @@ def _mouse_from_context_or_run_dir(self, run_dir: Path) -> str: return "Mouse" - def _date_from_context_or_run_dir(self, run_dir: Path) -> str: context = getattr(self, "recording_context", {}) or {} @@ -630,7 +620,6 @@ def _date_from_context_or_run_dir(self, run_dir: Path) -> str: except Exception: return datetime.now().strftime("%Y-%m-%d") - def _attempt_from_context(self, default: str = "1") -> str: context = getattr(self, "recording_context", {}) or {} @@ -648,7 +637,6 @@ def _attempt_from_context(self, default: str = "1") -> str: return default - @staticmethod def _sanitize(value: str) -> str: value = str(value).strip() @@ -656,7 +644,6 @@ def _sanitize(value: str) -> str: value = value.replace("_", "") return value or "unknown" - @staticmethod def _normalize_date(value: str) -> str | None: value = str(value) @@ -673,7 +660,6 @@ def _normalize_date(value: str) -> str | None: return None - def _date_from_run_dir_name(self, run_name: str) -> str | None: return self._normalize_date(run_name) @@ -724,7 +710,9 @@ def _strip_timestamp_json_suffix(name: str) -> str: return Path(name).stem def _find_video_files(self) -> list[Path]: - paths = self._paths_from_context_value(self.recording_context.get("video_files")) + 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) @@ -804,78 +792,17 @@ def _clear_legacy_pose_buffers(self) -> None: except Exception: logger.warning("Failed to clear legacy pose buffers after recording stop") - # ------------------------------------------------------------------ - # Cleanup - # ------------------------------------------------------------------ - - def stop(self, save: bool = False, file: str | Path | None = None) -> None: - """Cleanly stop processor resources.""" - if save: - try: - self.save(file) - except Exception: - logger.exception("Processor save during stop failed") - - self._close_teensy() - self._close_socket_connection() - self._close_listener() - - def close(self) -> None: - """Alias for generic cleanup.""" - self.stop(save=False) - - 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: - logger.exception("Failed to close Teensy cleanly") - finally: - try: - self.teensy = None - except Exception: - pass - - def _close_socket_connection(self) -> None: - try: - conn = getattr(self, "conn", None) - if conn is not None: - conn.close() - except Exception: - logger.exception("Failed to close processor socket connection") - finally: - try: - self.conn = None - except Exception: - pass - - def _close_listener(self) -> None: - try: - listener = getattr(self, "listener", None) - if listener is not None: - listener.close() - except Exception: - logger.exception("Failed to close processor listener") - finally: - try: - self.listener = None - except Exception: - pass - 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", ""), + "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", {}), } - } \ No newline at end of file + } From 372e67d9dfa6a0f2c6f69db19fbc5d88df559315 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?C=C3=A9lia=20Benquet?= <32598028+CeliaBenquet@users.noreply.github.com> Date: Fri, 24 Jul 2026 13:04:30 +0200 Subject: [PATCH 20/26] Improve compatibility to multiple cameras for autocompletion of the paths in transfer gui (#327) * Make processors compatible for dlclive gui update * Update dlcliveonly * Add HEAD_CONF_THRESHOLD to MyProcessor_socket and dlc_inference_w_pd_sync classes * Harden PD sync processor initialization cleanup Wrap processor initialization in a guarded try/except so partial startup failures trigger cleanup and re-raise as a clear RuntimeError. Add a dedicated stop() method (plus close() alias) to optionally save output and reliably release Teensy, socket connection, and listener resources, reducing leaked serial handles/ports after errors or shutdown. * Add default-path save for PD sync processor Implement a dedicated `save()` method in `dlc_inference_w_pd_sync` that persists latency data to a pickle file, supports an optional explicit path, and falls back to `self.save_path` when no file is passed. The change also initializes `save_path` on construction, ensures parent directories are created, and adds warning-based error handling for missing paths or save failures. * Enhance dlc_inference_w_pd_sync with legacy recording support and timestamp handling * Refactor dlc_inference_w_pd_sync for improved legacy support and enhanced logging * Refactor dlc_inference_w_pd_sync for DB compatibility and improved timestamp handling * Use direct datetime import in DLC sync Replaced `import datetime` with `from datetime import datetime` in `dlc_processor_socket_pd_sync.py` to align the import with direct `datetime` usage and avoid module/class ambiguity. * Use video prefix in DB compat base path Add a `_video_prefix()` helper that derives the base name from discovered video files, falling back to `filename_stem` and then `recording`. Update `_db_compat_base()` to use this prefix instead of the mouse identifier when building the output path, improving alignment with recording/video naming. * Update dlc_processor_socket_pd_sync.py * Update dlc_processor_socket_pd_sync.py * Enhance multi-camera support by extracting camera index from filenames and updating related file fetching logic * Implement review comments --------- Co-authored-by: Cyril Achard --- dj_pipeline/gui_transfer/modules/transfer.py | 12 +- .../gui_transfer/utils/session_files.py | 57 +++++- tests/unit/test_gui_transfer.py | 172 ++++++++++++++++++ 3 files changed, 236 insertions(+), 5 deletions(-) diff --git a/dj_pipeline/gui_transfer/modules/transfer.py b/dj_pipeline/gui_transfer/modules/transfer.py index f5d661398..093c16dfc 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, ) @@ -402,16 +403,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/tests/unit/test_gui_transfer.py b/tests/unit/test_gui_transfer.py index 6ca35585f..51a474402 100644 --- a/tests/unit/test_gui_transfer.py +++ b/tests/unit/test_gui_transfer.py @@ -265,6 +265,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"} From 9f4b1e2d4e6285757b05b2903bec6e7eea0b1ad1 Mon Sep 17 00:00:00 2001 From: Jaap de Ruyter van Steveninck <32810691+deruyter92@users.noreply.github.com> Date: Tue, 4 Aug 2026 10:45:18 +0200 Subject: [PATCH 21/26] Fix thread crash in `TeensyLatency` + prevent data-loss for `dlc_inference_w_pd_sync` (#329) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix `TeensyLatency`: join reader thread before closing serial port `close_serial()` was calling `self.ser.close()` immediately after signalling `_stop_reading()`, without waiting for the `read_on_thread` to exit its loop. * fix dlc_inference_w_pd_sync: save before cleanup when stopping active recording. it only saved data through the `on_recording_stopped` hook, which runs after the GUI's recording manager finalizes video files. This commit overrides `stop()` to save all three legacy outputs and clear buffers before cleanup. The save is guarded so it only runs when a recording was active and a `save_path` has been set. (the crash path, not the happy path) * add calls to renaming function in `dlc_inference_w_pd_sync.stop()` These were present in `on_recording_stopped` but not in `stop` * add centralized `_save_legacy_outputs` helper to deduplicate `on_recording_stopped` and the `stop` override method. * add missing timeout to `serial.Serial` call * TeensyLatency: skip empty reads to avoid wasted decode/parse cycles when idle --------- Co-authored-by: Célia Benquet <32598028+CeliaBenquet@users.noreply.github.com> --- .../dlc_utils/dlc_processor_socket_pd_sync.py | 43 +++++++++++++------ .../Teensy_latency/TeensyLatency.py | 16 ++++--- 2 files changed, 39 insertions(+), 20 deletions(-) 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 325fb588a..2fd3146b3 100644 --- a/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py +++ b/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py @@ -232,6 +232,24 @@ def on_recording_started(self, context: dict) -> None: 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 {}) @@ -239,23 +257,20 @@ def on_recording_stopped(self, context: dict) -> None: self.recording_context = previous_context self._legacy_recording_active = False + self._save_legacy_outputs() - proc_result = self.save() - logger.info("Processor legacy PROC save result: %r", proc_result) + def stop(self, save: bool = False, file=None): + """Save all buffered data before tearing down resources. - dlc_h5_result = self.save_legacy_dlc_h5() - logger.info("Processor legacy DLC h5 save result: %r", dlc_h5_result) - - npy_ts_result = self.save_legacy_timestamp_npy() - logger.info("Processor legacy timestamp npy save result: %r", npy_ts_result) - - video_copy_result = self.copy_legacy_video_files() - logger.info("Processor legacy video copy result: %r", video_copy_result) - - align_result = self.copy_processor_outputs_to_primary_legacy_base() - logger.info("Processor legacy output alignment result: %r", align_result) + 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 - self._clear_legacy_pose_buffers() + super().stop(save=False, file=file) # ------------------------------------------------------------------ # Primary PROC save diff --git a/mouse_task/latency_tests/Teensy_latency/TeensyLatency.py b/mouse_task/latency_tests/Teensy_latency/TeensyLatency.py index 07a0993c0..a5792316b 100644 --- a/mouse_task/latency_tests/Teensy_latency/TeensyLatency.py +++ b/mouse_task/latency_tests/Teensy_latency/TeensyLatency.py @@ -23,7 +23,10 @@ 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() + line_bytes = self.ser.readline() + if not line_bytes: + continue + line = line_bytes.decode("utf-8").rstrip() now = time.time() # Current time try: self._handle_line(line, now) @@ -32,15 +35,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() - self.ser.close() + if hasattr(self, '_reader_thread') and self._reader_thread.is_alive(): + self._reader_thread.join(timeout=2.0) + self.ser.close() \ No newline at end of file From 205d3902c664cfcf918f69c4a8265280f5befe16 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?C=C3=A9lia=20Benquet?= Date: Tue, 4 Aug 2026 11:30:47 +0200 Subject: [PATCH 22/26] Add fallback for dlclivegui imports in DLC processor modules --- mouse_task/dlc_utils/__init__.py | 40 ++++++++++++++++--- mouse_task/dlc_utils/dlcProcessor_dlconly.py | 10 ++++- mouse_task/dlc_utils/dlc_processor_socket.py | 10 ++++- .../dlc_utils/dlc_processor_socket_pd.py | 10 ++++- .../dlc_utils/dlc_processor_socket_pd_sync.py | 9 ++++- mouse_task/dlc_utils/simple_processor.py | 10 ++++- 6 files changed, 78 insertions(+), 11 deletions(-) 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 4c5d8f4ab..5e9853557 100644 --- a/mouse_task/dlc_utils/dlcProcessor_dlconly.py +++ b/mouse_task/dlc_utils/dlcProcessor_dlconly.py @@ -1,6 +1,14 @@ import numpy as np from dlclive.processor.processor import Processor -from dlclivegui.processors import PROCESSOR_REGISTRY, register_processor + +try: + from dlclivegui.processors import PROCESSOR_REGISTRY, register_processor +except ModuleNotFoundError: + PROCESSOR_REGISTRY = {} + + def register_processor(cls): + """No-op fallback when dlclivegui is not installed.""" + return cls from math import sqrt, acos, atan2, copysign, degrees import pickle diff --git a/mouse_task/dlc_utils/dlc_processor_socket.py b/mouse_task/dlc_utils/dlc_processor_socket.py index 7285ea6d5..8d9b0b34e 100644 --- a/mouse_task/dlc_utils/dlc_processor_socket.py +++ b/mouse_task/dlc_utils/dlc_processor_socket.py @@ -11,7 +11,15 @@ import numpy as np from numpy.typing import NDArray -from dlclivegui.processors import PROCESSOR_REGISTRY, register_processor # type: ignore[import-not-found] + +try: + from dlclivegui.processors import PROCESSOR_REGISTRY, register_processor # type: ignore[import-not-found] +except ModuleNotFoundError: + PROCESSOR_REGISTRY = {} + + def register_processor(cls): + """No-op fallback when dlclivegui is not installed.""" + return cls try: from dlc_utils.processor_with_signal import ProcessorWithSignal diff --git a/mouse_task/dlc_utils/dlc_processor_socket_pd.py b/mouse_task/dlc_utils/dlc_processor_socket_pd.py index 273ed9d0a..6ac72236a 100644 --- a/mouse_task/dlc_utils/dlc_processor_socket_pd.py +++ b/mouse_task/dlc_utils/dlc_processor_socket_pd.py @@ -5,7 +5,15 @@ from pathlib import Path from typing import Any, Dict -from dlclivegui.processors import PROCESSOR_REGISTRY, register_processor + +try: + from dlclivegui.processors import PROCESSOR_REGISTRY, register_processor +except ModuleNotFoundError: + PROCESSOR_REGISTRY = {} + + def register_processor(cls): + """No-op fallback when dlclivegui is not installed.""" + return cls try: from latency_tests.Teensy_latency.TeensyLatency import TeensyLatency 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 2fd3146b3..971c142df 100644 --- a/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py +++ b/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py @@ -19,7 +19,14 @@ import pandas as pd from numpy.typing import NDArray -from dlclivegui.processors import PROCESSOR_REGISTRY, register_processor # type: ignore[import-not-found] +try: + from dlclivegui.processors import PROCESSOR_REGISTRY, register_processor # type: ignore[import-not-found] +except ModuleNotFoundError: + PROCESSOR_REGISTRY = {} + + def register_processor(cls): + """No-op fallback when dlclivegui is not installed.""" + return cls try: diff --git a/mouse_task/dlc_utils/simple_processor.py b/mouse_task/dlc_utils/simple_processor.py index 8bf197862..137bb4eb3 100644 --- a/mouse_task/dlc_utils/simple_processor.py +++ b/mouse_task/dlc_utils/simple_processor.py @@ -1,7 +1,15 @@ from dlclive.processor.processor import Processor import pickle import time -from dlclivegui.processors import PROCESSOR_REGISTRY, register_processor # type: ignore[import-not-found] + +try: + from dlclivegui.processors import PROCESSOR_REGISTRY, register_processor # type: ignore[import-not-found] +except ModuleNotFoundError: + PROCESSOR_REGISTRY = {} + + def register_processor(cls): + """No-op fallback when dlclivegui is not installed.""" + return cls PROCESSOR_REGISTRY.pop("TeensyLaser", None) From 531a914e1a8504921f246b73ac982fe5cbdec34d Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?C=C3=A9lia=20Benquet?= <32598028+CeliaBenquet@users.noreply.github.com> Date: Wed, 5 Aug 2026 15:13:49 +0200 Subject: [PATCH 23/26] Improve gui transfer docs and fix move processed files to processed folder (#337) * fix `TeensyLatency`: join reader thread before closing serial port `close_serial()` was calling `self.ser.close()` immediately after signalling `_stop_reading()`, without waiting for the `read_on_thread` to exit its loop. * fix dlc_inference_w_pd_sync: save before cleanup when stopping active recording. it only saved data through the `on_recording_stopped` hook, which runs after the GUI's recording manager finalizes video files. This commit overrides `stop()` to save all three legacy outputs and clear buffers before cleanup. The save is guarded so it only runs when a recording was active and a `save_path` has been set. (the crash path, not the happy path) * add calls to renaming function in `dlc_inference_w_pd_sync.stop()` These were present in `on_recording_stopped` but not in `stop` * add centralized `_save_legacy_outputs` helper to deduplicate `on_recording_stopped` and the `stop` override method. * add missing timeout to `serial.Serial` call * TeensyLatency: skip empty reads to avoid wasted decode/parse cycles when idle * Improve gui transfer of processed files * Clarify documentation for remote path handling and improve conditional checks in Transfer class --------- Co-authored-by: Jaap de Ruyter van Steveninck <32810691+deruyter92@users.noreply.github.com> --- dj_pipeline/gui_transfer/README.md | 28 +++++++- dj_pipeline/gui_transfer/modules/transfer.py | 20 ++++-- docs/software/install_dj_pipeline.md | 6 ++ tests/unit/test_gui_transfer.py | 75 ++++++++++++++++++++ 4 files changed, 122 insertions(+), 7 deletions(-) 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 093c16dfc..2de8f1be7 100644 --- a/dj_pipeline/gui_transfer/modules/transfer.py +++ b/dj_pipeline/gui_transfer/modules/transfer.py @@ -90,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. @@ -212,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() @@ -227,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 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/tests/unit/test_gui_transfer.py b/tests/unit/test_gui_transfer.py index 51a474402..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" From c52d6d1132a7efcd44e57940ef60f1a5ca561edb Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?C=C3=A9lia=20Benquet?= <32598028+CeliaBenquet@users.noreply.github.com> Date: Fri, 7 Aug 2026 10:22:17 +0200 Subject: [PATCH 24/26] Enhance Teensy serial handling and refactor code for clarity (#330) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * Make processors compatible for dlclive gui update * Update dlcliveonly * Add HEAD_CONF_THRESHOLD to MyProcessor_socket and dlc_inference_w_pd_sync classes * Harden PD sync processor initialization cleanup Wrap processor initialization in a guarded try/except so partial startup failures trigger cleanup and re-raise as a clear RuntimeError. Add a dedicated stop() method (plus close() alias) to optionally save output and reliably release Teensy, socket connection, and listener resources, reducing leaked serial handles/ports after errors or shutdown. * Add default-path save for PD sync processor Implement a dedicated `save()` method in `dlc_inference_w_pd_sync` that persists latency data to a pickle file, supports an optional explicit path, and falls back to `self.save_path` when no file is passed. The change also initializes `save_path` on construction, ensures parent directories are created, and adds warning-based error handling for missing paths or save failures. * Enhance dlc_inference_w_pd_sync with legacy recording support and timestamp handling * Refactor dlc_inference_w_pd_sync for improved legacy support and enhanced logging * Refactor dlc_inference_w_pd_sync for DB compatibility and improved timestamp handling * Use direct datetime import in DLC sync Replaced `import datetime` with `from datetime import datetime` in `dlc_processor_socket_pd_sync.py` to align the import with direct `datetime` usage and avoid module/class ambiguity. * Use video prefix in DB compat base path Add a `_video_prefix()` helper that derives the base name from discovered video files, falling back to `filename_stem` and then `recording`. Update `_db_compat_base()` to use this prefix instead of the mouse identifier when building the output path, improving alignment with recording/video naming. * Update dlc_processor_socket_pd_sync.py * Update dlc_processor_socket_pd_sync.py * Enhance multi-camera support by extracting camera index from filenames and updating related file fetching logic * Improve Teensy serial reading with error handling and timeout * Refactor code for improved readability and error handling across multiple files * Remove autosave * Add warning on new session launching in vr4mice gui * Remove tmp files step * Add experiment lifecycle documentation to the developer guide * Copy PR#329 implementation * Make ctrl c safe if data not saved * Refactor DLC client to ensure proper socket closure * Udpate docs * Add tests for Teensy experiment GUI close behavior and latency exception handling * Add unit tests for UnityTask to simulate environment interactions * Refactor DLCClient to use socket connections with timeout handling and add regression tests for Teensy latency exceptions * Restart warning blocks action --------- Signed-off-by: Célia Benquet <32598028+CeliaBenquet@users.noreply.github.com> Co-authored-by: Cyril Achard --- _toc.yml | 1 + .../gui_transfer/utils/session_files.py | 4 +- docs/software_package/experiment_lifecycle.md | 112 ++++++++ mouse_task/dlc_utils/dlcProcessor_dlconly.py | 21 +- .../dlc_utils/dlc_processor_socket_pd_sync.py | 9 +- mouse_task/dlc_utils/simple_processor.py | 5 +- .../Teensy_latency/TeensyLatency.py | 20 +- mouse_task/task_active_sensing.py | 8 +- teensyexp/tasks_abc/dlc_deque_socket.py | 81 +++++- teensyexp/tasks_abc/dlc_socket.py | 106 ++++++-- teensyexp/teensy.py | 23 +- teensyexp/teensy_experiment.py | 126 ++++++--- tests/unit/test_dlc_socket_behavior.py | 249 ++++++++++++++++++ tests/unit/test_teensy_behavior.py | 181 +++++++++++++ tests/unit/test_unity_task_unit.py | 152 +++++++++++ 15 files changed, 998 insertions(+), 100 deletions(-) create mode 100644 docs/software_package/experiment_lifecycle.md create mode 100644 tests/unit/test_dlc_socket_behavior.py create mode 100644 tests/unit/test_teensy_behavior.py create mode 100644 tests/unit/test_unity_task_unit.py 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/utils/session_files.py b/dj_pipeline/gui_transfer/utils/session_files.py index 5eb191167..c463c32d3 100644 --- a/dj_pipeline/gui_transfer/utils/session_files.py +++ b/dj_pipeline/gui_transfer/utils/session_files.py @@ -158,7 +158,9 @@ def find_related_files(dataset_stem, path_by_key, get_type_fn, camera_number=Non ): continue file_key = get_type_fn(filepath.name) - candidates.setdefault(file_key, []).append((file_camera_number, filepath)) + candidates.setdefault(file_key, []).append( + (file_camera_number, filepath) + ) found = {} for file_key, matches in candidates.items(): 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/dlcProcessor_dlconly.py b/mouse_task/dlc_utils/dlcProcessor_dlconly.py index 5e9853557..afd197e86 100644 --- a/mouse_task/dlc_utils/dlcProcessor_dlconly.py +++ b/mouse_task/dlc_utils/dlcProcessor_dlconly.py @@ -32,20 +32,20 @@ class dlc_only(Processor): }, } - def __init__(self, con = 50, com=2): + 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: @@ -58,15 +58,14 @@ 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") @@ -84,7 +83,3 @@ def get_available_processors(): "params": getattr(dlc_only, "PROCESSOR_PARAMS", {}), } } - - - - \ No newline at end of file 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 971c142df..cfc6353c9 100644 --- a/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py +++ b/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py @@ -309,9 +309,10 @@ def save(self, file: str | Path | None = None) -> int: 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(self.save_latency_data(), f) + pickle.dump(save_dict, f) logger.info("Processor data saved to: %s", target) return 1 @@ -397,6 +398,12 @@ def _get_bodyparts_for_pose_width(self, flat_width: int) -> list[str] | None: # ------------------------------------------------------------------ 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: diff --git a/mouse_task/dlc_utils/simple_processor.py b/mouse_task/dlc_utils/simple_processor.py index 137bb4eb3..131358dd1 100644 --- a/mouse_task/dlc_utils/simple_processor.py +++ b/mouse_task/dlc_utils/simple_processor.py @@ -31,18 +31,15 @@ class TeensyLaser(Processor): }, } - def __init__( - self, com = 50, conn=2): + 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 diff --git a/mouse_task/latency_tests/Teensy_latency/TeensyLatency.py b/mouse_task/latency_tests/Teensy_latency/TeensyLatency.py index a5792316b..9908cc3b4 100644 --- a/mouse_task/latency_tests/Teensy_latency/TeensyLatency.py +++ b/mouse_task/latency_tests/Teensy_latency/TeensyLatency.py @@ -23,10 +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_bytes = self.ser.readline() - if not line_bytes: - continue - line = line_bytes.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) @@ -47,4 +57,4 @@ 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() \ No newline at end of file + 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_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 From 285a4131f6f4afce17f178434e0cd4e9a965f1f0 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?C=C3=A9lia=20Benquet?= Date: Fri, 7 Aug 2026 15:19:00 +0200 Subject: [PATCH 25/26] Run black --- dj_pipeline/gui_transfer/utils/session_files.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/dj_pipeline/gui_transfer/utils/session_files.py b/dj_pipeline/gui_transfer/utils/session_files.py index c463c32d3..5eb191167 100644 --- a/dj_pipeline/gui_transfer/utils/session_files.py +++ b/dj_pipeline/gui_transfer/utils/session_files.py @@ -158,9 +158,7 @@ def find_related_files(dataset_stem, path_by_key, get_type_fn, camera_number=Non ): continue file_key = get_type_fn(filepath.name) - candidates.setdefault(file_key, []).append( - (file_camera_number, filepath) - ) + candidates.setdefault(file_key, []).append((file_camera_number, filepath)) found = {} for file_key, matches in candidates.items(): From 6589282de8d87f94a4b3c54680be79c42c52da65 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?C=C3=A9lia=20Benquet?= Date: Fri, 7 Aug 2026 15:32:58 +0200 Subject: [PATCH 26/26] refactor: remove fallback for PROCESSOR_REGISTRY and register_processor imports --- mouse_task/dlc_utils/dlcProcessor_dlconly.py | 9 +------ mouse_task/dlc_utils/dlc_processor_socket.py | 9 +------ .../dlc_utils/dlc_processor_socket_pd.py | 9 +------ .../dlc_utils/dlc_processor_socket_pd_sync.py | 27 ++++++++++++------- mouse_task/dlc_utils/simple_processor.py | 9 +------ 5 files changed, 22 insertions(+), 41 deletions(-) diff --git a/mouse_task/dlc_utils/dlcProcessor_dlconly.py b/mouse_task/dlc_utils/dlcProcessor_dlconly.py index afd197e86..c8b3bc335 100644 --- a/mouse_task/dlc_utils/dlcProcessor_dlconly.py +++ b/mouse_task/dlc_utils/dlcProcessor_dlconly.py @@ -1,14 +1,7 @@ import numpy as np from dlclive.processor.processor import Processor -try: - from dlclivegui.processors import PROCESSOR_REGISTRY, register_processor -except ModuleNotFoundError: - PROCESSOR_REGISTRY = {} - - def register_processor(cls): - """No-op fallback when dlclivegui is not installed.""" - return cls +from dlclivegui.processors import PROCESSOR_REGISTRY, register_processor from math import sqrt, acos, atan2, copysign, degrees import pickle diff --git a/mouse_task/dlc_utils/dlc_processor_socket.py b/mouse_task/dlc_utils/dlc_processor_socket.py index 8d9b0b34e..8c4df4817 100644 --- a/mouse_task/dlc_utils/dlc_processor_socket.py +++ b/mouse_task/dlc_utils/dlc_processor_socket.py @@ -12,14 +12,7 @@ import numpy as np from numpy.typing import NDArray -try: - from dlclivegui.processors import PROCESSOR_REGISTRY, register_processor # type: ignore[import-not-found] -except ModuleNotFoundError: - PROCESSOR_REGISTRY = {} - - def register_processor(cls): - """No-op fallback when dlclivegui is not installed.""" - return cls +from dlclivegui.processors import PROCESSOR_REGISTRY, register_processor # type: ignore[import-not-found] try: from dlc_utils.processor_with_signal import ProcessorWithSignal diff --git a/mouse_task/dlc_utils/dlc_processor_socket_pd.py b/mouse_task/dlc_utils/dlc_processor_socket_pd.py index 6ac72236a..9203f15c1 100644 --- a/mouse_task/dlc_utils/dlc_processor_socket_pd.py +++ b/mouse_task/dlc_utils/dlc_processor_socket_pd.py @@ -6,14 +6,7 @@ from typing import Any, Dict -try: - from dlclivegui.processors import PROCESSOR_REGISTRY, register_processor -except ModuleNotFoundError: - PROCESSOR_REGISTRY = {} - - def register_processor(cls): - """No-op fallback when dlclivegui is not installed.""" - return cls +from dlclivegui.processors import PROCESSOR_REGISTRY, register_processor try: from latency_tests.Teensy_latency.TeensyLatency import TeensyLatency 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 cfc6353c9..9f8ce912b 100644 --- a/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py +++ b/mouse_task/dlc_utils/dlc_processor_socket_pd_sync.py @@ -19,14 +19,7 @@ import pandas as pd from numpy.typing import NDArray -try: - from dlclivegui.processors import PROCESSOR_REGISTRY, register_processor # type: ignore[import-not-found] -except ModuleNotFoundError: - PROCESSOR_REGISTRY = {} - - def register_processor(cls): - """No-op fallback when dlclivegui is not installed.""" - return cls +from dlclivegui.processors import PROCESSOR_REGISTRY, register_processor # type: ignore[import-not-found] try: @@ -59,6 +52,17 @@ def register_processor(cls): logger = logging.getLogger(__name__) +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. @@ -602,7 +606,12 @@ def _db_compat_base(self) -> Path: date = self._date_from_context_or_run_dir(run_dir) attempt = self._attempt_from_context(default="1") - return run_dir / f"vr4mice_{prefix}_{date}_{attempt}" + return run_dir / build_session_stem( + prefix, + date, + attempt, + namespace="vr4mice", + ) def _fallback_output_dir(self) -> Path: base_path = self._context_processor_base_path() diff --git a/mouse_task/dlc_utils/simple_processor.py b/mouse_task/dlc_utils/simple_processor.py index 131358dd1..91e4f980c 100644 --- a/mouse_task/dlc_utils/simple_processor.py +++ b/mouse_task/dlc_utils/simple_processor.py @@ -2,14 +2,7 @@ import pickle import time -try: - from dlclivegui.processors import PROCESSOR_REGISTRY, register_processor # type: ignore[import-not-found] -except ModuleNotFoundError: - PROCESSOR_REGISTRY = {} - - def register_processor(cls): - """No-op fallback when dlclivegui is not installed.""" - return cls +from dlclivegui.processors import PROCESSOR_REGISTRY, register_processor # type: ignore[import-not-found] PROCESSOR_REGISTRY.pop("TeensyLaser", None)