From a083d461d19de2768cdf29b148dcc7667dde555f Mon Sep 17 00:00:00 2001 From: andrianj Date: Fri, 29 May 2026 10:13:06 +0200 Subject: [PATCH 1/7] Improve-preprocessing-layout --- pyBer/gui_preprocessing.py | 123 ++++++++++++++++++++++++++++--------- pyBer/main.py | 40 +++++++----- 2 files changed, 120 insertions(+), 43 deletions(-) diff --git a/pyBer/gui_preprocessing.py b/pyBer/gui_preprocessing.py index 8c89942..bf399e9 100644 --- a/pyBer/gui_preprocessing.py +++ b/pyBer/gui_preprocessing.py @@ -847,13 +847,41 @@ class FileQueuePanel(QtWidgets.QGroupBox): def __init__(self, parent=None) -> None: super().__init__("Data", parent) + self.setObjectName("fileQueuePanel") self._current_dir_hint: str = "" self._build_ui() def _build_ui(self) -> None: v = QtWidgets.QVBoxLayout(self) - v.setSpacing(8) - v.setContentsMargins(8, 8, 8, 8) + v.setSpacing(10) + v.setContentsMargins(10, 10, 10, 10) + self.setStyleSheet( + """ + #fileQueuePanel QLabel[class="fieldLabel"] { + color: #d7e2f2; + font-weight: 650; + padding: 0 0 2px 1px; + } + #fileQueuePanel QLabel[class="pathHint"] { + color: #a9b7cb; + padding: 2px 1px 0 1px; + } + #fileQueuePanel QGroupBox#fileQueueSelectionBox { + margin-top: 12px; + } + #fileQueuePanel QGroupBox#fileQueueSelectionBox::title { + left: 10px; + padding: 0 8px; + } + #fileQueuePanel QListWidget { + padding: 6px; + } + #fileQueuePanel QListWidget::item { + min-height: 22px; + padding: 4px 6px; + } + """ + ) # Top actions top_row = QtWidgets.QHBoxLayout() @@ -867,9 +895,10 @@ def _build_ui(self) -> None: top_row.addWidget(self.btn_folder) # File list fills available height - self.list_files = PlaceholderListWidget("Drop files here or click Open File") + self.list_files = PlaceholderListWidget("Drop Doric/HDF5/CSV files here\nor click Open File") self.list_files.setSelectionMode(QtWidgets.QAbstractItemView.SelectionMode.ExtendedSelection) - self.list_files.setMinimumHeight(180) + self.list_files.setMinimumHeight(210) + self.list_files.setUniformItemSizes(True) self.btn_remove_file = QtWidgets.QPushButton("Remove selected") self.btn_remove_file.setProperty("class", "blueSecondarySmall") @@ -879,46 +908,69 @@ def _build_ui(self) -> None: # Selection block self.grp_sel = QtWidgets.QGroupBox("Selection") - form = QtWidgets.QGridLayout(self.grp_sel) - form.setContentsMargins(8, 8, 8, 8) - form.setHorizontalSpacing(6) - form.setVerticalSpacing(6) + self.grp_sel.setObjectName("fileQueueSelectionBox") + form = QtWidgets.QVBoxLayout(self.grp_sel) + form.setContentsMargins(12, 14, 12, 12) + form.setSpacing(8) self.combo_channel = QtWidgets.QComboBox() - self.combo_channel.setMinimumWidth(60) - _compact_combo(self.combo_channel, min_chars=6) + self.combo_channel.setMinimumWidth(180) + _compact_combo(self.combo_channel, min_chars=18) self.combo_trigger = QtWidgets.QComboBox() - self.combo_trigger.setMinimumWidth(60) - _compact_combo(self.combo_trigger, min_chars=6) + self.combo_trigger.setMinimumWidth(180) + _compact_combo(self.combo_trigger, min_chars=18) self.combo_trigger.addItem("") self.edit_time_start = QtWidgets.QLineEdit() self.edit_time_end = QtWidgets.QLineEdit() for ed in (self.edit_time_start, self.edit_time_end): - ed.setPlaceholderText("Start (s)" if ed is self.edit_time_start else "End (s)") + ed.setPlaceholderText("Start" if ed is self.edit_time_start else "End") val = QtGui.QDoubleValidator(0.0, 1e9, 3, ed) val.setLocale(_system_locale()) ed.setValidator(val) + ed.setMinimumWidth(82) ed.setSizePolicy(QtWidgets.QSizePolicy.Policy.Expanding, QtWidgets.QSizePolicy.Policy.Fixed) - form.addWidget(QtWidgets.QLabel("Channel"), 0, 0) - form.addWidget(self.combo_channel, 0, 1, 1, 3) - form.addWidget(QtWidgets.QLabel("Analog/Digital channel"), 1, 0) - form.addWidget(self.combo_trigger, 1, 1, 1, 3) - form.addWidget(QtWidgets.QLabel("Time window"), 2, 0) - form.addWidget(self.edit_time_start, 2, 1) - form.addWidget(QtWidgets.QLabel("to"), 2, 2) - form.addWidget(self.edit_time_end, 2, 3) + self.combo_channel.setToolTip("Signal channel to preprocess.") + self.combo_trigger.setToolTip("Optional trigger channel used for overlay and export alignment.") + self.edit_time_start.setToolTip("Optional window start in seconds.") + self.edit_time_end.setToolTip("Optional window end in seconds.") + + def _field(label_text: str, field: QtWidgets.QWidget) -> QtWidgets.QWidget: + box = QtWidgets.QWidget() + lay = QtWidgets.QVBoxLayout(box) + lay.setContentsMargins(0, 0, 0, 0) + lay.setSpacing(3) + lab = QtWidgets.QLabel(label_text) + lab.setProperty("class", "fieldLabel") + lay.addWidget(lab) + lay.addWidget(field) + return box + + time_row = QtWidgets.QWidget() + time_lay = QtWidgets.QHBoxLayout(time_row) + time_lay.setContentsMargins(0, 0, 0, 0) + time_lay.setSpacing(6) + to_label = QtWidgets.QLabel("to") + to_label.setAlignment(QtCore.Qt.AlignmentFlag.AlignCenter) + to_label.setMinimumWidth(20) + time_lay.addWidget(self.edit_time_start, 1) + time_lay.addWidget(to_label, 0) + time_lay.addWidget(self.edit_time_end, 1) + + form.addWidget(_field("Signal channel", self.combo_channel)) + form.addWidget(_field("Trigger channel", self.combo_trigger)) + form.addWidget(_field("Time window (s)", time_row)) self.btn_cutting = QtWidgets.QPushButton("Cutting / Sectioning") self.btn_cutting.setProperty("class", "blueSecondarySmall") self.btn_cutting.setSizePolicy(QtWidgets.QSizePolicy.Policy.Expanding, QtWidgets.QSizePolicy.Policy.Fixed) - form.addWidget(self.btn_cutting, 3, 0, 1, 4) - form.setColumnStretch(1, 1) - form.setColumnStretch(3, 1) + form.addWidget(self.btn_cutting) self.lbl_hint = QtWidgets.QLabel("") - self.lbl_hint.setProperty("class", "hint") + self.lbl_hint.setProperty("class", "pathHint") + self.lbl_hint.setWordWrap(True) + self.lbl_hint.setMaximumHeight(42) self.lbl_hint.setTextInteractionFlags(QtCore.Qt.TextInteractionFlag.TextSelectableByMouse) v.addLayout(top_row) @@ -966,21 +1018,36 @@ def _build_ui(self) -> None: self.btn_qc_batch.clicked.connect(self.batchQcRequested.emit) def set_path_hint(self, text: str) -> None: - self.lbl_hint.setText(text) + self.lbl_hint.setText(self._format_path_hint(text)) + self.lbl_hint.setToolTip(text or "") if text and os.path.isdir(text): self._current_dir_hint = text def path_hint(self) -> str: - return self.lbl_hint.text() + return self.lbl_hint.toolTip() or self.lbl_hint.text() def set_current_dir_hint(self, dir_path: str) -> None: self._current_dir_hint = dir_path or "" if dir_path: - self.lbl_hint.setText(dir_path) + self.lbl_hint.setText(self._format_path_hint(dir_path)) + self.lbl_hint.setToolTip(dir_path) def current_dir_hint(self) -> str: return self._current_dir_hint + def _format_path_hint(self, path: str) -> str: + text = str(path or "").strip() + if not text: + return "" + try: + parent = os.path.basename(os.path.dirname(text)) + name = os.path.basename(text) + if name and parent: + return f"Folder: {parent}/{name}" + return f"Folder: {name or text}" + except Exception: + return text + def add_file(self, path: str) -> None: item = QtWidgets.QListWidgetItem(os.path.basename(path)) item.setToolTip(path) diff --git a/pyBer/main.py b/pyBer/main.py index 8528699..0e81c8a 100644 --- a/pyBer/main.py +++ b/pyBer/main.py @@ -1470,7 +1470,7 @@ def _build_ui(self) -> None: self.artifact_panel.installEventFilter(self) self.addDockWidget(QtCore.Qt.DockWidgetArea.LeftDockWidgetArea, self.art_dock) - # Left pane: data browser + # Data browser: mounted immediately to the right of the toolbar rail. self.file_panel.setMinimumWidth(260) self.file_panel.setMaximumWidth(340) self.file_panel.setSizePolicy(QtWidgets.QSizePolicy.Policy.Fixed, QtWidgets.QSizePolicy.Policy.Expanding) @@ -1669,24 +1669,25 @@ def _build_ui(self) -> None: self._pre_drawer_splitter.setStretchFactor(0, 0) self._pre_drawer_splitter.setStretchFactor(1, 1) self._pre_drawer_splitter.setSizes([0, 1400]) - center_h.addWidget(self._pre_drawer_splitter, stretch=1) + content_widget = self._pre_drawer_splitter else: - center_h.addWidget(center_panel, stretch=1) + content_widget = center_panel - # Main splitter: data panel + visuals. Parameter popups are floating by default. + # Main splitter: data browser + visuals, both to the right of the toolbar rail. self.pre_splitter = QtWidgets.QSplitter(QtCore.Qt.Orientation.Horizontal) self.pre_splitter.setObjectName("preprocessing_splitter") self.pre_splitter.addWidget(self.file_panel) - self.pre_splitter.addWidget(center_widget) + self.pre_splitter.addWidget(content_widget) self.pre_splitter.setChildrenCollapsible(False) self.pre_splitter.setStretchFactor(0, 0) self.pre_splitter.setStretchFactor(1, 1) self.pre_splitter.setSizes([350, 1350]) self.pre_splitter.splitterMoved.connect(self._save_splitter_sizes) + center_h.addWidget(self.pre_splitter, stretch=1) pre_layout = QtWidgets.QVBoxLayout(self.pre_tab) pre_layout.setContentsMargins(10, 10, 10, 10) - pre_layout.addWidget(self.pre_splitter) + pre_layout.addWidget(center_widget) # Postprocessing tab self.post_tab = PostProcessingPanel() @@ -3997,7 +3998,7 @@ def _restore_settings(self) -> None: if self._force_fixed_dock_layouts: # Fixed mode: always enforce deterministic defaults. try: - self.pre_splitter.setSizes([300, 1200]) + self._set_pre_splitter_sizes(data_width=300, center_width=1200) except Exception: pass try: @@ -4008,15 +4009,15 @@ def _restore_settings(self) -> None: vals = [int(x) for x in splitter_sizes] if self._use_pg_dockarea_pre_layout: if len(vals) >= 3: - self.pre_splitter.setSizes([vals[0], max(640, vals[1] + vals[2])]) + self._set_pre_splitter_sizes(vals[0], max(640, vals[1] + vals[2])) elif len(vals) == 2: - self.pre_splitter.setSizes(vals[:2]) + self._set_pre_splitter_sizes(vals[0], vals[1]) elif len(vals) >= 3: left = max(260, vals[0]) center = max(640, vals[1] + vals[2]) - self.pre_splitter.setSizes([left, center]) + self._set_pre_splitter_sizes(left, center) elif len(vals) == 2: - self.pre_splitter.setSizes(vals[:2]) + self._set_pre_splitter_sizes(vals[0], vals[1]) except Exception: pass try: @@ -4045,16 +4046,16 @@ def _restore_settings(self) -> None: vals = [int(x) for x in splitter_sizes] if self._use_pg_dockarea_pre_layout: if len(vals) >= 3: - self.pre_splitter.setSizes([vals[0], max(640, vals[1] + vals[2])]) + self._set_pre_splitter_sizes(vals[0], max(640, vals[1] + vals[2])) elif len(vals) == 2: - self.pre_splitter.setSizes(vals[:2]) + self._set_pre_splitter_sizes(vals[0], vals[1]) elif len(vals) >= 3: # Migrate old 3-pane [left, center, right] into [left, center+right]. left = max(260, vals[0]) center = max(640, vals[1] + vals[2]) - self.pre_splitter.setSizes([left, center]) + self._set_pre_splitter_sizes(left, center) elif len(vals) == 2: - self.pre_splitter.setSizes(vals[:2]) + self._set_pre_splitter_sizes(vals[0], vals[1]) except Exception: pass @@ -4136,6 +4137,15 @@ def _save_settings(self) -> None: except Exception: pass + def _set_pre_splitter_sizes(self, data_width: int, center_width: int) -> None: + """Apply logical [data, center] sizes to the preprocessing splitter.""" + try: + data = max(0, int(data_width)) + center = max(640, int(center_width)) + self.pre_splitter.setSizes([data, center]) + except Exception: + pass + def _save_splitter_sizes(self, *_args) -> None: """Save the current splitter sizes to settings.""" try: From 8760b0cae7bac51d08186d036830d55eb411af79 Mon Sep 17 00:00:00 2001 From: andrianj Date: Fri, 29 May 2026 10:44:40 +0200 Subject: [PATCH 2/7] Refine-artifact-controls --- pyBer/gui_preprocessing.py | 33 ++++++-- pyBer/main.py | 155 ++++++++++++++++++++++++++++--------- pyBer/onboarding.py | 6 +- 3 files changed, 144 insertions(+), 50 deletions(-) diff --git a/pyBer/gui_preprocessing.py b/pyBer/gui_preprocessing.py index bf399e9..6537033 100644 --- a/pyBer/gui_preprocessing.py +++ b/pyBer/gui_preprocessing.py @@ -549,10 +549,20 @@ def __init__(self, parent=None) -> None: def _build_ui(self) -> None: layout = QtWidgets.QVBoxLayout(self) + table_min_height = 260 auto_group = QtWidgets.QGroupBox("Auto-detected (threshold)") + auto_group.setSizePolicy( + QtWidgets.QSizePolicy.Policy.Expanding, + QtWidgets.QSizePolicy.Policy.Expanding, + ) auto_layout = QtWidgets.QVBoxLayout(auto_group) self.table_auto = QtWidgets.QTableWidget(0, 6) + self.table_auto.setMinimumHeight(table_min_height) + self.table_auto.setSizePolicy( + QtWidgets.QSizePolicy.Policy.Expanding, + QtWidgets.QSizePolicy.Policy.Expanding, + ) self.table_auto.setHorizontalHeaderLabels(["ID", "Remove", "Source", "Core (s)", "Cut start", "Cut end"]) self.table_auto.horizontalHeader().setStretchLastSection(True) self.table_auto.verticalHeader().setVisible(False) @@ -563,12 +573,21 @@ def _build_ui(self) -> None: self.table_auto.setColumnWidth(2, 66) self.table_auto.setColumnWidth(3, 132) self.table_auto.setColumnWidth(4, 82) - auto_layout.addWidget(self.table_auto) - layout.addWidget(auto_group) + auto_layout.addWidget(self.table_auto, 1) + layout.addWidget(auto_group, 1) manual_group = QtWidgets.QGroupBox("Manual artifacts") + manual_group.setSizePolicy( + QtWidgets.QSizePolicy.Policy.Expanding, + QtWidgets.QSizePolicy.Policy.Expanding, + ) manual_layout = QtWidgets.QVBoxLayout(manual_group) self.table = QtWidgets.QTableWidget(0, 3) + self.table.setMinimumHeight(table_min_height) + self.table.setSizePolicy( + QtWidgets.QSizePolicy.Policy.Expanding, + QtWidgets.QSizePolicy.Policy.Expanding, + ) self.table.setHorizontalHeaderLabels(["ID", "Start (s)", "End (s)"]) self.table.horizontalHeader().setStretchLastSection(True) self.table.verticalHeader().setVisible(False) @@ -608,7 +627,7 @@ def _build_ui(self) -> None: self.btn_close = QtWidgets.QPushButton("Close") btnrow.addWidget(self.btn_close) manual_layout.addLayout(btnrow) - layout.addWidget(manual_group) + layout.addWidget(manual_group, 1) self.btn_add.clicked.connect(self._on_add) self.btn_update.clicked.connect(self._on_update_selected) @@ -1933,7 +1952,6 @@ def mk_spin(minw=60) -> QtWidgets.QSpinBox: self.btn_load_config = QtWidgets.QPushButton("Load config") self.btn_reset_defaults = QtWidgets.QPushButton("Reset defaults") for b in ( - self.btn_artifacts_panel, self.btn_qc, self.btn_qc_batch, self.btn_export, @@ -1946,6 +1964,7 @@ def mk_spin(minw=60) -> QtWidgets.QSpinBox: b.setProperty("class", "compactSmall") b.setSizePolicy(QtWidgets.QSizePolicy.Policy.Expanding, QtWidgets.QSizePolicy.Policy.Fixed) self.btn_export.setProperty("class", "compactPrimarySmall") + self.btn_artifacts_panel.setVisible(False) self.btn_save_config.clicked.connect(self._save_config) self.btn_load_config.clicked.connect(self._load_config) self.btn_reset_defaults.clicked.connect(self._reset_defaults) @@ -2026,8 +2045,7 @@ def mk_spin(minw=60) -> QtWidgets.QSpinBox: qc_grid.setHorizontalSpacing(6) qc_grid.setVerticalSpacing(6) qc_grid.addWidget(self.btn_export, 0, 0, 1, 2) - qc_grid.addWidget(self.btn_artifacts_panel, 1, 0) - qc_grid.addWidget(self.btn_advanced, 1, 1) + qc_grid.addWidget(self.btn_advanced, 1, 0, 1, 2) qc_grid.addWidget(self.btn_qc, 2, 0) qc_grid.addWidget(self.btn_qc_batch, 2, 1) qc_grid.addWidget(self.btn_metadata, 3, 0) @@ -2869,6 +2887,7 @@ def _build_ui(self) -> None: self.btn_redo.setToolTip("Redo last undone preprocessing action (Ctrl+Y)") self.btn_redo.setFixedSize(34, 30) self.btn_artifacts = QtWidgets.QPushButton("Artifacts") + self.btn_artifacts.setVisible(False) self.btn_box_select = QtWidgets.QPushButton("Box select") self.btn_box_select.setCheckable(True) self.btn_thresholds = QtWidgets.QPushButton("Thresholds: ON") @@ -2878,7 +2897,6 @@ def _build_ui(self) -> None: for b in ( self.btn_add_region, self.btn_clear_regions, - self.btn_artifacts, self.btn_box_select, self.btn_thresholds, ): @@ -2887,7 +2905,6 @@ def _build_ui(self) -> None: tools.addWidget(self.btn_clear_regions) tools.addWidget(self.btn_undo) tools.addWidget(self.btn_redo) - tools.addWidget(self.btn_artifacts) tools.addWidget(self.btn_box_select) tools.addWidget(self.btn_thresholds) tools.addStretch(1) diff --git a/pyBer/main.py b/pyBer/main.py index 0e81c8a..401203d 100644 --- a/pyBer/main.py +++ b/pyBer/main.py @@ -130,7 +130,6 @@ def _is_user_site_path(path: str) -> bool: app_qss, _make_icon, _paint_database, - _paint_list, _paint_sliders, _paint_filter, _paint_wave, @@ -207,7 +206,7 @@ def _to_bool(value: object, default: bool = False) -> bool: _POST_DOCK_PREFIX = "post." _FORCE_FIXED_DOCK_LAYOUTS = False _USE_PG_DOCKAREA_PRE_LAYOUT = True -_PRE_DOCKAREA_PRIMARY_ORDER = ("artifacts_list", "artifacts", "filtering", "baseline", "output", "export") +_PRE_DOCKAREA_PRIMARY_ORDER = ("artifacts", "filtering", "baseline", "output", "export") _PRE_DOCKAREA_OPTIONAL_ORDER = ("qc", "config") _PRE_DOCKAREA_DEFAULT_VISIBLE = frozenset(_PRE_DOCKAREA_PRIMARY_ORDER) _CSV_NONE_LABEL = "(none)" @@ -967,7 +966,14 @@ def __init__(self, qc: Dict[str, object], parent=None) -> None: r = float(qc.get("r", np.nan)) r2 = r * r if np.isfinite(r) else np.nan if np.isfinite(r): - self._add_plot_text_topleft(self.plot_corr, f"r={r:.3g} r2={r2:.3g}") + self._add_plot_text_topleft( + self.plot_corr, + f"r={r:.3g} r2={r2:.3g}", + color=(255, 213, 95), + corner="topright", + fill=(12, 16, 24, 205), + border=(255, 213, 95, 150), + ) self.plot_corr.setLabel("left", "Signal dF/F (%)") self.plot_corr.setLabel("bottom", "Isobestic dF/F (%)") else: @@ -1114,19 +1120,42 @@ def _add_filled_band( plot.addItem(upper_curve) plot.addItem(lower_curve) - def _add_plot_text_topleft(self, plot: pg.PlotWidget, text: str) -> None: + def _add_plot_text_topleft( + self, + plot: pg.PlotWidget, + text: str, + *, + color: Tuple[int, int, int] = (220, 220, 220), + corner: str = "topleft", + fill: Optional[Tuple[int, int, int, int]] = None, + border: Optional[Tuple[int, int, int, int]] = None, + ) -> None: if not text: return vb = plot.getViewBox() if not vb: return (x0, x1), (y0, y1) = vb.viewRange() - if not np.isfinite(x0) or not np.isfinite(y1): - return - pad_x = (x1 - x0) * 0.02 - pad_y = (y1 - y0) * 0.05 - item = pg.TextItem(text, color=(220, 220, 220), anchor=(0, 1)) - item.setPos(x0 + pad_x, y1 - pad_y) + if not all(np.isfinite(v) for v in (x0, x1, y0, y1)): + return + pad_x = (x1 - x0) * 0.03 + pad_y = (y1 - y0) * 0.08 + corner_norm = str(corner or "topleft").strip().lower() + if corner_norm == "topright": + anchor = (1, 1) + pos = (x1 - pad_x, y1 - pad_y) + else: + anchor = (0, 1) + pos = (x0 + pad_x, y1 - pad_y) + item = pg.TextItem( + text, + color=color, + anchor=anchor, + fill=pg.mkBrush(fill) if fill is not None else None, + border=pg.mkPen(border, width=1.0) if border is not None else None, + ) + item.setZValue(50) + item.setPos(*pos) plot.addItem(item) def _save_images(self) -> None: @@ -1530,8 +1559,7 @@ def _build_ui(self) -> None: self.btn_plot_style.setMenu(self.menu_plot_style) # Inline parameter section buttons (same row as workflow actions). - self.btn_section_artifacts_list = QtWidgets.QToolButton(); self.btn_section_artifacts_list.setText("Artifact list") - self.btn_section_artifacts = QtWidgets.QToolButton(); self.btn_section_artifacts.setText("Artifact setup") + self.btn_section_artifacts = QtWidgets.QToolButton(); self.btn_section_artifacts.setText("Artifacts") self.btn_section_filtering = QtWidgets.QToolButton(); self.btn_section_filtering.setText("Filtering") self.btn_section_baseline = QtWidgets.QToolButton(); self.btn_section_baseline.setText("Baseline") self.btn_section_output = QtWidgets.QToolButton(); self.btn_section_output.setText("Output") @@ -1539,7 +1567,6 @@ def _build_ui(self) -> None: self.btn_section_export = QtWidgets.QToolButton(); self.btn_section_export.setText("Export") self.btn_section_config = QtWidgets.QToolButton(); self.btn_section_config.setText("Configuration") self._section_buttons: Dict[str, QtWidgets.QPushButton] = { - "artifacts_list": self.btn_section_artifacts_list, "artifacts": self.btn_section_artifacts, "filtering": self.btn_section_filtering, "baseline": self.btn_section_baseline, @@ -1557,8 +1584,7 @@ def _build_ui(self) -> None: # ----- Modern shell: vertical icon rail + thin transport bar ------ # Configure section buttons as icon-only rail buttons. _rail_section_meta = { - "artifacts_list": ("Artifacts", "Detected and manual artifacts list", _paint_list), - "artifacts": ("Artifact", "Artifact detection thresholds", _paint_sliders), + "artifacts": ("Artifacts", "Detection thresholds and artifact list", _paint_sliders), "filtering": ("Filtering", "Low-pass and smoothing options", _paint_filter), "baseline": ("Baseline", "Baseline estimation across recording", _paint_wave), "output": ("Output", "Choose dFF / dF / z-score formula", _paint_chart), @@ -1603,7 +1629,7 @@ def _build_ui(self) -> None: sep.setObjectName("railSeparator") sep.setFrameShape(QtWidgets.QFrame.Shape.HLine) rail_layout.addWidget(sep) - for key in ("artifacts_list", "artifacts", "filtering", "baseline", + for key in ("artifacts", "filtering", "baseline", "output", "qc", "export", "config"): rail_layout.addWidget(self._section_buttons[key], 0, QtCore.Qt.AlignmentFlag.AlignHCenter) @@ -1859,8 +1885,7 @@ def _setup_section_popups(self) -> None: self.param_panel.card_actions.setVisible(False) section_widgets: Dict[str, QtWidgets.QWidget] = { - "artifacts_list": self.artifact_panel, - "artifacts": self.param_panel.card_artifacts, + "artifacts": self._build_artifacts_section_widget(), "filtering": self.param_panel.card_filtering, "baseline": self.param_panel.card_baseline, "output": self.param_panel.card_output, @@ -1869,8 +1894,7 @@ def _setup_section_popups(self) -> None: "config": self._build_config_actions_widget(), } section_titles: Dict[str, str] = { - "artifacts_list": "Artifact list", - "artifacts": "Artifact setup", + "artifacts": "Artifacts", "filtering": "Filtering", "baseline": "Baseline", "output": "Output", @@ -1939,6 +1963,45 @@ def _setup_section_popups(self) -> None: widget.installEventFilter(self) self._section_docks[key] = dock + def _build_artifacts_section_widget(self) -> QtWidgets.QWidget: + panel = QtWidgets.QWidget() + layout = QtWidgets.QVBoxLayout(panel) + layout.setContentsMargins(0, 0, 0, 0) + layout.setSpacing(10) + + self.param_panel.card_artifacts.setSizePolicy( + QtWidgets.QSizePolicy.Policy.Expanding, + QtWidgets.QSizePolicy.Policy.Fixed, + ) + layout.addWidget(self.param_panel.card_artifacts) + + try: + self.artifact_panel.btn_close.setVisible(False) + except Exception: + pass + table_min_heights = { + "table_auto": 260, + "table": 260, + } + for table_name, min_height in table_min_heights.items(): + try: + table = getattr(self.artifact_panel, table_name) + table.setMinimumHeight(min_height) + table.setSizePolicy( + QtWidgets.QSizePolicy.Policy.Expanding, + QtWidgets.QSizePolicy.Policy.Expanding, + ) + except Exception: + pass + self.artifact_panel.setSizePolicy( + QtWidgets.QSizePolicy.Policy.Expanding, + QtWidgets.QSizePolicy.Policy.Expanding, + ) + self.artifact_panel.show() + layout.addWidget(self.artifact_panel, 1) + + return panel + def _pre_dockarea_dock(self, key: str) -> Optional[Dock]: return self._pre_dockarea_docks.get(key) @@ -2037,7 +2100,7 @@ def _arrange_pre_dockarea_default(self) -> None: if self._pre_dockarea is None: return ordered = self._pre_dockarea_ordered_keys() - root = self._pre_dockarea_dock("artifacts_list") + root = self._pre_dockarea_dock("artifacts") if root is None and ordered: root = self._pre_dockarea_dock(ordered[0]) if root is None: @@ -2063,6 +2126,8 @@ def _set_pre_dockarea_visible(self, key: str, visible: bool) -> None: if dock is None: return if visible: + if key == "artifacts": + self.artifact_panel.show() self._arrange_pre_dockarea_default() dock.show() try: @@ -2092,8 +2157,6 @@ def _save_pre_dockarea_layout_state(self) -> None: left_i = _dock_area_to_int(QtCore.Qt.DockWidgetArea.LeftDockWidgetArea, 1) for key, dock in self._pre_dockarea_docks.items(): - if key == "artifacts_list": - continue try: base = f"pre_section_docks/{key}" self.settings.setValue(f"{base}/visible", bool(dock.isVisible())) @@ -2103,7 +2166,7 @@ def _save_pre_dockarea_layout_state(self) -> None: continue try: art_base = "pre_artifact_dock_state" - art_vis = bool(visible.get("artifacts_list", False)) + art_vis = bool(visible.get("artifacts", False)) self.settings.setValue(f"{art_base}/visible", art_vis) self.settings.setValue(f"{art_base}/floating", False) self.settings.setValue(f"{art_base}/area", left_i) @@ -2128,20 +2191,25 @@ def _restore_pre_dockarea_layout_state(self) -> None: except Exception: visible_map = {} + legacy_artifacts_visible = visible_map.pop("artifacts_list", None) + if legacy_artifacts_visible is not None and "artifacts" in self._pre_dockarea_docks: + visible_map["artifacts"] = bool(visible_map.get("artifacts", False) or legacy_artifacts_visible) + if not visible_map: for key in self._pre_dockarea_docks.keys(): - if key == "artifacts_list": - raw = self.settings.value("pre_artifact_dock_state/visible", None) - if raw is not None: - visible_map[key] = _to_bool(raw, False) - continue raw = self.settings.value(f"pre_section_docks/{key}/visible", None) if raw is not None: visible_map[key] = _to_bool(raw, False) + if key == "artifacts": + legacy_raw = self.settings.value("pre_artifact_dock_state/visible", None) + if legacy_raw is not None: + visible_map[key] = bool(visible_map.get(key, False) or _to_bool(legacy_raw, False)) if not visible_map: visible_map = self._pre_dockarea_default_visible_map() - active = str(self.settings.value(_PRE_DOCKAREA_ACTIVE_KEY, "artifacts_list") or "artifacts_list") + active = str(self.settings.value(_PRE_DOCKAREA_ACTIVE_KEY, "artifacts") or "artifacts") + if active == "artifacts_list": + active = "artifacts" if not bool(visible_map.get(active, False)): active = next((key for key in self._pre_dockarea_ordered_keys() if bool(visible_map.get(key, False))), "") @@ -2213,10 +2281,8 @@ def _build_qc_actions_widget(self) -> QtWidgets.QWidget: v.setSpacing(6) self.param_panel.btn_qc.setProperty("class", "blueSecondarySmall") self.param_panel.btn_qc_batch.setProperty("class", "blueSecondarySmall") - self.param_panel.btn_artifacts_panel.setProperty("class", "blueSecondarySmall") v.addWidget(self.param_panel.btn_qc) v.addWidget(self.param_panel.btn_qc_batch) - v.addWidget(self.param_panel.btn_artifacts_panel) v.addWidget(self.param_panel.lbl_fs) v.addStretch(1) return panel @@ -2310,8 +2376,7 @@ def _force_hide_pre_drawer_initially(self) -> None: pass _PRE_SECTION_TITLES = { - "artifacts_list": "Artifact list", - "artifacts": "Artifact setup", + "artifacts": "Artifacts", "filtering": "Filtering", "baseline": "Baseline", "output": "Output", @@ -2385,6 +2450,8 @@ def _toggle_section_popup(self, key: str, checked: bool) -> None: except Exception: pass dock.show() + if key == "artifacts": + self.artifact_panel.show() try: dock.raiseDock() except Exception: @@ -3418,7 +3485,7 @@ def _read_section_settings(prefix: str, keys: List[str]) -> Dict[str, Dict[str, return out if self._use_pg_dockarea_pre_layout and self._pre_dockarea_docks: - pre_sections = [k for k in self._pre_dockarea_docks.keys() if k != "artifacts_list"] + pre_sections = list(self._pre_dockarea_docks.keys()) else: pre_sections = list(self._section_docks.keys()) post_sections = [] @@ -3907,7 +3974,7 @@ def _has_saved_pre_layout_state(self) -> bool: return True keys = list(self._section_docks.keys()) if self._use_pg_dockarea_pre_layout and self._pre_dockarea_docks: - keys = [k for k in self._pre_dockarea_docks.keys() if k != "artifacts_list"] + keys = list(self._pre_dockarea_docks.keys()) for key in keys: if self.settings.contains(f"pre_section_docks/{key}/visible"): return True @@ -6942,22 +7009,34 @@ def _contains(target: Tuple[float, float], arr: List[Tuple[float, float]]) -> bo def _toggle_artifacts_panel(self) -> None: if self._use_pg_dockarea_pre_layout: self._setup_section_popups() - dock = self._pre_dockarea_dock("artifacts_list") + dock = self._pre_dockarea_dock("artifacts") if dock is None: return if dock.isVisible(): dock.hide() else: + self.artifact_panel.show() dock.show() try: dock.raiseDock() except Exception: pass - self._last_opened_section = "artifacts_list" + self._last_opened_section = "artifacts" self._sync_section_button_states_from_docks() self._save_panel_layout_state() return + section_dock = self._section_docks.get("artifacts") + if isinstance(section_dock, QtWidgets.QDockWidget): + if section_dock.isVisible(): + section_dock.setVisible(False) + else: + self.artifact_panel.show() + section_dock.setVisible(True) + section_dock.raise_() + self._save_panel_layout_state() + return + if isinstance(self.art_dock, QtWidgets.QDockWidget): if self.art_dock.isVisible(): self.art_dock.setVisible(False) diff --git a/pyBer/onboarding.py b/pyBer/onboarding.py index 0d10b85..281e9c3 100644 --- a/pyBer/onboarding.py +++ b/pyBer/onboarding.py @@ -602,8 +602,7 @@ def _before(_w: QtWidgets.QWidget) -> None: ] for key, title, body in [ - ("artifacts_list", "Artifact list", "Inspect detected and manual artifact windows, jump to them, and remove entries."), - ("artifacts", "Artifact setup", "Choose artifact thresholds and how artifacts are handled: interpolate, cut, low-pass locally, or leave unchanged."), + ("artifacts", "Artifacts", "Choose artifact thresholds, set artifact handling, and inspect detected or manual artifact windows."), ("filtering", "Filtering", "Set low-pass and smoothing options for signal and reference traces."), ("baseline", "Baseline", "Configure baseline estimation before computing dF/F, dF, or z-score outputs."), ("output", "Output", "Choose the processed trace formula and preview the resulting channel."), @@ -960,8 +959,7 @@ def _keyboard_cheatsheet_html() -> str: } _PRE_SECTION_META: Dict[str, Tuple[str, str, str, str]] = { - "artifacts_list": ("L", "#f5c542", "Artifact list", "Inspect and edit detected / manual artifacts"), - "artifacts": ("A", "#ee6471", "Artifact setup", "Detection thresholds and manual selection"), + "artifacts": ("A", "#ee6471", "Artifacts", "Detection thresholds and detected / manual artifact list"), "filtering": ("F", "#7d4df2", "Filtering", "Low-pass + smoothing for the photometry trace"), "baseline": ("B", "#4b9df8", "Baseline", "Baseline estimation across the recording"), "output": ("O", "#5dd39e", "Output", "Choose dFF / dF / z-score formula"), From ca48b7ea373f268a90411ea3ee8f252eddf156df Mon Sep 17 00:00:00 2001 From: andrianj Date: Fri, 29 May 2026 19:22:30 +0200 Subject: [PATCH 3/7] Add-continuous-PSTH-alignment --- pyBer/gui_postprocessing.py | 584 +++++++++++++++++++++++++++++++++++- pyBer/gui_preprocessing.py | 94 +++++- pyBer/main.py | 136 ++++++--- 3 files changed, 750 insertions(+), 64 deletions(-) diff --git a/pyBer/gui_postprocessing.py b/pyBer/gui_postprocessing.py index 7350db6..ddc3297 100644 --- a/pyBer/gui_postprocessing.py +++ b/pyBer/gui_postprocessing.py @@ -428,6 +428,319 @@ def _load_behavior_ethovision( } +_RULE_NUMBER_RE = r"[-+]?(?:\d+(?:\.\d*)?|\.\d+)(?:[eE][-+]?\d+)?" + + +def _parse_rule_float(text: str) -> Optional[float]: + raw = str(text or "").strip() + if not re.fullmatch(_RULE_NUMBER_RE, raw): + return None + try: + return float(raw) + except Exception: + return None + + +def _compare_rule_values(values: np.ndarray, op: str, threshold: float) -> np.ndarray: + if op == ">": + return values > threshold + if op == ">=": + return values >= threshold + if op == "<": + return values < threshold + if op == "<=": + return values <= threshold + raise ValueError(f"Unsupported threshold operator: {op}") + + +def _invert_rule_operator(op: str) -> str: + return {">": "<", ">=": "<=", "<": ">", "<=": ">="}.get(op, op) + + +def _continuous_rule_mask(values: np.ndarray, variable_name: str, rule_text: str) -> np.ndarray: + """Return a boolean mask for a simple threshold rule without using eval.""" + values = np.asarray(values, float) + rule = str(rule_text or "").strip() + if not rule: + raise ValueError("Enter a threshold rule.") + + and_parts = re.split(r"\s+(?:and|&&)\s+", rule, flags=re.IGNORECASE) + if len(and_parts) > 1: + mask = np.ones(values.shape, dtype=bool) + for part in and_parts: + mask &= _continuous_rule_mask(values, variable_name, part) + return mask + + or_parts = re.split(r"\s+(?:or|\|\|)\s+", rule, flags=re.IGNORECASE) + if len(or_parts) > 1: + mask = np.zeros(values.shape, dtype=bool) + for part in or_parts: + mask |= _continuous_rule_mask(values, variable_name, part) + return mask + + ops = list(re.finditer(r"(>=|<=|>|<)", rule)) + numbers = [float(m.group(0)) for m in re.finditer(_RULE_NUMBER_RE, rule)] + + if len(ops) >= 2 and len(numbers) >= 2: + # Users often type band rules as either a < x < b or a > x > b. + low, high = sorted(numbers[:2]) + inclusive = any(m.group(1) in {">=", "<="} for m in ops[:2]) + if inclusive: + return (values >= low) & (values <= high) + return (values > low) & (values < high) + + if len(ops) == 1: + match = ops[0] + left = rule[:match.start()].strip() + right = rule[match.end():].strip() + left_num = _parse_rule_float(left) + right_num = _parse_rule_float(right) + op = match.group(1) + if left_num is None and right_num is not None: + return _compare_rule_values(values, op, right_num) + if left_num is not None and right_num is None: + return _compare_rule_values(values, _invert_rule_operator(op), left_num) + if left_num is not None and right_num is not None: + return np.full(values.shape, bool(_compare_rule_values(np.asarray([left_num]), op, right_num)[0]), dtype=bool) + + if len(numbers) >= 2: + low, high = sorted(numbers[:2]) + return (values > low) & (values < high) + if len(numbers) == 1: + return values > numbers[0] + + raise ValueError(f"Could not parse threshold rule for {variable_name}.") + + +def _continuous_threshold_events( + time: np.ndarray, + values: np.ndarray, + rule_text: str, + variable_name: str, +) -> Tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]: + values = np.asarray(values, float).reshape(-1) + time = np.asarray(time, float).reshape(-1) + if time.size == 0 and values.size: + time = np.arange(values.size, dtype=float) + n = min(time.size, values.size) + if n <= 0: + return np.array([], float), np.array([], float), np.array([], float), np.array([], bool) + time = time[:n] + values = values[:n] + finite = np.isfinite(time) & np.isfinite(values) + mask = np.zeros(n, dtype=bool) + if np.any(finite): + mask[finite] = _continuous_rule_mask(values[finite], variable_name, rule_text) + prev_high = np.r_[False, mask[:-1]] + next_high = np.r_[mask[1:], False] + on_idx = np.where(mask & ~prev_high)[0] + off_idx = np.where(mask & ~next_high)[0] + on = time[on_idx] + off = time[off_idx] + m = min(on.size, off.size) + dur = off[:m] - on[:m] if m else np.array([], float) + if on.size != off.size: + dur = np.full(on.shape, np.nan, dtype=float) + else: + dur = np.maximum(dur, 0.0) + return on, off, dur, mask + + +def _continuous_behavior_name(variable_name: str, rule_text: str, align_text: str) -> str: + suffix = "offset" if str(align_text).strip().lower().endswith("offset") else "onset" + compact_rule = re.sub(r"\s+", " ", str(rule_text or "").strip()) + if len(compact_rule) > 42: + compact_rule = compact_rule[:39] + "..." + base = f"{variable_name} [{compact_rule}] {suffix}".strip() + return base or f"continuous {suffix}" + + +class ContinuousAlignDialog(QtWidgets.QDialog): + def __init__( + self, + behavior_sources: Dict[str, Dict[str, Any]], + parent: Optional[QtWidgets.QWidget] = None, + ) -> None: + super().__init__(parent) + self.setWindowTitle("Align to continuous") + self.resize(760, 560) + self._sources = { + str(k): v + for k, v in (behavior_sources or {}).items() + if isinstance(v, dict) and bool(v.get("trajectory") or {}) + } + self._rule_touched = False + self._name_touched = False + + root = QtWidgets.QVBoxLayout(self) + root.setContentsMargins(12, 12, 12, 12) + root.setSpacing(10) + + form = QtWidgets.QFormLayout() + form.setLabelAlignment(QtCore.Qt.AlignmentFlag.AlignLeft | QtCore.Qt.AlignmentFlag.AlignTop) + form.setRowWrapPolicy(QtWidgets.QFormLayout.RowWrapPolicy.WrapLongRows) + + self.combo_source = QtWidgets.QComboBox() + self.combo_variable = QtWidgets.QComboBox() + self.edit_rule = QtWidgets.QLineEdit() + self.edit_rule.setPlaceholderText("velocity > 6 or 2.1 < velocity < 5.5") + self.combo_align = QtWidgets.QComboBox() + self.combo_align.addItems(["Align to onset", "Align to offset"]) + self.edit_name = QtWidgets.QLineEdit() + self.cb_apply_all = QtWidgets.QCheckBox("Apply to all loaded files with this variable") + self.cb_apply_all.setChecked(True) + self.lbl_status = QtWidgets.QLabel("") + self.lbl_status.setProperty("class", "hint") + self.lbl_status.setWordWrap(True) + + _compact_combo(self.combo_source, min_chars=18) + _compact_combo(self.combo_variable, min_chars=18) + _compact_combo(self.combo_align, min_chars=10) + + form.addRow("Behavior file", self.combo_source) + form.addRow("Continuous variable", self.combo_variable) + form.addRow("Threshold rule", self.edit_rule) + form.addRow("Align to", self.combo_align) + form.addRow("Behavior name", self.edit_name) + form.addRow("", self.cb_apply_all) + root.addLayout(form) + + self.plot = pg.PlotWidget(title="Continuous threshold preview") + _opt_plot(self.plot) + self.plot.setMinimumHeight(260) + self.plot.setLabel("bottom", "Time (s)") + self.plot.setLabel("left", "Value") + root.addWidget(self.plot, stretch=1) + root.addWidget(self.lbl_status) + + buttons = QtWidgets.QDialogButtonBox( + QtWidgets.QDialogButtonBox.StandardButton.Ok | QtWidgets.QDialogButtonBox.StandardButton.Cancel + ) + self.btn_ok = buttons.button(QtWidgets.QDialogButtonBox.StandardButton.Ok) + if self.btn_ok is not None: + self.btn_ok.setText("Create alignment") + self.btn_ok.setProperty("class", "compactPrimary") + buttons.accepted.connect(self.accept) + buttons.rejected.connect(self.reject) + root.addWidget(buttons) + + self._populate_sources() + self.combo_source.currentIndexChanged.connect(self._populate_variables) + self.combo_source.currentIndexChanged.connect(self._update_default_rule) + self.combo_variable.currentIndexChanged.connect(self._update_default_rule) + self.combo_variable.currentIndexChanged.connect(self._update_preview) + self.combo_align.currentIndexChanged.connect(self._update_default_name) + self.combo_align.currentIndexChanged.connect(self._update_preview) + self.edit_rule.textEdited.connect(self._on_rule_edited) + self.edit_name.textEdited.connect(self._on_name_edited) + self._populate_variables() + self._update_default_rule() + self._update_preview() + + def _populate_sources(self) -> None: + self.combo_source.clear() + for stem in sorted(self._sources.keys()): + self.combo_source.addItem(stem, stem) + + def _selected_source(self) -> Tuple[str, Dict[str, Any]]: + key = str(self.combo_source.currentData() or self.combo_source.currentText() or "") + return key, self._sources.get(key, {}) + + def _populate_variables(self) -> None: + current = self.combo_variable.currentText() + self.combo_variable.blockSignals(True) + try: + self.combo_variable.clear() + _key, info = self._selected_source() + for name in sorted(str(k) for k in (info.get("trajectory") or {}).keys()): + self.combo_variable.addItem(name, name) + idx = self.combo_variable.findText(current) + if idx >= 0: + self.combo_variable.setCurrentIndex(idx) + finally: + self.combo_variable.blockSignals(False) + self._update_preview() + + def _selected_variable_data(self) -> Tuple[str, np.ndarray, np.ndarray]: + _key, info = self._selected_source() + variable = str(self.combo_variable.currentData() or self.combo_variable.currentText() or "") + trajectory = info.get("trajectory") or {} + values = np.asarray(trajectory.get(variable, np.array([], float)), float) + time = np.asarray(info.get("trajectory_time", np.array([], float)), float) + if time.size == 0 and values.size: + time = np.arange(values.size, dtype=float) + return variable, time, values + + def _update_default_rule(self) -> None: + variable, _time, values = self._selected_variable_data() + if not variable: + return + if not self._rule_touched or not self.edit_rule.text().strip(): + finite = values[np.isfinite(values)] + threshold = float(np.nanmedian(finite)) if finite.size else 0.0 + self.edit_rule.setText(f"{variable} > {threshold:.4g}") + self._update_default_name() + self._update_preview() + + def _update_default_name(self) -> None: + variable = str(self.combo_variable.currentData() or self.combo_variable.currentText() or "continuous") + if not self.edit_name.text().strip() or not getattr(self, "_name_touched", False): + self.edit_name.setText(_continuous_behavior_name(variable, self.edit_rule.text(), self.combo_align.currentText())) + + def _on_rule_edited(self, _text: str) -> None: + self._rule_touched = True + self._name_touched = False + self._update_default_name() + self._update_preview() + + def _on_name_edited(self, _text: str) -> None: + self._name_touched = True + + def _update_preview(self) -> None: + variable, time, values = self._selected_variable_data() + rule = self.edit_rule.text().strip() + self.plot.clear() + if self.btn_ok is not None: + self.btn_ok.setEnabled(False) + if not variable: + self.lbl_status.setText("Load a CSV or EthoVision file with continuous numeric columns first.") + return + n = min(time.size, values.size) + if n <= 0: + self.lbl_status.setText("Selected variable has no numeric samples.") + return + time = time[:n] + values = values[:n] + try: + on, off, dur, mask = _continuous_threshold_events(time, values, rule, variable) + except Exception as exc: + self.plot.plot(time, values, pen=pg.mkPen((120, 170, 220), width=1)) + self.lbl_status.setText(str(exc)) + return + self.plot.plot(time, values, pen=pg.mkPen((100, 190, 255), width=1)) + if mask.size == values.size and np.any(mask): + masked = np.where(mask, values, np.nan) + self.plot.plot(time, masked, pen=pg.mkPen((255, 180, 70), width=2)) + event_count = int(on.size if self.combo_align.currentText().endswith("onset") else off.size) + sample_count = int(np.sum(mask)) if mask.size else 0 + self.lbl_status.setText( + f"{sample_count} sample(s) pass the rule. {event_count} event(s) will be created." + ) + if self.btn_ok is not None: + self.btn_ok.setEnabled(event_count > 0) + + def config(self) -> Dict[str, object]: + source_key, _info = self._selected_source() + return { + "source_key": source_key, + "variable": str(self.combo_variable.currentData() or self.combo_variable.currentText() or ""), + "rule": self.edit_rule.text().strip(), + "align": self.combo_align.currentText(), + "name": self.edit_name.text().strip(), + "apply_all": self.cb_apply_all.isChecked(), + } + + def _compute_psth_matrix( t: np.ndarray, y: np.ndarray, @@ -496,6 +809,7 @@ def __init__(self, parent=None) -> None: self._processed: List[ProcessedTrial] = [] self._dio_cache: Dict[Tuple[str, str], Tuple[np.ndarray, np.ndarray]] = {} # (path,dio)->(t,x) self._behavior_sources: Dict[str, Dict[str, Any]] = {} # stem->behavior data + self._continuous_align_rules: Dict[str, Dict[str, object]] = {} self._last_mat: Optional[np.ndarray] = None self._last_tvec: Optional[np.ndarray] = None self._last_events: Optional[np.ndarray] = None @@ -632,6 +946,7 @@ def _build_ui(self) -> None: self.btn_use_current = QtWidgets.QPushButton("Use current preprocessed selection") self.btn_use_current.setProperty("class", "compactPrimary") self.btn_use_current.setSizePolicy(QtWidgets.QSizePolicy.Policy.Ignored, QtWidgets.QSizePolicy.Policy.Fixed) + self.btn_use_current.setVisible(False) self.btn_load_processed_single = QtWidgets.QPushButton("Load processed file (CSV/H5)") self.btn_load_processed_single.setProperty("class", "compactSmall") self.btn_load_processed_single.setSizePolicy( @@ -639,7 +954,6 @@ def _build_ui(self) -> None: QtWidgets.QSizePolicy.Policy.Fixed, ) single_layout.addWidget(self.lbl_current) - single_layout.addWidget(self.btn_use_current) single_layout.addWidget(self.btn_load_processed_single) group_layout = QtWidgets.QVBoxLayout(tab_group) @@ -654,9 +968,9 @@ def _build_ui(self) -> None: self.btn_refresh_dio = QtWidgets.QPushButton("Refresh A/D channel list") self.btn_refresh_dio.setProperty("class", "compactSmall") self.btn_refresh_dio.setSizePolicy(QtWidgets.QSizePolicy.Policy.Ignored, QtWidgets.QSizePolicy.Policy.Fixed) + self.btn_refresh_dio.setVisible(False) vsrc.addWidget(self.tab_sources) - vsrc.addWidget(self.btn_refresh_dio) grp_align = QtWidgets.QGroupBox("Behavior / Events") grp_align.setSizePolicy(QtWidgets.QSizePolicy.Policy.Preferred, QtWidgets.QSizePolicy.Policy.Expanding) @@ -721,7 +1035,7 @@ def _build_ui(self) -> None: # Preprocessed files list self.list_preprocessed = FileDropList() - self.list_preprocessed.setMinimumHeight(180) + self.list_preprocessed.setMinimumHeight(260) self.list_preprocessed.setSizePolicy( QtWidgets.QSizePolicy.Policy.Expanding, QtWidgets.QSizePolicy.Policy.Expanding, @@ -730,7 +1044,7 @@ def _build_ui(self) -> None: # Behaviors list self.list_behaviors = FileDropList() - self.list_behaviors.setMinimumHeight(180) + self.list_behaviors.setMinimumHeight(260) self.list_behaviors.setSizePolicy( QtWidgets.QSizePolicy.Policy.Expanding, QtWidgets.QSizePolicy.Policy.Expanding, @@ -800,14 +1114,17 @@ def _build_ui(self) -> None: self.spin_transition_gap.setValue(1.0) self.spin_transition_gap.setDecimals(2) + self.lbl_behavior_name = QtWidgets.QLabel("Behavior name") + self.lbl_behavior_align = QtWidgets.QLabel("Behavior align") self.lbl_trans_from = QtWidgets.QLabel("Transition from") self.lbl_trans_to = QtWidgets.QLabel("Transition to") self.lbl_trans_gap = QtWidgets.QLabel("Transition gap (s)") - fal.addRow("Behavior name", self.combo_behavior_name) - fal.addRow("Behavior align", self.combo_behavior_align) - fal.addRow(self.lbl_trans_from, self.combo_behavior_from) - fal.addRow(self.lbl_trans_to, self.combo_behavior_to) - fal.addRow(self.lbl_trans_gap, self.spin_transition_gap) + self.lbl_continuous_align = QtWidgets.QLabel("Continuous") + self.btn_continuous_align = QtWidgets.QPushButton("Align to continuous") + self.btn_continuous_align.setProperty("class", "compactSmall") + self.lbl_continuous_align_status = QtWidgets.QLabel("No continuous rule") + self.lbl_continuous_align_status.setProperty("class", "hint") + self.lbl_continuous_align_status.setWordWrap(True) # ── Shared QSS for PSTH subsection headers ── @@ -900,6 +1217,29 @@ def _dual_row(lbl_a: str, w_a, lbl_b: str, w_b): metric_post_widget = _dual_row("Start:", self.spin_metric_post0, "End:", self.spin_metric_post1) global_widget = _dual_row("Start:", self.spin_global_start, "End:", self.spin_global_end) + # ═══════════════════════════════════════════════════════ + # Section 0: Alignment + # ═══════════════════════════════════════════════════════ + grp_align_psth = QtWidgets.QGroupBox("Alignment") + grp_align_psth.setStyleSheet(_psth_section_qss) + fa_psth = QtWidgets.QFormLayout(grp_align_psth) + fa_psth.setRowWrapPolicy(QtWidgets.QFormLayout.RowWrapPolicy.WrapLongRows) + fa_psth.setLabelAlignment(QtCore.Qt.AlignmentFlag.AlignLeft | QtCore.Qt.AlignmentFlag.AlignTop) + fa_psth.addRow(self.lbl_behavior_name, self.combo_behavior_name) + fa_psth.addRow(self.lbl_behavior_align, self.combo_behavior_align) + fa_psth.addRow(self.lbl_trans_from, self.combo_behavior_from) + fa_psth.addRow(self.lbl_trans_to, self.combo_behavior_to) + fa_psth.addRow(self.lbl_trans_gap, self.spin_transition_gap) + continuous_row = QtWidgets.QHBoxLayout() + continuous_row.setContentsMargins(0, 0, 0, 0) + continuous_row.setSpacing(6) + continuous_row.addWidget(self.btn_continuous_align, 0) + continuous_row.addWidget(self.lbl_continuous_align_status, 1) + continuous_widget = QtWidgets.QWidget() + continuous_widget.setLayout(continuous_row) + fa_psth.addRow(self.lbl_continuous_align, continuous_widget) + self._continuous_align_widget = continuous_widget + # ═══════════════════════════════════════════════════════ # Section 1 — Window & Baseline # ═══════════════════════════════════════════════════════ @@ -992,6 +1332,7 @@ def _dual_row(lbl_a: str, w_a, lbl_b: str, w_b): _psth_vbox = QtWidgets.QVBoxLayout(grp_opt) _psth_vbox.setContentsMargins(0, 0, 0, 0) _psth_vbox.setSpacing(4) + _psth_vbox.addWidget(grp_align_psth) _psth_vbox.addWidget(grp_window) _psth_vbox.addWidget(grp_filt) _psth_vbox.addWidget(grp_include) @@ -1856,8 +2197,8 @@ def _set_text_with_banner(text: str) -> None: self.btn_setup_load.setProperty("class", "compactPrimarySmall") self.btn_setup_refresh = QtWidgets.QPushButton("Refresh A/D") self.btn_setup_refresh.setProperty("class", "compactSmall") + self.btn_setup_refresh.setVisible(False) setup_btn_row.addWidget(self.btn_setup_load) - setup_btn_row.addWidget(self.btn_setup_refresh) setup_btn_row.addStretch(1) setup_wrap = QtWidgets.QWidget() setup_wrap.setLayout(setup_btn_row) @@ -1954,6 +2295,8 @@ def _set_text_with_banner(text: str) -> None: self.menu_recent_projects = self.menu_action_recent.addMenu("Projects") self.menu_action_recent.aboutToShow.connect(self._refresh_recent_postprocessing_menus) self.act_refresh_dio = self.menu_action_load.addAction("Refresh A/D channel list") + self.act_load_current.setVisible(False) + self.act_refresh_dio.setVisible(False) self.menu_action_load.addSeparator() self.act_open_plot_style = self.menu_action_load.addAction("Plot style...") self.btn_action_load.setMenu(self.menu_action_load) @@ -2591,6 +2934,7 @@ def _set_text_with_banner(text: str) -> None: self.combo_align.currentIndexChanged.connect(self._update_align_ui) self.combo_behavior_file_type.currentIndexChanged.connect(self._update_align_ui) self.combo_behavior_align.currentIndexChanged.connect(self._update_align_ui) + self.btn_continuous_align.clicked.connect(self._open_continuous_align_dialog) self.combo_align.currentIndexChanged.connect(self._refresh_behavior_list) self.combo_align.currentIndexChanged.connect(self._compute_psth) for w in ( @@ -3817,12 +4161,30 @@ def _set_resample_from_processed(self) -> None: self.spin_resample.setValue(fs) self._update_status_strip() + def _refresh_dio_channels_from_processed(self) -> None: + names: List[str] = [] + seen: set[str] = set() + for proc in self._processed: + name = str(getattr(proc, "dio_name", "") or "").strip() + if not name or name in seen: + continue + dio = getattr(proc, "dio", None) + try: + has_data = dio is not None and np.asarray(dio).size > 0 + except Exception: + has_data = dio is not None + if has_data: + seen.add(name) + names.append(name) + self.receive_dio_list(names) + @QtCore.Slot(list) def receive_current_processed(self, processed_list: List[ProcessedTrial]) -> None: self._processed = processed_list or [] if not self._autosave_restoring: self._project_dirty = True - # update trace preview with first entry + self._update_file_lists() + self._refresh_dio_channels_from_processed() self._refresh_behavior_list() self._refresh_sync_sources() self._set_resample_from_processed() @@ -3841,6 +4203,8 @@ def append_processed(self, processed_list: List[ProcessedTrial]) -> None: self._processed.extend(processed_list) if not self._autosave_restoring: self._project_dirty = True + self._update_file_lists() + self._refresh_dio_channels_from_processed() self._refresh_behavior_list() self._refresh_sync_sources() self._set_resample_from_processed() @@ -4022,6 +4386,7 @@ def _load_behavior_paths(self, paths: List[str], replace: bool) -> None: "Behavior load warning", f"No behavior or trajectory numeric columns detected in {os.path.basename(p)} for the selected file type.", ) + info.setdefault("event_behaviors", {}) info["source_path"] = str(p) self._behavior_sources[stem] = info loaded_any = True @@ -4043,6 +4408,95 @@ def _load_behavior_paths(self, paths: List[str], replace: bool) -> None: self._sync_temporal_modeling_context() self._refresh_sync_sources() + def _open_continuous_align_dialog(self) -> None: + has_continuous = any(bool((info or {}).get("trajectory") or {}) for info in (self._behavior_sources or {}).values()) + if not has_continuous: + QtWidgets.QMessageBox.information( + self, + "Align to continuous", + "Load a behavior CSV or EthoVision file with continuous numeric columns first.", + ) + return + dlg = ContinuousAlignDialog(self._behavior_sources, self) + if dlg.exec() != QtWidgets.QDialog.DialogCode.Accepted: + return + try: + count = self._apply_continuous_alignment_rule(dlg.config()) + except Exception as exc: + QtWidgets.QMessageBox.warning(self, "Align to continuous", str(exc)) + return + self.statusUpdate.emit(f"Created continuous alignment for {count} behavior file(s).", 5000) + + def _apply_continuous_alignment_rule(self, config: Dict[str, object]) -> int: + source_key = str(config.get("source_key") or "").strip() + variable = str(config.get("variable") or "").strip() + rule = str(config.get("rule") or "").strip() + align = str(config.get("align") or "Align to onset").strip() + name = str(config.get("name") or "").strip() or _continuous_behavior_name(variable, rule, align) + apply_all = bool(config.get("apply_all", True)) + if not variable: + raise ValueError("Choose a continuous variable.") + if not rule: + raise ValueError("Enter a threshold rule.") + + targets: List[Tuple[str, Dict[str, Any]]] = [] + if apply_all: + for stem, info in (self._behavior_sources or {}).items(): + if variable in (info.get("trajectory") or {}): + targets.append((stem, info)) + elif source_key in self._behavior_sources: + targets.append((source_key, self._behavior_sources[source_key])) + if not targets: + raise ValueError(f"No loaded behavior file contains '{variable}'.") + + total_events = 0 + updated = 0 + first_align_index = self.combo_behavior_align.findText(align) + for stem, info in targets: + trajectory = info.get("trajectory") or {} + values = np.asarray(trajectory.get(variable, np.array([], float)), float) + time = np.asarray(info.get("trajectory_time", np.array([], float)), float) + on, off, dur, _mask = _continuous_threshold_events(time, values, rule, variable) + events = off if align.endswith("offset") else on + if events.size == 0: + continue + event_store = info.setdefault("event_behaviors", {}) + event_store[name] = { + "on": np.asarray(on, float), + "off": np.asarray(off, float), + "dur": np.asarray(dur, float), + "rule": rule, + "variable": variable, + "align": align, + "source": stem, + } + self._continuous_align_rules[name] = { + "source_key": stem, + "variable": variable, + "rule": rule, + "align": align, + "name": name, + "apply_all": apply_all, + } + total_events += int(events.size) + updated += 1 + + if updated == 0: + raise ValueError("The threshold rule did not create any onset or offset events.") + + self._refresh_behavior_list() + idx_name = self.combo_behavior_name.findText(name) + if idx_name >= 0: + self.combo_behavior_name.setCurrentIndex(idx_name) + if first_align_index >= 0: + self.combo_behavior_align.setCurrentIndex(first_align_index) + self.lbl_continuous_align_status.setText(f"{name}: {total_events} event(s)") + if not self._autosave_restoring: + self._project_dirty = True + self._compute_psth() + self._refresh_sync_sources() + return updated + def _sync_temporal_modeling_context(self) -> None: if not hasattr(self, "section_temporal"): return @@ -4094,6 +4548,7 @@ def _load_processed_paths(self, paths: List[str], replace: bool) -> None: self.lbl_group.setText(f"{len(self._processed)} file(s) loaded") self._push_recent_paths("postprocess_recent_processed_paths", paths) self._update_file_lists() + self._refresh_dio_channels_from_processed() self._refresh_sync_sources() self._set_resample_from_processed() self._compute_psth() @@ -4140,6 +4595,8 @@ def _update_align_ui(self) -> None: self.combo_behavior_file_type.setVisible(use_beh) self.btn_load_beh.setEnabled(use_beh) self.btn_load_beh.setVisible(use_beh) + self.lbl_behavior_name.setEnabled(use_beh) + self.lbl_behavior_name.setVisible(use_beh) self.combo_behavior_name.setEnabled(use_beh) self.combo_behavior_name.setVisible(use_beh) show_time_panel = bool(use_beh and self._behavior_sources_need_generated_time()) @@ -4147,8 +4604,20 @@ def _update_align_ui(self) -> None: self.grp_behavior_time.setVisible(show_time_panel) # Behavior align combo + transition settings + self.lbl_behavior_align.setEnabled(use_beh) + self.lbl_behavior_align.setVisible(use_beh) self.combo_behavior_align.setEnabled(use_beh) self.combo_behavior_align.setVisible(use_beh) + for w in ( + self.lbl_continuous_align, + self.btn_continuous_align, + self.lbl_continuous_align_status, + getattr(self, "_continuous_align_widget", None), + ): + if w is None: + continue + w.setEnabled(use_beh) + w.setVisible(use_beh) is_transition = use_beh and self.combo_behavior_align.currentText().startswith("Transition") for w in ( self.combo_behavior_from, @@ -4896,6 +5365,7 @@ def _project_dirty_fingerprint(self) -> str: "kind": str(data.get("kind", "") or ""), "row_count": int(data.get("row_count", 0) or 0), "behaviors": sorted(str(k) for k in (data.get("behaviors", {}) or {}).keys()), + "event_behaviors": sorted(str(k) for k in (data.get("event_behaviors", {}) or {}).keys()), "trajectory": sorted(str(k) for k in (data.get("trajectory", {}) or {}).keys()), } ) @@ -5549,6 +6019,10 @@ def _load_processed_h5(self, path: str) -> Optional[ProcessedTrial]: ) def _refresh_behavior_list(self) -> None: + prev_name = self.combo_behavior_name.currentText().strip() + prev_analysis = self.combo_behavior_analysis.currentText().strip() if hasattr(self, "combo_behavior_analysis") else "" + prev_from = self.combo_behavior_from.currentText().strip() + prev_to = self.combo_behavior_to.currentText().strip() self.combo_behavior_name.clear() if hasattr(self, "combo_behavior_analysis"): self.combo_behavior_analysis.clear() @@ -5560,6 +6034,8 @@ def _refresh_behavior_list(self) -> None: except Exception: pass if not self._behavior_sources: + self.combo_behavior_from.clear() + self.combo_behavior_to.clear() self._refresh_spatial_columns() self._compute_spatial_heatmap() self._update_data_availability() @@ -5569,6 +6045,8 @@ def _refresh_behavior_list(self) -> None: for info in self._behavior_sources.values(): behaviors = info.get("behaviors") or {} behavior_names.update(str(k) for k in behaviors.keys()) + event_behaviors = info.get("event_behaviors") or {} + behavior_names.update(str(k) for k in event_behaviors.keys()) behaviors = sorted(list(behavior_names)) for name in behaviors: self.combo_behavior_name.addItem(name) @@ -5579,6 +6057,18 @@ def _refresh_behavior_list(self) -> None: for name in behaviors: self.combo_behavior_from.addItem(name) self.combo_behavior_to.addItem(name) + for combo, previous in ( + (self.combo_behavior_name, prev_name), + (self.combo_behavior_from, prev_from), + (self.combo_behavior_to, prev_to), + ): + idx = combo.findText(previous) + if idx >= 0: + combo.setCurrentIndex(idx) + if hasattr(self, "combo_behavior_analysis"): + idx = self.combo_behavior_analysis.findText(prev_analysis) + if idx >= 0: + self.combo_behavior_analysis.setCurrentIndex(idx) # Update the lists with numbered items self._update_file_lists() @@ -7010,6 +7500,26 @@ def _export_sync_aligned_files(self) -> None: self.statusUpdate.emit(f"Aligned time export complete: {out_dir}", 5000) def _extract_behavior_events(self, info: Dict[str, Any], behavior_name: str) -> Tuple[np.ndarray, np.ndarray, np.ndarray]: + event_behaviors = info.get("event_behaviors") or {} + if behavior_name in event_behaviors: + event_info = event_behaviors.get(behavior_name) or {} + if isinstance(event_info, dict): + on = np.asarray(event_info.get("on", np.array([], float)), float) + off = np.asarray(event_info.get("off", on), float) + dur = np.asarray(event_info.get("dur", np.array([], float)), float) + else: + on = np.asarray(event_info, float) + off = on.copy() + dur = np.full(on.shape, np.nan, dtype=float) + on = on[np.isfinite(on)] + off = off[np.isfinite(off)] + on = np.sort(np.unique(on)) + off = np.sort(np.unique(off)) if off.size else on.copy() + if dur.size != on.size: + m = min(on.size, off.size) + dur = off[:m] - on[:m] if m else np.array([], float) + return on, off, np.asarray(dur, float) + behaviors = info.get("behaviors") or {} if behavior_name not in behaviors: return np.array([], float), np.array([], float), np.array([], float) @@ -9924,6 +10434,24 @@ def _save_project_h5(self, path: str) -> None: ds = behaviors_group.create_dataset(f"item_{b_idx:04d}", data=data, **kwargs) ds.attrs["name"] = str(name) + event_behaviors_group = entry.create_group("event_behaviors") + event_behaviors = source.get("event_behaviors") or {} + for e_idx, (name, event_info) in enumerate(event_behaviors.items()): + event_entry = event_behaviors_group.create_group(f"item_{e_idx:04d}") + event_entry.attrs["name"] = str(name) + if isinstance(event_info, dict): + for attr_name in ("rule", "variable", "align", "source"): + if event_info.get(attr_name) is not None: + event_entry.attrs[attr_name] = str(event_info.get(attr_name)) + self._write_h5_numeric(event_entry, "on", np.asarray(event_info.get("on", np.array([], float)), float)) + self._write_h5_numeric(event_entry, "off", np.asarray(event_info.get("off", np.array([], float)), float)) + self._write_h5_numeric(event_entry, "dur", np.asarray(event_info.get("dur", np.array([], float)), float)) + else: + values = np.asarray(event_info, float) + self._write_h5_numeric(event_entry, "on", values) + self._write_h5_numeric(event_entry, "off", values) + self._write_h5_numeric(event_entry, "dur", np.full(values.shape, np.nan, dtype=float)) + trajectory_group = entry.create_group("trajectory") trajectory = source.get("trajectory") or {} for t_idx, (name, values) in enumerate(trajectory.items()): @@ -10082,6 +10610,7 @@ def _aligned(values: Optional[np.ndarray], fill_nan: bool = True) -> np.ndarray: float, ), "behaviors": {}, + "event_behaviors": {}, "trajectory": {}, "trajectory_time": np.asarray( self._read_h5_numeric(entry, "trajectory_time") @@ -10105,6 +10634,26 @@ def _aligned(values: Optional[np.ndarray], fill_nan: bool = True) -> np.ndarray: name = self._h5_text(ds.attrs.get("name", b_key), b_key) info["behaviors"][name] = np.asarray(ds[()], float) + event_behaviors_group = entry.get("event_behaviors") + if isinstance(event_behaviors_group, h5py.Group): + for e_key in sorted(event_behaviors_group.keys()): + event_entry = event_behaviors_group.get(e_key) + if not isinstance(event_entry, h5py.Group): + continue + name = self._h5_text(event_entry.attrs.get("name", e_key), e_key) + on = self._read_h5_numeric(event_entry, "on") + off = self._read_h5_numeric(event_entry, "off") + dur = self._read_h5_numeric(event_entry, "dur") + info["event_behaviors"][name] = { + "on": np.asarray(on if on is not None else np.array([], float), float), + "off": np.asarray(off if off is not None else np.array([], float), float), + "dur": np.asarray(dur if dur is not None else np.array([], float), float), + "rule": self._h5_text(event_entry.attrs.get("rule", ""), ""), + "variable": self._h5_text(event_entry.attrs.get("variable", ""), ""), + "align": self._h5_text(event_entry.attrs.get("align", ""), ""), + "source": self._h5_text(event_entry.attrs.get("source", stem), stem), + } + trajectory_group = entry.get("trajectory") if isinstance(trajectory_group, h5py.Group): for t_key in sorted(trajectory_group.keys()): @@ -10273,12 +10822,16 @@ def _reset_project_state(self) -> None: self._clear_cached_analysis_outputs() self._processed = [] self._behavior_sources = {} + self._continuous_align_rules = {} self._sync_results_by_file = {} self._last_sync_preview = None self._pending_project_recompute_from_current = False self._dio_cache.clear() + self.receive_dio_list([]) self.lbl_group.setText("(none)") self.lbl_beh.setText("(none)") + if hasattr(self, "lbl_continuous_align_status"): + self.lbl_continuous_align_status.setText("No continuous rule") self.lbl_behavior_msg.setText("") self.lbl_signal_msg.setText("") if hasattr(self, "txt_sync_report"): @@ -10415,6 +10968,7 @@ def _load_project_from_path(self, path: str, from_autosave: bool = False) -> boo self.lbl_beh.setText(f"{len(self._behavior_sources)} file(s) loaded [{mode_label}]") self._update_file_lists() + self._refresh_dio_channels_from_processed() self._refresh_behavior_list() self._refresh_sync_sources() self._set_resample_from_processed() @@ -10665,6 +11219,7 @@ def _collect_settings(self) -> Dict[str, object]: "behavior_from": self.combo_behavior_from.currentText(), "behavior_to": self.combo_behavior_to.currentText(), "transition_gap": float(self.spin_transition_gap.value()), + "continuous_align_rules": copy.deepcopy(self._continuous_align_rules), "window_pre": float(self.spin_pre.value()), "window_post": float(self.spin_post.value()), "baseline_start": float(self.spin_b0.value()), @@ -10788,6 +11343,9 @@ def _set_combo_data(combo: QtWidgets.QComboBox, val: object) -> None: _set_combo(self.combo_behavior_to, data.get("behavior_to")) if "transition_gap" in data: self.spin_transition_gap.setValue(float(data["transition_gap"])) + rules = data.get("continuous_align_rules") + if isinstance(rules, dict): + self._continuous_align_rules = copy.deepcopy(rules) if "window_pre" in data: self.spin_pre.setValue(float(data["window_pre"])) if "window_post" in data: @@ -11411,6 +11969,10 @@ def _get_all_behavior_names(self) -> List[str]: for beh in behaviors: if beh not in names: names.append(beh) + event_behaviors = info.get("event_behaviors") or {} + for beh in event_behaviors: + if beh not in names: + names.append(beh) return names def _compute_psth_for_behavior(self, behavior_name: str) -> Tuple[Optional[np.ndarray], Optional[np.ndarray], List[str]]: diff --git a/pyBer/gui_preprocessing.py b/pyBer/gui_preprocessing.py index 6537033..5770b94 100644 --- a/pyBer/gui_preprocessing.py +++ b/pyBer/gui_preprocessing.py @@ -4,6 +4,8 @@ from typing import Callable, Dict, List, Optional, Tuple import json import os +import subprocess +import sys import numpy as np from PySide6 import QtCore, QtWidgets, QtGui @@ -549,6 +551,8 @@ def __init__(self, parent=None) -> None: def _build_ui(self) -> None: layout = QtWidgets.QVBoxLayout(self) + layout.setContentsMargins(6, 6, 6, 6) + layout.setSpacing(10) table_min_height = 260 auto_group = QtWidgets.QGroupBox("Auto-detected (threshold)") @@ -557,22 +561,31 @@ def _build_ui(self) -> None: QtWidgets.QSizePolicy.Policy.Expanding, ) auto_layout = QtWidgets.QVBoxLayout(auto_group) + auto_layout.setContentsMargins(8, 16, 8, 8) self.table_auto = QtWidgets.QTableWidget(0, 6) self.table_auto.setMinimumHeight(table_min_height) self.table_auto.setSizePolicy( QtWidgets.QSizePolicy.Policy.Expanding, QtWidgets.QSizePolicy.Policy.Expanding, ) - self.table_auto.setHorizontalHeaderLabels(["ID", "Remove", "Source", "Core (s)", "Cut start", "Cut end"]) - self.table_auto.horizontalHeader().setStretchLastSection(True) + self.table_auto.setHorizontalHeaderLabels(["ID", "Use", "Src", "Core", "Start", "End"]) + auto_header = self.table_auto.horizontalHeader() + auto_header.setStretchLastSection(False) + auto_header.setHighlightSections(False) + auto_header.setMinimumSectionSize(28) self.table_auto.verticalHeader().setVisible(False) + self.table_auto.verticalHeader().setDefaultSectionSize(26) self.table_auto.setSelectionBehavior(QtWidgets.QAbstractItemView.SelectionBehavior.SelectRows) self.table_auto.setSelectionMode(QtWidgets.QAbstractItemView.SelectionMode.ExtendedSelection) - self.table_auto.setColumnWidth(0, 42) - self.table_auto.setColumnWidth(1, 72) - self.table_auto.setColumnWidth(2, 66) - self.table_auto.setColumnWidth(3, 132) - self.table_auto.setColumnWidth(4, 82) + self.table_auto.setHorizontalScrollBarPolicy(QtCore.Qt.ScrollBarPolicy.ScrollBarAlwaysOff) + self.table_auto.setAlternatingRowColors(True) + self.table_auto.setShowGrid(False) + self.table_auto.setWordWrap(False) + for col, width in ((0, 36), (1, 44), (2, 48)): + auto_header.setSectionResizeMode(col, QtWidgets.QHeaderView.ResizeMode.Fixed) + self.table_auto.setColumnWidth(col, width) + for col in (3, 4, 5): + auto_header.setSectionResizeMode(col, QtWidgets.QHeaderView.ResizeMode.Stretch) auto_layout.addWidget(self.table_auto, 1) layout.addWidget(auto_group, 1) @@ -582,17 +595,30 @@ def _build_ui(self) -> None: QtWidgets.QSizePolicy.Policy.Expanding, ) manual_layout = QtWidgets.QVBoxLayout(manual_group) + manual_layout.setContentsMargins(8, 16, 8, 8) self.table = QtWidgets.QTableWidget(0, 3) self.table.setMinimumHeight(table_min_height) self.table.setSizePolicy( QtWidgets.QSizePolicy.Policy.Expanding, QtWidgets.QSizePolicy.Policy.Expanding, ) - self.table.setHorizontalHeaderLabels(["ID", "Start (s)", "End (s)"]) - self.table.horizontalHeader().setStretchLastSection(True) + self.table.setHorizontalHeaderLabels(["ID", "Start", "End"]) + manual_header = self.table.horizontalHeader() + manual_header.setStretchLastSection(False) + manual_header.setHighlightSections(False) + manual_header.setMinimumSectionSize(32) + manual_header.setSectionResizeMode(0, QtWidgets.QHeaderView.ResizeMode.Fixed) + manual_header.setSectionResizeMode(1, QtWidgets.QHeaderView.ResizeMode.Stretch) + manual_header.setSectionResizeMode(2, QtWidgets.QHeaderView.ResizeMode.Stretch) + self.table.setColumnWidth(0, 42) self.table.verticalHeader().setVisible(False) + self.table.verticalHeader().setDefaultSectionSize(26) self.table.setSelectionBehavior(QtWidgets.QAbstractItemView.SelectionBehavior.SelectRows) self.table.setSelectionMode(QtWidgets.QAbstractItemView.SelectionMode.ExtendedSelection) + self.table.setHorizontalScrollBarPolicy(QtCore.Qt.ScrollBarPolicy.ScrollBarAlwaysOff) + self.table.setAlternatingRowColors(True) + self.table.setShowGrid(False) + self.table.setWordWrap(False) manual_layout.addWidget(self.table) @@ -603,7 +629,8 @@ def _build_ui(self) -> None: ed.setDecimals(3) ed.setRange(-1e9, 1e9) ed.setKeyboardTracking(False) - ed.setMinimumWidth(140) + ed.setMinimumWidth(96) + ed.setSizePolicy(QtWidgets.QSizePolicy.Policy.Expanding, QtWidgets.QSizePolicy.Policy.Fixed) self.btn_add = QtWidgets.QPushButton("Add") self.btn_update = QtWidgets.QPushButton("Update selected") @@ -851,6 +878,7 @@ class FileQueuePanel(QtWidgets.QGroupBox): openFileRequested = QtCore.Signal() openFolderRequested = QtCore.Signal() selectionChanged = QtCore.Signal() + sendToPostprocessingRequested = QtCore.Signal(list) channelChanged = QtCore.Signal(str) triggerChanged = QtCore.Signal(str) @@ -918,6 +946,8 @@ def _build_ui(self) -> None: self.list_files.setSelectionMode(QtWidgets.QAbstractItemView.SelectionMode.ExtendedSelection) self.list_files.setMinimumHeight(210) self.list_files.setUniformItemSizes(True) + self.list_files.setContextMenuPolicy(QtCore.Qt.ContextMenuPolicy.CustomContextMenu) + self.list_files.customContextMenuRequested.connect(self._show_file_context_menu) self.btn_remove_file = QtWidgets.QPushButton("Remove selected") self.btn_remove_file.setProperty("class", "blueSecondarySmall") @@ -1148,6 +1178,50 @@ def _remove_selected_files(self) -> None: def _update_remove_button(self) -> None: self.btn_remove_file.setEnabled(len(self.list_files.selectedItems()) > 0) + def _show_file_context_menu(self, pos: QtCore.QPoint) -> None: + item = self.list_files.itemAt(pos) + if item is not None and not item.isSelected(): + self.list_files.clearSelection() + item.setSelected(True) + self.list_files.setCurrentItem(item) + paths = self.selected_paths() + if not paths: + return + + menu = QtWidgets.QMenu(self) + act_reveal = menu.addAction("Reveal in Explorer") + act_send = menu.addAction("Load in postprocessing") + menu.addSeparator() + act_remove = menu.addAction("Remove from list") + chosen = menu.exec(self.list_files.viewport().mapToGlobal(pos)) + if chosen is act_reveal: + self._reveal_path_in_file_manager(paths[0]) + elif chosen is act_send: + self.sendToPostprocessingRequested.emit(paths) + elif chosen is act_remove: + self._remove_selected_files() + + def _reveal_path_in_file_manager(self, path: str) -> None: + target = os.path.normpath(str(path or "")) + if not target: + return + folder = target if os.path.isdir(target) else os.path.dirname(target) + try: + if sys.platform.startswith("win"): + if os.path.exists(target): + subprocess.Popen(["explorer", "/select,", target]) + elif folder and os.path.isdir(folder): + subprocess.Popen(["explorer", folder]) + return + if sys.platform == "darwin" and os.path.exists(target): + subprocess.Popen(["open", "-R", target]) + return + if folder and os.path.isdir(folder): + QtGui.QDesktopServices.openUrl(QtCore.QUrl.fromLocalFile(folder)) + except Exception: + if folder and os.path.isdir(folder): + QtGui.QDesktopServices.openUrl(QtCore.QUrl.fromLocalFile(folder)) + class SectionParamsDialog(QtWidgets.QDialog): def __init__(self, params: ProcessingParams, parent=None) -> None: diff --git a/pyBer/main.py b/pyBer/main.py index 401203d..6fce0a9 100644 --- a/pyBer/main.py +++ b/pyBer/main.py @@ -1750,6 +1750,7 @@ def _build_ui(self) -> None: self.file_panel.advancedOptionsRequested.connect(self._open_advanced_options) self.file_panel.qcRequested.connect(self._run_qc_dialog) self.file_panel.batchQcRequested.connect(self._run_batch_qc) + self.file_panel.sendToPostprocessingRequested.connect(self._send_preprocessing_paths_to_postprocessing) # Parameters: changes and actions self.param_panel.paramsChanged.connect(self._on_params_changed) @@ -2411,11 +2412,16 @@ def _update_pre_drawer_visibility(self) -> None: sizes = splitter.sizes() if len(sizes) >= 2: if any_checked: - if sizes[0] < 60: - total = sum(sizes) or 1 - drawer_w = max(420, int(total * 0.28)) + total = sum(sizes) or 1 + if active_key == "artifacts": + target_w = max(520, int(total * 0.34)) + else: + target_w = max(420, int(total * 0.28)) + drawer_w = min(target_w, max(420, total - 640)) + if sizes[0] < 60 or (active_key == "artifacts" and sizes[0] < drawer_w - 20): + delta_w = max(0, drawer_w - max(0, sizes[0])) sizes[0] = drawer_w - sizes[1] = max(400, sizes[1] - drawer_w) + sizes[1] = max(400, sizes[1] - delta_w) splitter.setSizes(sizes) else: if sizes[0] > 0: @@ -5725,6 +5731,11 @@ def _on_main_tab_changed(self, index: int) -> None: _LOG.exception("Failed to handle main tab switch") finally: self._restore_window_state_after_tab_switch(was_fullscreen, was_maximized) + try: + if self.tabs.currentWidget() is self.post_tab: + QtCore.QTimer.singleShot(0, self._post_get_current_dio_list) + except Exception: + pass self._handling_main_tab_change = False if self._force_fixed_dock_layouts: QtCore.QTimer.singleShot(0, self._enforce_fixed_layout_for_active_tab) @@ -5962,6 +5973,7 @@ def _on_file_selection_changed(self) -> None: self._current_channel = None self._current_trigger = None self.plots.set_title("No file loaded") + self._post_get_current_dio_list() self._update_plot_status() return @@ -5998,6 +6010,7 @@ def _on_file_selection_changed(self) -> None: # update post tab selection context self.post_tab.set_current_source_label(os.path.basename(path), self._current_channel or "") + self._post_get_current_dio_list() self._update_plot_status() def _on_channel_changed(self, ch: str) -> None: @@ -7306,19 +7319,32 @@ def _export_one(proc: ProcessedTrial, suffix: str = "") -> None: # ---------------- Postprocessing bridge ---------------- - @QtCore.Slot() - def _post_get_current_processed(self): - # Determine selection context: if multiple selected, provide multiple processed outputs if available - paths = self._selected_paths() - if not paths: - paths = [self._current_path] if self._current_path else [] + def _postprocessing_bridge_paths(self, paths: Optional[List[str]] = None) -> List[str]: + raw_paths = [str(p or "") for p in (paths or []) if str(p or "")] + if not raw_paths: + raw_paths = self._selected_paths() + if not raw_paths: + raw_paths = [self._current_path] if self._current_path else [] + out: List[str] = [] + seen: set[str] = set() + for path in raw_paths: + if not path or path in seen: + continue + seen.add(path) + out.append(path) + return out + def _processed_trials_for_postprocessing_paths(self, paths: List[str]) -> List[ProcessedTrial]: out: List[ProcessedTrial] = [] + try: + params = self.param_panel.get_params() + except Exception: + params = ProcessingParams() + start_s, end_s = self._time_window_bounds() for p in paths: doric = self._loaded_files.get(p) if not doric: continue - # Use current channel when available for all selected files if self._current_channel and self._current_channel in doric.channels: ch = self._current_channel else: @@ -7326,45 +7352,69 @@ def _post_get_current_processed(self): key = (p, ch) if key in self._last_processed: out.append(self._last_processed[key]) - else: - # compute on-demand (fast due to decimation), using current params - try: - params = self.param_panel.get_params() - trial = doric.make_trial(ch, trigger_name=self._current_trigger) - trial = self._apply_time_window(trial) - start_s, end_s = self._time_window_bounds() - manual = self._clip_regions_to_window(self._manual_regions_by_key.get(key, []), start_s, end_s) - manual_exclude = self._clip_regions_to_window(self._manual_exclude_by_key.get(key, []), start_s, end_s) - proc = self.processor.process_trial( - trial, - params, - manual_regions_sec=manual, - manual_exclude_regions_sec=manual_exclude, - preview_mode=False, - ) - cutouts = self._cutout_regions_by_key.get(key, []) - proc = self._apply_cutouts_to_processed(proc, cutouts) - self._last_processed[key] = proc - out.append(proc) - except Exception: - pass - - self.post_tab.receive_current_processed(out) - - @QtCore.Slot() - def _post_get_current_dio_list(self): - # Analog/digital channel list for current/selected files: union. - paths = self._selected_paths() - if not paths: - paths = [self._current_path] if self._current_path else [] + continue + try: + trial = doric.make_trial(ch, trigger_name=self._current_trigger) + trial = self._apply_time_window(trial) + manual = self._clip_regions_to_window(self._manual_regions_by_key.get(key, []), start_s, end_s) + manual_exclude = self._clip_regions_to_window(self._manual_exclude_by_key.get(key, []), start_s, end_s) + proc = self.processor.process_trial( + trial, + params, + manual_regions_sec=manual, + manual_exclude_regions_sec=manual_exclude, + preview_mode=False, + ) + cutouts = self._cutout_regions_by_key.get(key, []) + proc = self._apply_cutouts_to_processed(proc, cutouts) + self._last_processed[key] = proc + out.append(proc) + except Exception: + pass + return out - dio = set() + def _send_dio_list_for_paths_to_postprocessing(self, paths: List[str]) -> None: + dio: set[str] = set() for p in paths: f = self._loaded_files.get(p) if f: dio |= set(f.trigger_by_name.keys()) self.post_tab.receive_dio_list(sorted(dio)) + @QtCore.Slot(list) + def _send_preprocessing_paths_to_postprocessing(self, paths: List[str]) -> None: + source_paths = self._postprocessing_bridge_paths(paths) + processed = self._processed_trials_for_postprocessing_paths(source_paths) + if not processed: + self._show_status_message("No selected preprocessing file could be loaded into postprocessing.", 6000) + return + self.post_tab.receive_current_processed(processed) + self._send_dio_list_for_paths_to_postprocessing(source_paths) + first = processed[0] + self.post_tab.set_current_source_label( + os.path.basename(getattr(first, "path", "") or ""), + str(getattr(first, "channel_id", "") or ""), + ) + try: + idx = self.tabs.indexOf(self.post_tab) + if idx >= 0: + self.tabs.setCurrentIndex(idx) + except Exception: + pass + self._show_status_message(f"Loaded {len(processed)} preprocessing file(s) into postprocessing.", 6000) + + @QtCore.Slot() + def _post_get_current_processed(self): + paths = self._postprocessing_bridge_paths() + out = self._processed_trials_for_postprocessing_paths(paths) + self.post_tab.receive_current_processed(out) + self._send_dio_list_for_paths_to_postprocessing(paths) + + @QtCore.Slot() + def _post_get_current_dio_list(self): + paths = self._postprocessing_bridge_paths() + self._send_dio_list_for_paths_to_postprocessing(paths) + @QtCore.Slot(str, str) def _post_get_dio_data_for_path(self, path: str, dio_name: str): """ From 18d7fb22b2a955a913444d66b931b60e7248b05f Mon Sep 17 00:00:00 2001 From: andrianj Date: Mon, 1 Jun 2026 17:01:47 +0200 Subject: [PATCH 4/7] Improve sync and modeling tools --- pyBer/gui_postprocessing.py | 1564 +++++++++++++++++++++++++++++++++-- pyBer/styles.py | 45 + pyBer/temporal_modeling.py | 243 +++++- 3 files changed, 1749 insertions(+), 103 deletions(-) diff --git a/pyBer/gui_postprocessing.py b/pyBer/gui_postprocessing.py index ddc3297..c4958a1 100644 --- a/pyBer/gui_postprocessing.py +++ b/pyBer/gui_postprocessing.py @@ -3,6 +3,7 @@ import os import re +import sys import json import copy import logging @@ -41,6 +42,7 @@ _FIXED_POST_VISIBLE_SECTIONS = frozenset({"setup", "spatial", "psth", "export", "temporal"}) _FIXED_POST_RIGHT_TAB_ORDER = ("setup", "psth", "spatial", "temporal", "export") _POST_RIGHT_PANEL_MIN_WIDTH = 420 +_SYNC_TOOL_ROOT = Path(r"C:\Analysis\app_project\barcode_reader\video_barcode_extractor") _FIXED_POST_RIGHT_TAB_TITLES: Dict[str, str] = { "setup": "Setup", "psth": "PSTH", @@ -80,6 +82,117 @@ def _is_doric_channel_align(text: str) -> bool: return "doric" in (text or "").strip().lower() +def _ensure_sync_tool_import_path() -> bool: + root = _SYNC_TOOL_ROOT + if not root.is_dir(): + return False + root_text = str(root) + if root_text not in sys.path: + sys.path.insert(0, root_text) + return True + + +def _load_sync_table(path: str): + import pandas as pd + + ext = os.path.splitext(path)[1].lower() + if ext in {".h5", ".hdf5"}: + return _load_sync_h5_table(path) + if ext in {".xlsx", ".xls"}: + return pd.read_excel(path, engine="openpyxl") + if _ensure_sync_tool_import_path(): + try: + from vbe.core.csv_loader import load_csv_robust + return load_csv_robust(path) + except Exception: + pass + return pd.read_csv(path, sep=None, engine="python", comment="#") + + +def _load_sync_h5_table(path: str): + import pandas as pd + + columns: Dict[str, np.ndarray] = {} + lengths: List[int] = [] + + def _unique_name(name: str) -> str: + clean = str(name or "value").strip().replace("\\", "/") + if clean not in columns: + return clean + i = 2 + while f"{clean}_{i}" in columns: + i += 1 + return f"{clean}_{i}" + + def _collect(group: h5py.Group, prefix: str = "") -> None: + for name, obj in group.items(): + label = str(name) + if isinstance(obj, h5py.Dataset): + try: + raw = np.asarray(obj[()]) + except Exception: + continue + if raw.ndim != 1 or raw.size < 2: + continue + try: + arr = raw.astype(float) + except Exception: + continue + if int(np.sum(np.isfinite(arr))) < 2: + continue + attr_label = "" + try: + attr_label = str(obj.attrs.get("label", "") or "") + except Exception: + attr_label = "" + col = _unique_name(f"{prefix}{attr_label or label}") + columns[col] = arr + lengths.append(int(arr.size)) + elif isinstance(obj, h5py.Group): + _collect(obj, f"{prefix}{label}/") + + with h5py.File(path, "r") as h5: + root = h5.get("data") + if isinstance(root, h5py.Group): + _collect(root) + else: + _collect(h5) + + if not columns: + raise ValueError("No 1D numeric datasets were found in the H5 file.") + target_len = max(set(lengths), key=lengths.count) + table: Dict[str, np.ndarray] = {} + for name, arr in columns.items(): + if int(arr.size) == int(target_len): + table[name] = np.asarray(arr, float) + if "time" not in table and "data/time" in table: + table["time"] = table["data/time"] + return pd.DataFrame(table) + + +def _numeric_columns_from_df(df, time_col: Optional[str] = None) -> Dict[str, np.ndarray]: + columns: Dict[str, np.ndarray] = {} + for c in df.columns: + name = str(c).strip() + if not name or (time_col and name == str(time_col)): + continue + arr = _numeric_column_array(df, name) + if arr.size == 0: + continue + finite = arr[np.isfinite(arr)] + if finite.size >= 2: + columns[name] = arr + return columns + + +def _time_array_for_signal_table(df) -> Tuple[np.ndarray, str]: + time_col = _detect_time_column(df) + if time_col: + return _numeric_column_array(df, time_col), str(time_col) + n = int(len(df.index)) if hasattr(df, "index") else 0 + return np.arange(n, dtype=float), "sample_index" + + class FileDropList(QtWidgets.QListWidget): filesDropped = QtCore.Signal(list) orderChanged = QtCore.Signal() @@ -741,6 +854,261 @@ def config(self) -> Dict[str, object]: } +class SyncRoiPreview(QtWidgets.QWidget): + roiChanged = QtCore.Signal(tuple) + + def __init__(self, parent: Optional[QtWidgets.QWidget] = None) -> None: + super().__init__(parent) + self.setMinimumSize(360, 240) + self.setSizePolicy(QtWidgets.QSizePolicy.Policy.Expanding, QtWidgets.QSizePolicy.Policy.Expanding) + self._pixmap = QtGui.QPixmap() + self._image_size = QtCore.QSize(0, 0) + self._roi = QtCore.QRect(20, 20, 80, 80) + self._drag_origin: Optional[QtCore.QPoint] = None + self._target_rect = QtCore.QRect() + + def set_frame_bgr(self, frame: np.ndarray) -> None: + arr = np.asarray(frame) + if arr.ndim == 3 and arr.shape[2] >= 3: + rgb = arr[:, :, :3][:, :, ::-1].copy() + h, w = rgb.shape[:2] + img = QtGui.QImage(rgb.data, w, h, int(rgb.strides[0]), QtGui.QImage.Format.Format_RGB888).copy() + elif arr.ndim == 2: + gray = np.asarray(arr, dtype=np.uint8).copy() + h, w = gray.shape + img = QtGui.QImage(gray.data, w, h, int(gray.strides[0]), QtGui.QImage.Format.Format_Grayscale8).copy() + else: + return + self._pixmap = QtGui.QPixmap.fromImage(img) + self._image_size = QtCore.QSize(int(w), int(h)) + if self._roi.width() <= 0 or self._roi.height() <= 0: + self._roi = QtCore.QRect(0, 0, max(1, w // 8), max(1, h // 8)) + self.update() + + def roi(self) -> Tuple[int, int, int, int]: + r = self._bounded_roi(self._roi) + return int(r.x()), int(r.y()), int(r.width()), int(r.height()) + + def set_roi(self, x: int, y: int, w: int, h: int, emit_signal: bool = False) -> None: + self._roi = self._bounded_roi(QtCore.QRect(int(x), int(y), max(1, int(w)), max(1, int(h)))) + self.update() + if emit_signal: + self.roiChanged.emit(self.roi()) + + def _bounded_roi(self, roi: QtCore.QRect) -> QtCore.QRect: + if self._image_size.width() <= 0 or self._image_size.height() <= 0: + return QtCore.QRect(max(0, roi.x()), max(0, roi.y()), max(1, roi.width()), max(1, roi.height())) + x = max(0, min(int(roi.x()), self._image_size.width() - 1)) + y = max(0, min(int(roi.y()), self._image_size.height() - 1)) + w = max(1, min(int(roi.width()), self._image_size.width() - x)) + h = max(1, min(int(roi.height()), self._image_size.height() - y)) + return QtCore.QRect(x, y, w, h) + + def _image_to_widget(self, point: QtCore.QPoint) -> QtCore.QPointF: + if self._image_size.width() <= 0 or self._target_rect.width() <= 0: + return QtCore.QPointF(float(point.x()), float(point.y())) + sx = self._target_rect.width() / max(1, self._image_size.width()) + sy = self._target_rect.height() / max(1, self._image_size.height()) + return QtCore.QPointF(self._target_rect.x() + point.x() * sx, self._target_rect.y() + point.y() * sy) + + def _widget_to_image(self, point: QtCore.QPoint) -> QtCore.QPoint: + if self._image_size.width() <= 0 or self._target_rect.width() <= 0: + return QtCore.QPoint(0, 0) + x = (point.x() - self._target_rect.x()) * self._image_size.width() / max(1, self._target_rect.width()) + y = (point.y() - self._target_rect.y()) * self._image_size.height() / max(1, self._target_rect.height()) + return QtCore.QPoint( + max(0, min(self._image_size.width() - 1, int(round(x)))), + max(0, min(self._image_size.height() - 1, int(round(y)))), + ) + + def paintEvent(self, event: QtGui.QPaintEvent) -> None: + painter = QtGui.QPainter(self) + painter.fillRect(self.rect(), QtGui.QColor("#141922")) + if self._pixmap.isNull(): + painter.setPen(QtGui.QColor("#94a3b8")) + painter.drawText(self.rect(), QtCore.Qt.AlignmentFlag.AlignCenter, "Load a video frame") + return + scaled = self._pixmap.scaled(self.size(), QtCore.Qt.AspectRatioMode.KeepAspectRatio, QtCore.Qt.TransformationMode.SmoothTransformation) + x = (self.width() - scaled.width()) // 2 + y = (self.height() - scaled.height()) // 2 + self._target_rect = QtCore.QRect(x, y, scaled.width(), scaled.height()) + painter.drawPixmap(self._target_rect, scaled) + roi = self._bounded_roi(self._roi) + top_left = self._image_to_widget(roi.topLeft()).toPoint() + bottom_right = self._image_to_widget(roi.bottomRight()).toPoint() + draw_rect = QtCore.QRect(top_left, bottom_right).normalized() + painter.setPen(QtGui.QPen(QtGui.QColor("#f9e154"), 2)) + painter.setBrush(QtGui.QColor(249, 225, 84, 36)) + painter.drawRect(draw_rect) + + def mousePressEvent(self, event: QtGui.QMouseEvent) -> None: + if event.button() == QtCore.Qt.MouseButton.LeftButton and not self._pixmap.isNull(): + self._drag_origin = self._widget_to_image(event.position().toPoint()) + self._roi = QtCore.QRect(self._drag_origin, QtCore.QSize(1, 1)) + self.update() + + def mouseMoveEvent(self, event: QtGui.QMouseEvent) -> None: + if self._drag_origin is None: + return + end = self._widget_to_image(event.position().toPoint()) + self._roi = self._bounded_roi(QtCore.QRect(self._drag_origin, end).normalized()) + self.roiChanged.emit(self.roi()) + self.update() + + def mouseReleaseEvent(self, event: QtGui.QMouseEvent) -> None: + if event.button() == QtCore.Qt.MouseButton.LeftButton and self._drag_origin is not None: + end = self._widget_to_image(event.position().toPoint()) + self._roi = self._bounded_roi(QtCore.QRect(self._drag_origin, end).normalized()) + self._drag_origin = None + self.roiChanged.emit(self.roi()) + self.update() + + +class SyncLedExtractDialog(QtWidgets.QDialog): + def __init__(self, parent: Optional[QtWidgets.QWidget] = None) -> None: + super().__init__(parent) + self.setWindowTitle("Sync LED extraction") + self.resize(820, 620) + self._video_path = "" + self._fps = 30.0 + self._n_frames = 0 + self._updating_roi = False + + root = QtWidgets.QVBoxLayout(self) + root.setContentsMargins(12, 12, 12, 12) + root.setSpacing(10) + + file_row = QtWidgets.QHBoxLayout() + self.edit_video_path = QtWidgets.QLineEdit() + self.edit_video_path.setReadOnly(True) + self.btn_browse_video = QtWidgets.QPushButton("Load video") + self.btn_browse_video.setProperty("class", "compactPrimarySmall") + file_row.addWidget(self.edit_video_path, 1) + file_row.addWidget(self.btn_browse_video) + root.addLayout(file_row) + + body = QtWidgets.QHBoxLayout() + self.preview = SyncRoiPreview() + body.addWidget(self.preview, 1) + form_wrap = QtWidgets.QWidget() + form = QtWidgets.QFormLayout(form_wrap) + form.setRowWrapPolicy(QtWidgets.QFormLayout.RowWrapPolicy.WrapLongRows) + self.combo_channel = QtWidgets.QComboBox() + self.combo_channel.addItems(["Grayscale", "Red", "Green", "Blue"]) + self.spin_x = QtWidgets.QSpinBox(); self.spin_x.setRange(0, 100000) + self.spin_y = QtWidgets.QSpinBox(); self.spin_y.setRange(0, 100000) + self.spin_w = QtWidgets.QSpinBox(); self.spin_w.setRange(1, 100000); self.spin_w.setValue(80) + self.spin_h = QtWidgets.QSpinBox(); self.spin_h.setRange(1, 100000); self.spin_h.setValue(80) + self.spin_start = QtWidgets.QSpinBox(); self.spin_start.setRange(0, 100000000) + self.spin_end = QtWidgets.QSpinBox(); self.spin_end.setRange(0, 100000000) + form.addRow("Channel", self.combo_channel) + form.addRow("ROI x", self.spin_x) + form.addRow("ROI y", self.spin_y) + form.addRow("ROI width", self.spin_w) + form.addRow("ROI height", self.spin_h) + form.addRow("Start frame", self.spin_start) + form.addRow("End frame", self.spin_end) + body.addWidget(form_wrap, 0) + root.addLayout(body, 1) + + self.lbl_info = QtWidgets.QLabel("Select the LED ROI on the preview frame.") + self.lbl_info.setProperty("class", "hint") + self.lbl_info.setWordWrap(True) + root.addWidget(self.lbl_info) + + buttons = QtWidgets.QDialogButtonBox( + QtWidgets.QDialogButtonBox.StandardButton.Ok | QtWidgets.QDialogButtonBox.StandardButton.Cancel + ) + self.btn_ok = buttons.button(QtWidgets.QDialogButtonBox.StandardButton.Ok) + if self.btn_ok is not None: + self.btn_ok.setText("Extract signal") + self.btn_ok.setEnabled(False) + self.btn_ok.setProperty("class", "compactPrimary") + buttons.accepted.connect(self.accept) + buttons.rejected.connect(self.reject) + root.addWidget(buttons) + + self.btn_browse_video.clicked.connect(self._browse_video) + self.preview.roiChanged.connect(self._set_roi_spins) + for spin in (self.spin_x, self.spin_y, self.spin_w, self.spin_h): + spin.valueChanged.connect(self._set_preview_roi_from_spins) + + def _browse_video(self) -> None: + path, _ = QtWidgets.QFileDialog.getOpenFileName( + self, + "Load sync video", + str(Path.home()), + "Video files (*.mp4 *.avi *.mkv *.mov *.m4v *.wmv);;All files (*.*)", + ) + if path: + self.load_video(path) + + def load_video(self, path: str) -> None: + try: + import cv2 + except Exception as exc: + QtWidgets.QMessageBox.warning(self, "Sync LED extraction", f"OpenCV is required for LED extraction:\n{exc}") + return + cap = cv2.VideoCapture(path) + if not cap.isOpened(): + QtWidgets.QMessageBox.warning(self, "Sync LED extraction", "Could not open the selected video.") + return + fps = float(cap.get(cv2.CAP_PROP_FPS) or 30.0) + n_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT) or 0) + ok, frame = cap.read() + cap.release() + if not ok or frame is None: + QtWidgets.QMessageBox.warning(self, "Sync LED extraction", "Could not read a preview frame from the video.") + return + self._video_path = path + self._fps = fps if np.isfinite(fps) and fps > 0 else 30.0 + self._n_frames = max(0, n_frames) + self.edit_video_path.setText(path) + self.preview.set_frame_bgr(frame) + h, w = frame.shape[:2] + default_w = max(8, min(80, w // 6)) + default_h = max(8, min(80, h // 6)) + self._set_roi_spins((max(0, w // 2 - default_w // 2), max(0, h // 2 - default_h // 2), default_w, default_h)) + self.spin_start.setMaximum(max(0, self._n_frames - 1)) + self.spin_end.setMaximum(max(1, self._n_frames)) + self.spin_start.setValue(0) + self.spin_end.setValue(max(1, self._n_frames)) + self.lbl_info.setText(f"{Path(path).name}: {self._n_frames} frame(s), {self._fps:.3f} fps") + if self.btn_ok is not None: + self.btn_ok.setEnabled(True) + + def _set_roi_spins(self, roi: Tuple[int, int, int, int]) -> None: + if self._updating_roi: + return + self._updating_roi = True + try: + x, y, w, h = [int(v) for v in roi] + self.spin_x.setValue(x) + self.spin_y.setValue(y) + self.spin_w.setValue(max(1, w)) + self.spin_h.setValue(max(1, h)) + self.preview.set_roi(x, y, w, h, emit_signal=False) + finally: + self._updating_roi = False + + def _set_preview_roi_from_spins(self) -> None: + if self._updating_roi: + return + self.preview.set_roi(self.spin_x.value(), self.spin_y.value(), self.spin_w.value(), self.spin_h.value(), emit_signal=False) + + def config(self) -> Dict[str, object]: + x, y, w, h = self.preview.roi() + return { + "video_path": self._video_path, + "fps": float(self._fps), + "n_frames": int(self._n_frames), + "roi": (int(x), int(y), int(w), int(h)), + "channel": self.combo_channel.currentText(), + "start_frame": int(self.spin_start.value()), + "end_frame": int(self.spin_end.value()), + } + + def _compute_psth_matrix( t: np.ndarray, y: np.ndarray, @@ -809,6 +1177,7 @@ def __init__(self, parent=None) -> None: self._processed: List[ProcessedTrial] = [] self._dio_cache: Dict[Tuple[str, str], Tuple[np.ndarray, np.ndarray]] = {} # (path,dio)->(t,x) self._behavior_sources: Dict[str, Dict[str, Any]] = {} # stem->behavior data + self._sync_external_sources: Dict[str, Dict[str, Any]] = {} self._continuous_align_rules: Dict[str, Dict[str, object]] = {} self._last_mat: Optional[np.ndarray] = None self._last_tvec: Optional[np.ndarray] = None @@ -826,6 +1195,7 @@ def __init__(self, parent=None) -> None: self._psth_excluded_files: Dict[str, Dict[str, object]] = {} self._sync_results_by_file: Dict[str, Dict[str, object]] = {} self._last_sync_preview: Optional[SyncResult] = None + self._sync_dialog: Optional[QtWidgets.QDialog] = None self._known_dio_channels: List[str] = [] # Per-file / group data for Individual vs Group visual modes self._per_file_mats: Dict[str, Tuple[np.ndarray, np.ndarray]] = {} # file_id -> (tvec, mat) @@ -2050,7 +2420,7 @@ def _set_text_with_banner(text: str) -> None: pass self.lbl_spatial_msg.setText = _set_text_with_banner # type: ignore[assignment] - # ---- Time synchronization panel ---- + # ---- Sync popup ---- self.combo_sync_behavior_file = QtWidgets.QComboBox() self.combo_sync_behavior_file.addItem("Auto-match behavior file", "") _compact_combo(self.combo_sync_behavior_file, min_chars=12) @@ -2092,19 +2462,25 @@ def _set_text_with_banner(text: str) -> None: _compact_combo(self.combo_sync_export_format, min_chars=8) self.btn_sync_refresh = QtWidgets.QPushButton("Refresh sources") self.btn_sync_refresh.setProperty("class", "compactSmall") - self.btn_sync_preview = QtWidgets.QPushButton("Preview selected file") + self.btn_sync_load_file = QtWidgets.QPushButton("Open Sync file") + self.btn_sync_load_file.setProperty("class", "compactSmall") + self.btn_sync_load_processed = QtWidgets.QPushButton("Open processed") + self.btn_sync_load_processed.setProperty("class", "compactSmall") + self.btn_sync_extract_led = QtWidgets.QPushButton("Extract") + self.btn_sync_extract_led.setProperty("class", "compactPrimarySmall") + self.btn_sync_preview = QtWidgets.QPushButton("Preview") self.btn_sync_preview.setProperty("class", "compactSmall") - self.btn_sync_apply_selected = QtWidgets.QPushButton("Apply selected") + self.btn_sync_apply_selected = QtWidgets.QPushButton("Auto-align") self.btn_sync_apply_selected.setProperty("class", "compactPrimarySmall") self.btn_sync_apply_batch = QtWidgets.QPushButton("Apply batch") self.btn_sync_apply_batch.setProperty("class", "compactPrimarySmall") - self.btn_sync_export = QtWidgets.QPushButton("Export aligned files") + self.btn_sync_export = QtWidgets.QPushButton("Export aligned") self.btn_sync_export.setProperty("class", "compactSmall") self.sync_status = _PyberInlineStatus("", "info") self.txt_sync_report = QtWidgets.QTextEdit() self.txt_sync_report.setReadOnly(True) self.txt_sync_report.setMinimumHeight(92) - self.txt_sync_report.setPlaceholderText("Preview or apply time synchronization to see fit diagnostics.") + self.txt_sync_report.setPlaceholderText("Preview or apply Sync to see fit diagnostics.") self.tbl_sync_results = QtWidgets.QTableWidget(0, 6) self.tbl_sync_results.setHorizontalHeaderLabels(["File", "Pairs", "RMS ms", "Median lag ms", "Drift ppm", "Status"]) self.tbl_sync_results.verticalHeader().setVisible(False) @@ -2126,28 +2502,243 @@ def _set_text_with_banner(text: str) -> None: self.plot_sync_residual.setLabel("bottom", "Matched pulse") self.plot_sync_residual.setLabel("left", "Residual (ms)") - grp_sync = QtWidgets.QGroupBox("Time synchronization") - sync_layout = QtWidgets.QVBoxLayout(grp_sync) - sync_layout.setContentsMargins(8, 8, 8, 8) - sync_layout.setSpacing(8) - - sub_sync_sources = _PyberSubsection( - "Sources", - "Pair an external camera/behavior TTL or barcode column with the photometry DIO sync signal.", - ) - f_sync_sources = _form() - _add_row(f_sync_sources, "Behavior file", self.combo_sync_behavior_file) - _add_row(f_sync_sources, "Camera sync", self.combo_sync_camera_column) - _add_row(f_sync_sources, "Camera mode", self.combo_sync_camera_mode) - _add_row(f_sync_sources, "Photometry sync", self.combo_sync_fiber_source) - _add_row(f_sync_sources, "Photometry mode", self.combo_sync_fiber_mode) - sub_sync_sources.add_layout(f_sync_sources) - - sub_sync_fit = _PyberSubsection( - "Alignment", - "Linear regression estimates global clock drift; interpolation follows pulse-to-pulse timing.", - ) - f_sync_fit = _form() + def _sync_panel(name: str = "syncPanel") -> QtWidgets.QFrame: + panel = QtWidgets.QFrame() + panel.setObjectName(name) + return panel + + def _sync_form(panel: QtWidgets.QWidget) -> QtWidgets.QFormLayout: + form = QtWidgets.QFormLayout(panel) + form.setRowWrapPolicy(QtWidgets.QFormLayout.RowWrapPolicy.WrapLongRows) + form.setLabelAlignment(QtCore.Qt.AlignmentFlag.AlignLeft | QtCore.Qt.AlignmentFlag.AlignTop) + form.setHorizontalSpacing(12) + form.setVerticalSpacing(8) + form.setContentsMargins(10, 10, 10, 10) + return form + + self.plot_sync_signal = pg.PlotWidget(title="Sync signal") + _opt_plot(self.plot_sync_signal) + self.plot_sync_signal.setMinimumHeight(260) + self.plot_sync_signal.setLabel("bottom", "Time (s)") + self.plot_sync_signal.setLabel("left", "Signal") + self.lbl_sync_quality = QtWidgets.QLabel("SNR (dB): -\nEdges: -\nEdge interval CV: -\nSaturated frac: -") + self.lbl_sync_quality.setProperty("class", "hint") + self.lbl_sync_quality.setWordWrap(True) + + self.edit_sync_filter = QtWidgets.QLineEdit() + self.edit_sync_filter.setPlaceholderText("Filter videos or Sync files") + self.edit_sync_filter.setClearButtonEnabled(True) + self.list_sync_sources = QtWidgets.QListWidget() + self.list_sync_sources.setSelectionMode(QtWidgets.QAbstractItemView.SelectionMode.SingleSelection) + self.list_sync_sources.setHorizontalScrollBarPolicy(QtCore.Qt.ScrollBarPolicy.ScrollBarAlwaysOff) + self.lbl_sync_source_badge = QtWidgets.QLabel("0 sources") + self.lbl_sync_source_badge.setObjectName("BadgeLabel") + self.lbl_sync_source_badge.setAlignment(QtCore.Qt.AlignmentFlag.AlignCenter) + self.lbl_sync_current_source = QtWidgets.QLabel("No Sync source selected") + self.lbl_sync_current_source.setWordWrap(True) + self.lbl_sync_current_source.setProperty("class", "hint") + + self.sync_led_preview = SyncRoiPreview() + self.sync_led_preview.setMinimumHeight(300) + self.edit_sync_video_path = QtWidgets.QLineEdit() + self.edit_sync_video_path.setReadOnly(True) + self.edit_sync_video_path.setPlaceholderText("No video selected") + self.btn_sync_open_video = QtWidgets.QPushButton("Open video") + self.btn_sync_open_video.setProperty("class", "compactPrimarySmall") + self.btn_sync_open_video_inline = QtWidgets.QPushButton("Browse") + self.btn_sync_open_video_inline.setProperty("class", "compactSmall") + self.btn_sync_refresh_side = QtWidgets.QPushButton("Refresh") + self.btn_sync_load_file_side = QtWidgets.QPushButton("Sync file") + self.btn_sync_load_processed_side = QtWidgets.QPushButton("Processed") + self.btn_sync_open_video_side = QtWidgets.QPushButton("Video") + self.btn_sync_remove_source = QtWidgets.QPushButton("Clear selected") + for btn in ( + self.btn_sync_refresh_side, + self.btn_sync_load_file_side, + self.btn_sync_load_processed_side, + self.btn_sync_open_video_side, + self.btn_sync_remove_source, + ): + btn.setProperty("class", "compactSmall") + self.combo_sync_led_channel = QtWidgets.QComboBox() + self.combo_sync_led_channel.addItems(["Grayscale", "Red", "Green", "Blue"]) + self.spin_sync_led_x = QtWidgets.QSpinBox(); self.spin_sync_led_x.setRange(0, 100000) + self.spin_sync_led_y = QtWidgets.QSpinBox(); self.spin_sync_led_y.setRange(0, 100000) + self.spin_sync_led_w = QtWidgets.QSpinBox(); self.spin_sync_led_w.setRange(1, 100000); self.spin_sync_led_w.setValue(80) + self.spin_sync_led_h = QtWidgets.QSpinBox(); self.spin_sync_led_h.setRange(1, 100000); self.spin_sync_led_h.setValue(80) + self.spin_sync_led_start = QtWidgets.QSpinBox(); self.spin_sync_led_start.setRange(0, 100000000) + self.spin_sync_led_end = QtWidgets.QSpinBox(); self.spin_sync_led_end.setRange(0, 100000000) + self.spin_sync_led_frame = QtWidgets.QSpinBox(); self.spin_sync_led_frame.setRange(0, 100000000) + self.btn_sync_led_prev = QtWidgets.QPushButton("<") + self.btn_sync_led_next = QtWidgets.QPushButton(">") + self.btn_sync_led_fit = QtWidgets.QPushButton("Fit") + for btn in (self.btn_sync_led_prev, self.btn_sync_led_next, self.btn_sync_led_fit): + btn.setProperty("class", "compactSmall") + self.lbl_sync_led_info = QtWidgets.QLabel("Load a video, draw the LED ROI, then extract a Sync signal.") + self.lbl_sync_led_info.setProperty("class", "hint") + self.lbl_sync_led_info.setWordWrap(True) + self._sync_led_video_path = "" + self._sync_led_fps = 30.0 + self._sync_led_n_frames = 0 + self._sync_led_updating_roi = False + + sync_root = QtWidgets.QWidget() + sync_root_layout = QtWidgets.QVBoxLayout(sync_root) + sync_root_layout.setContentsMargins(0, 0, 0, 0) + sync_root_layout.setSpacing(0) + + sync_menu = QtWidgets.QMenuBar() + menu_file = sync_menu.addMenu("File") + menu_extract = sync_menu.addMenu("Extract") + menu_align = sync_menu.addMenu("Align") + menu_export = sync_menu.addMenu("Export") + menu_view = sync_menu.addMenu("View") + sync_menu.addMenu("Help") + act_sync_open_csv = menu_file.addAction("Open Sync CSV/H5") + act_sync_open_processed = menu_file.addAction("Open processed file") + act_sync_open_video = menu_file.addAction("Open video") + act_sync_remove_source = menu_file.addAction("Clear selected source") + act_sync_extract = menu_extract.addAction("Extract LED signal") + act_sync_preview = menu_align.addAction("Preview selected") + act_sync_apply = menu_align.addAction("Auto-align selected") + act_sync_apply_batch = menu_align.addAction("Apply batch") + act_sync_export = menu_export.addAction("Export aligned files") + act_sync_refresh = menu_view.addAction("Refresh sources") + act_sync_open_csv.triggered.connect(self._load_sync_signal_files) + act_sync_open_processed.triggered.connect(self._load_processed_files) + act_sync_open_video.triggered.connect(self._sync_browse_led_video) + act_sync_remove_source.triggered.connect(self._remove_selected_sync_source) + act_sync_extract.triggered.connect(self._open_sync_led_extract_dialog) + act_sync_preview.triggered.connect(lambda _checked=False: self._preview_time_sync()) + act_sync_apply.triggered.connect(lambda _checked=False: self._apply_time_sync_selected()) + act_sync_apply_batch.triggered.connect(lambda _checked=False: self._apply_time_sync_batch()) + act_sync_export.triggered.connect(lambda _checked=False: self._export_sync_aligned_files()) + act_sync_refresh.triggered.connect(self._refresh_sync_sources) + sync_root_layout.addWidget(sync_menu, 0) + + sync_toolbar = QtWidgets.QFrame() + sync_toolbar.setObjectName("transportBar") + sync_toolbar_layout = QtWidgets.QHBoxLayout(sync_toolbar) + sync_toolbar_layout.setContentsMargins(12, 8, 12, 8) + sync_toolbar_layout.setSpacing(8) + sync_toolbar_layout.addWidget(self.btn_sync_load_file) + sync_toolbar_layout.addWidget(self.btn_sync_load_processed) + sync_toolbar_layout.addWidget(self.btn_sync_open_video) + sync_toolbar_layout.addSpacing(10) + sync_toolbar_layout.addWidget(self.btn_sync_extract_led) + sync_toolbar_layout.addWidget(self.btn_sync_preview) + sync_toolbar_layout.addWidget(self.btn_sync_apply_selected) + sync_toolbar_layout.addWidget(self.btn_sync_apply_batch) + sync_toolbar_layout.addSpacing(10) + sync_toolbar_layout.addWidget(self.combo_sync_export_format) + sync_toolbar_layout.addWidget(self.btn_sync_export) + sync_toolbar_layout.addStretch(1) + sync_root_layout.addWidget(sync_toolbar, 0) + + main_splitter = QtWidgets.QSplitter(QtCore.Qt.Orientation.Horizontal) + main_splitter.setChildrenCollapsible(False) + + left_panel = _sync_panel("SidePanel") + left_panel.setMinimumWidth(240) + left_panel.setMaximumWidth(320) + left_layout = QtWidgets.QVBoxLayout(left_panel) + left_layout.setContentsMargins(10, 10, 10, 10) + left_layout.setSpacing(8) + left_head = QtWidgets.QHBoxLayout() + left_title = QtWidgets.QLabel("Files") + left_title.setObjectName("panelTitle") + left_head.addWidget(left_title) + left_head.addStretch(1) + left_head.addWidget(self.lbl_sync_source_badge) + left_layout.addLayout(left_head) + left_layout.addWidget(self.edit_sync_filter) + left_layout.addWidget(self.list_sync_sources, 1) + left_actions = QtWidgets.QGridLayout() + left_actions.setContentsMargins(0, 0, 0, 0) + left_actions.setHorizontalSpacing(6) + left_actions.setVerticalSpacing(6) + left_actions.addWidget(self.btn_sync_refresh_side, 0, 0) + left_actions.addWidget(self.btn_sync_load_file_side, 0, 1) + left_actions.addWidget(self.btn_sync_load_processed_side, 1, 0) + left_actions.addWidget(self.btn_sync_open_video_side, 1, 1) + left_actions.addWidget(self.btn_sync_remove_source, 2, 0, 1, 2) + left_layout.addLayout(left_actions) + left_layout.addWidget(self.lbl_sync_current_source) + + center_panel = QtWidgets.QWidget() + center_panel.setMinimumWidth(440) + center_layout = QtWidgets.QVBoxLayout(center_panel) + center_layout.setContentsMargins(10, 10, 10, 10) + center_layout.setSpacing(8) + video_row = QtWidgets.QHBoxLayout() + video_row.setContentsMargins(0, 0, 0, 0) + video_row.setSpacing(6) + video_row.addWidget(self.edit_sync_video_path, 1) + video_row.addWidget(self.btn_sync_open_video_inline, 0) + center_layout.addLayout(video_row) + center_layout.addWidget(self.sync_led_preview, 1) + frame_row = QtWidgets.QHBoxLayout() + frame_row.setContentsMargins(0, 0, 0, 0) + frame_row.setSpacing(6) + frame_row.addWidget(self.btn_sync_led_prev) + frame_row.addWidget(self.btn_sync_led_next) + frame_row.addWidget(self.btn_sync_led_fit) + frame_row.addSpacing(8) + frame_row.addWidget(QtWidgets.QLabel("Frame:")) + frame_row.addWidget(self.spin_sync_led_frame) + frame_row.addStretch(1) + center_layout.addLayout(frame_row) + roi_panel = _sync_panel() + roi_form = _sync_form(roi_panel) + roi_widget = QtWidgets.QWidget() + roi_grid = QtWidgets.QGridLayout(roi_widget) + roi_grid.setContentsMargins(0, 0, 0, 0) + roi_grid.setHorizontalSpacing(6) + roi_grid.setVerticalSpacing(6) + roi_grid.addWidget(QtWidgets.QLabel("x"), 0, 0) + roi_grid.addWidget(self.spin_sync_led_x, 0, 1) + roi_grid.addWidget(QtWidgets.QLabel("y"), 0, 2) + roi_grid.addWidget(self.spin_sync_led_y, 0, 3) + roi_grid.addWidget(QtWidgets.QLabel("w"), 1, 0) + roi_grid.addWidget(self.spin_sync_led_w, 1, 1) + roi_grid.addWidget(QtWidgets.QLabel("h"), 1, 2) + roi_grid.addWidget(self.spin_sync_led_h, 1, 3) + range_widget = QtWidgets.QWidget() + range_row = QtWidgets.QHBoxLayout(range_widget) + range_row.setContentsMargins(0, 0, 0, 0) + range_row.setSpacing(6) + range_row.addWidget(QtWidgets.QLabel("Start frame")) + range_row.addWidget(self.spin_sync_led_start) + range_row.addWidget(QtWidgets.QLabel("End frame")) + range_row.addWidget(self.spin_sync_led_end) + roi_form.addRow("ROI (x,y,w,h)", roi_widget) + roi_form.addRow("Channel", self.combo_sync_led_channel) + roi_form.addRow("Extract range", range_widget) + center_layout.addWidget(roi_panel, 0) + center_layout.addWidget(self.lbl_sync_led_info, 0) + + right_tabs = QtWidgets.QTabWidget() + right_tabs.setDocumentMode(True) + right_tabs.setMinimumWidth(560) + + signal_page = QtWidgets.QWidget() + signal_layout = QtWidgets.QVBoxLayout(signal_page) + signal_layout.setContentsMargins(10, 10, 10, 10) + signal_layout.setSpacing(8) + signal_subtabs = QtWidgets.QTabWidget() + signal_subtabs.setDocumentMode(True) + camera_panel = _sync_panel() + f_sync_camera = _sync_form(camera_panel) + _add_row(f_sync_camera, "Source", self.combo_sync_behavior_file) + _add_row(f_sync_camera, "Signal", self.combo_sync_camera_column) + _add_row(f_sync_camera, "Mode", self.combo_sync_camera_mode) + photometry_panel = _sync_panel() + f_sync_photometry = _sync_form(photometry_panel) + _add_row(f_sync_photometry, "Target signal", self.combo_sync_fiber_source) + _add_row(f_sync_photometry, "Mode", self.combo_sync_fiber_mode) + _add_row(f_sync_photometry, "Refresh", self.btn_sync_refresh) + threshold_panel = _sync_panel() + f_sync_fit = _sync_form(threshold_panel) _add_row(f_sync_fit, "Method", self.combo_sync_method) _add_row(f_sync_fit, "Threshold", self.cb_sync_auto_threshold) _add_row(f_sync_fit, "Manual level", self.spin_sync_threshold) @@ -2155,36 +2746,76 @@ def _set_text_with_banner(text: str) -> None: _add_row(f_sync_fit, "Pulse offset", self.spin_sync_max_offset) _add_row(f_sync_fit, "Use result", self.cb_sync_use_aligned) _add_row(f_sync_fit, "Auto recompute", self.cb_sync_auto_recompute) - sub_sync_fit.add_layout(f_sync_fit) - - sub_sync_batch = _PyberSubsection( - "Batch and export", - "The selected settings are reused across matched files in the queue and saved in the project.", - ) - sync_btn_grid = QtWidgets.QGridLayout() - sync_btn_grid.setContentsMargins(0, 0, 0, 0) - sync_btn_grid.setHorizontalSpacing(6) - sync_btn_grid.setVerticalSpacing(6) - sync_btn_grid.addWidget(self.btn_sync_refresh, 0, 0) - sync_btn_grid.addWidget(self.btn_sync_preview, 0, 1) - sync_btn_grid.addWidget(self.btn_sync_apply_selected, 1, 0) - sync_btn_grid.addWidget(self.btn_sync_apply_batch, 1, 1) - sync_export_row = QtWidgets.QHBoxLayout() - sync_export_row.setContentsMargins(0, 0, 0, 0) - sync_export_row.setSpacing(6) - sync_export_row.addWidget(self.combo_sync_export_format, 1) - sync_export_row.addWidget(self.btn_sync_export, 1) - sub_sync_batch.add_layout(sync_btn_grid) - sub_sync_batch.add_layout(sync_export_row) - - sync_layout.addWidget(sub_sync_sources) - sync_layout.addWidget(sub_sync_fit) - sync_layout.addWidget(sub_sync_batch) - sync_layout.addWidget(self.sync_status) - sync_layout.addWidget(self.txt_sync_report) - sync_layout.addWidget(self.tbl_sync_results) - sync_layout.addWidget(self.plot_sync_map) - sync_layout.addWidget(self.plot_sync_residual) + signal_subtabs.addTab(camera_panel, "Camera") + signal_subtabs.addTab(photometry_panel, "Photometry") + signal_subtabs.addTab(threshold_panel, "Threshold") + signal_layout.addWidget(signal_subtabs, 0) + signal_layout.addWidget(self.plot_sync_signal, 1) + quality_panel = _sync_panel() + quality_layout = QtWidgets.QVBoxLayout(quality_panel) + quality_layout.setContentsMargins(10, 8, 10, 8) + quality_layout.addWidget(self.lbl_sync_quality) + signal_layout.addWidget(quality_panel, 0) + + alignment_page = QtWidgets.QWidget() + alignment_layout = QtWidgets.QVBoxLayout(alignment_page) + alignment_layout.setContentsMargins(10, 10, 10, 10) + alignment_layout.setSpacing(8) + alignment_subtabs = QtWidgets.QTabWidget() + alignment_subtabs.setDocumentMode(True) + event_page = QtWidgets.QWidget() + event_layout = QtWidgets.QVBoxLayout(event_page) + event_layout.setContentsMargins(0, 0, 0, 0) + self.plot_sync_map.setMinimumHeight(360) + event_layout.addWidget(self.plot_sync_map, 1) + residual_page = QtWidgets.QWidget() + residual_layout = QtWidgets.QVBoxLayout(residual_page) + residual_layout.setContentsMargins(0, 0, 0, 0) + self.plot_sync_residual.setMinimumHeight(360) + residual_layout.addWidget(self.plot_sync_residual, 1) + report_page = QtWidgets.QWidget() + report_layout = QtWidgets.QVBoxLayout(report_page) + report_layout.setContentsMargins(0, 0, 0, 0) + self.txt_sync_report.setMinimumHeight(220) + report_layout.addWidget(self.sync_status, 0) + report_layout.addWidget(self.txt_sync_report, 1) + files_page = QtWidgets.QWidget() + files_layout = QtWidgets.QVBoxLayout(files_page) + files_layout.setContentsMargins(0, 0, 0, 0) + self.tbl_sync_results.setMinimumHeight(260) + files_layout.addWidget(self.tbl_sync_results, 1) + alignment_subtabs.addTab(event_page, "Event Map") + alignment_subtabs.addTab(residual_page, "Residuals") + alignment_subtabs.addTab(report_page, "Report") + alignment_subtabs.addTab(files_page, "Files") + alignment_layout.addWidget(alignment_subtabs, 1) + + right_tabs.addTab(signal_page, "Signal") + right_tabs.addTab(alignment_page, "Alignment") + + main_splitter.addWidget(left_panel) + main_splitter.addWidget(center_panel) + main_splitter.addWidget(right_tabs) + main_splitter.setStretchFactor(0, 0) + main_splitter.setStretchFactor(1, 1) + main_splitter.setStretchFactor(2, 2) + main_splitter.setSizes([260, 520, 760]) + sync_root_layout.addWidget(main_splitter, 1) + + self._sync_content = sync_root + self._sync_dialog = QtWidgets.QDialog(self) + self._sync_dialog.setWindowTitle("Sync") + self._sync_dialog.setWindowModality(QtCore.Qt.WindowModality.NonModal) + self._sync_dialog.resize(1500, 880) + sync_dialog_layout = QtWidgets.QVBoxLayout(self._sync_dialog) + sync_dialog_layout.setContentsMargins(0, 0, 0, 0) + sync_dialog_layout.setSpacing(0) + sync_dialog_layout.addWidget(self._sync_content, 1) + + self._sync_main_splitter = main_splitter + self._sync_right_tabs = right_tabs + self._sync_signal_tabs = signal_subtabs + self._sync_alignment_tabs = alignment_subtabs self.section_setup = QtWidgets.QWidget() setup_layout = QtWidgets.QVBoxLayout(self.section_setup) @@ -2242,7 +2873,10 @@ def _set_text_with_banner(text: str) -> None: sync_section_layout = QtWidgets.QVBoxLayout(self.section_sync) sync_section_layout.setContentsMargins(6, 6, 6, 6) sync_section_layout.setSpacing(8) - sync_section_layout.addWidget(grp_sync) + self.btn_open_sync_dialog = QtWidgets.QPushButton("Open Sync") + self.btn_open_sync_dialog.setProperty("class", "compactPrimarySmall") + sync_section_layout.addWidget(self.btn_open_sync_dialog) + sync_section_layout.addStretch(1) self.section_temporal = TemporalModelingWidget() self.section_temporal.statusMessage.connect( @@ -2868,12 +3502,41 @@ def _set_text_with_banner(text: str) -> None: self.btn_load_processed.clicked.connect(self._load_processed_files) self.btn_load_processed_single.clicked.connect(self._load_processed_files_single) self.btn_sync_refresh.clicked.connect(self._refresh_sync_sources) + self.btn_sync_refresh_side.clicked.connect(self._refresh_sync_sources) + self.btn_sync_load_file.clicked.connect(self._load_sync_signal_files) + self.btn_sync_load_file_side.clicked.connect(self._load_sync_signal_files) + self.btn_sync_load_processed.clicked.connect(self._load_processed_files) + self.btn_sync_load_processed_side.clicked.connect(self._load_processed_files) + self.btn_sync_open_video.clicked.connect(self._sync_browse_led_video) + self.btn_sync_open_video_inline.clicked.connect(self._sync_browse_led_video) + self.btn_sync_open_video_side.clicked.connect(self._sync_browse_led_video) + self.btn_sync_remove_source.clicked.connect(self._remove_selected_sync_source) + self.btn_sync_extract_led.clicked.connect(self._open_sync_led_extract_dialog) + self.btn_open_sync_dialog.clicked.connect(self._open_sync_dialog) self.btn_sync_preview.clicked.connect(lambda _checked=False: self._preview_time_sync()) self.btn_sync_apply_selected.clicked.connect(lambda _checked=False: self._apply_time_sync_selected()) self.btn_sync_apply_batch.clicked.connect(lambda _checked=False: self._apply_time_sync_batch()) self.btn_sync_export.clicked.connect(lambda _checked=False: self._export_sync_aligned_files()) self.cb_sync_auto_threshold.toggled.connect(self._on_sync_auto_threshold_toggled) + self.cb_sync_auto_threshold.toggled.connect(lambda _checked=False: self._refresh_sync_signal_preview()) self.combo_sync_behavior_file.currentIndexChanged.connect(self._refresh_sync_camera_columns) + self.combo_sync_behavior_file.currentIndexChanged.connect(lambda _idx=0: self._refresh_sync_source_list()) + self.combo_sync_behavior_file.currentIndexChanged.connect(lambda _idx=0: self._refresh_sync_signal_preview()) + self.combo_sync_camera_column.currentIndexChanged.connect(lambda _idx=0: self._refresh_sync_signal_preview()) + self.combo_sync_camera_mode.currentIndexChanged.connect(lambda _idx=0: self._refresh_sync_signal_preview()) + self.combo_sync_fiber_source.currentIndexChanged.connect(lambda _idx=0: self._refresh_sync_signal_preview()) + self.combo_sync_fiber_mode.currentIndexChanged.connect(lambda _idx=0: self._refresh_sync_signal_preview()) + self.spin_sync_threshold.valueChanged.connect(lambda _value=0.0: self._refresh_sync_signal_preview()) + self.spin_sync_min_interval.valueChanged.connect(lambda _value=0.0: self._refresh_sync_signal_preview()) + self.edit_sync_filter.textChanged.connect(lambda _text="": self._refresh_sync_source_list()) + self.list_sync_sources.itemSelectionChanged.connect(self._on_sync_source_list_selected) + self.sync_led_preview.roiChanged.connect(self._sync_led_roi_to_spins) + for spin in (self.spin_sync_led_x, self.spin_sync_led_y, self.spin_sync_led_w, self.spin_sync_led_h): + spin.valueChanged.connect(self._sync_led_spins_to_roi) + self.spin_sync_led_frame.valueChanged.connect(self._sync_seek_led_frame) + self.btn_sync_led_prev.clicked.connect(lambda _checked=False: self.spin_sync_led_frame.setValue(max(0, self.spin_sync_led_frame.value() - 1))) + self.btn_sync_led_next.clicked.connect(lambda _checked=False: self.spin_sync_led_frame.setValue(self.spin_sync_led_frame.value() + 1)) + self.btn_sync_led_fit.clicked.connect(self._sync_led_fit_roi) self.cb_sync_use_aligned.toggled.connect(lambda _checked=False: self._on_sync_use_aligned_changed()) self.list_preprocessed.filesDropped.connect(self._on_preprocessed_files_dropped) self.list_preprocessed.orderChanged.connect(self._sync_processed_order_from_list) @@ -3095,7 +3758,6 @@ def _theme_plot_background(self) -> Tuple[int, int, int]: def _section_widget_map(self) -> Dict[str, Tuple[str, QtWidgets.QWidget]]: return { "setup": ("Setup", self.section_setup), - "sync": ("Time Synchronization", self.section_sync), "psth": ("PSTH", self.section_psth), "spatial": ("Spatial", self.section_spatial), "temporal": ("Temporal Modeling", self.section_temporal), @@ -3513,7 +4175,6 @@ def _setup_section_popups(self) -> None: self._dock_host = host section_map: Dict[str, Tuple[str, QtWidgets.QWidget]] = { "setup": ("Setup", self.section_setup), - "sync": ("Time Synchronization", self.section_sync), "psth": ("PSTH", self.section_psth), "signal": ("Signal Event Analyzer", self.section_signal), "behavior": ("Behavior Analysis", self.section_behavior), @@ -3693,6 +4354,11 @@ def _force_hide_post_drawer_initially(self) -> None: pass def _toggle_section_popup(self, key: str, checked: bool) -> None: + if key == "sync": + if checked: + self._open_sync_dialog() + self._set_section_button_checked("sync", False) + return if self._use_pg_dockarea_layout: self._setup_dockarea_sections() dock = self._dockarea_dock(key) @@ -4497,6 +5163,611 @@ def _apply_continuous_alignment_rule(self, config: Dict[str, object]) -> int: self._refresh_sync_sources() return updated + def _open_sync_dialog(self) -> None: + if self._sync_dialog is None: + return + self._refresh_sync_sources() + self._refresh_sync_source_list() + self._refresh_sync_signal_preview() + self._update_data_availability() + self._sync_dialog.show() + self._sync_dialog.raise_() + self._sync_dialog.activateWindow() + + def _add_sync_external_source( + self, + key: str, + time: np.ndarray, + columns: Dict[str, np.ndarray], + *, + source_path: str = "", + kind: str = "signal_file", + time_col: str = "", + ) -> str: + clean_key = str(key or "").strip() or "sync_source" + base_key = clean_key + i = 2 + while clean_key in self._sync_external_sources: + clean_key = f"{base_key}_{i}" + i += 1 + t = np.asarray(time, float).reshape(-1) + numeric_columns: Dict[str, np.ndarray] = {} + for name, values in (columns or {}).items(): + arr = np.asarray(values, float).reshape(-1) + n = min(t.size, arr.size) + if n >= 2: + numeric_columns[str(name)] = arr[:n] + if not numeric_columns: + raise ValueError("No numeric sync signal columns were found.") + self._sync_external_sources[clean_key] = { + "time": t, + "columns": numeric_columns, + "source_path": str(source_path or ""), + "kind": str(kind or "signal_file"), + "time_col": str(time_col or ""), + } + self._refresh_sync_sources() + idx = self.combo_sync_behavior_file.findData(f"external::{clean_key}") + if idx >= 0: + self.combo_sync_behavior_file.setCurrentIndex(idx) + self._refresh_sync_camera_columns() + if not self._autosave_restoring: + self._project_dirty = True + self._update_data_availability() + return clean_key + + def _load_sync_signal_files(self) -> None: + paths, _ = QtWidgets.QFileDialog.getOpenFileNames( + self, + "Load Sync signal, barcode, or processed file", + self._export_start_dir(), + "Sync files (*.csv *.txt *.tsv *.xlsx *.xls *.h5 *.hdf5);;All files (*.*)", + ) + if not paths: + return + loaded = 0 + errors: List[str] = [] + for path in paths: + try: + df = _load_sync_table(path) + time, time_col = _time_array_for_signal_table(df) + cols = _numeric_columns_from_df(df, time_col=None if time_col == "sample_index" else time_col) + key = os.path.splitext(os.path.basename(path))[0] + self._add_sync_external_source(key, time, cols, source_path=path, kind="signal_file", time_col=time_col) + loaded += 1 + except Exception as exc: + errors.append(f"{os.path.basename(path)}: {exc}") + if errors: + QtWidgets.QMessageBox.warning( + self, + "Sync", + "Some Sync files could not be loaded:\n" + "\n".join(errors[:8]), + ) + if loaded: + self.sync_status.set(f"Loaded {loaded} Sync signal file(s).", "ok") + self.statusUpdate.emit(f"Loaded {loaded} Sync signal file(s).", 5000) + + def _remove_selected_sync_source(self) -> None: + if not hasattr(self, "list_sync_sources"): + return + item = self.list_sync_sources.currentItem() + data = item.data(QtCore.Qt.ItemDataRole.UserRole) if item is not None else None + if not isinstance(data, str) or not data: + return + removed = "" + if data.startswith("external::"): + key = data.split("::", 1)[1] + if key in self._sync_external_sources: + self._sync_external_sources.pop(key, None) + removed = key + if self.combo_sync_behavior_file.currentData() == data: + self.combo_sync_behavior_file.setCurrentIndex(0) + if isinstance(self.combo_sync_fiber_source.currentData(), str) and str(self.combo_sync_fiber_source.currentData()).startswith(f"external::{key}::"): + self.combo_sync_fiber_source.setCurrentIndex(0) + elif data in self._behavior_sources: + self._behavior_sources.pop(data, None) + removed = data + if self.combo_sync_behavior_file.currentData() == data: + self.combo_sync_behavior_file.setCurrentIndex(0) + self._refresh_behavior_list() + if not removed: + return + if not self._autosave_restoring: + self._project_dirty = True + self._refresh_sync_sources() + self._update_data_availability() + self.sync_status.set(f"Removed Sync source: {removed}", "ok") + self.statusUpdate.emit(f"Removed Sync source: {removed}", 4000) + + def _refresh_sync_source_list(self) -> None: + if not hasattr(self, "list_sync_sources"): + return + current_data = self.combo_sync_behavior_file.currentData() if hasattr(self, "combo_sync_behavior_file") else "" + query = self.edit_sync_filter.text().strip().lower() if hasattr(self, "edit_sync_filter") else "" + rows: List[Tuple[str, str, str, str]] = [] + for stem, info in (self._behavior_sources or {}).items(): + path = str((info or {}).get("path", "") or "") + detail = os.path.basename(path) if path else "behavior source" + rows.append(("Behavior", str(stem), str(stem), detail)) + for key, info in (self._sync_external_sources or {}).items(): + kind = str((info or {}).get("kind", "") or "") + label_kind = "LED video" if kind == "led_video" else "Sync file" + path = str((info or {}).get("source_path", "") or "") + detail = os.path.basename(path) if path else label_kind + rows.append((label_kind, str(key), f"external::{key}", detail)) + + filtered: List[Tuple[str, str, str, str]] = [] + for kind, label, data, detail in rows: + haystack = f"{kind} {label} {detail}".lower() + if not query or query in haystack: + filtered.append((kind, label, data, detail)) + + self.list_sync_sources.blockSignals(True) + try: + self.list_sync_sources.clear() + selected_row = -1 + for row, (kind, label, data, detail) in enumerate(filtered): + item = QtWidgets.QListWidgetItem(f"{label}\n{kind}: {detail}") + item.setData(QtCore.Qt.ItemDataRole.UserRole, data) + item.setToolTip(detail) + self.list_sync_sources.addItem(item) + if isinstance(current_data, str) and data == current_data: + selected_row = row + if not filtered: + item = QtWidgets.QListWidgetItem("No Sync sources loaded") + item.setFlags(item.flags() & ~QtCore.Qt.ItemFlag.ItemIsSelectable & ~QtCore.Qt.ItemFlag.ItemIsEnabled) + self.list_sync_sources.addItem(item) + elif selected_row >= 0: + self.list_sync_sources.setCurrentRow(selected_row) + finally: + self.list_sync_sources.blockSignals(False) + + if hasattr(self, "lbl_sync_source_badge"): + self.lbl_sync_source_badge.setText(f"{len(rows)} source(s)") + if hasattr(self, "lbl_sync_current_source"): + item = self.list_sync_sources.currentItem() + if item is not None and item.data(QtCore.Qt.ItemDataRole.UserRole): + self.lbl_sync_current_source.setText(item.text().replace("\n", " | ")) + elif current_data: + self.lbl_sync_current_source.setText(str(current_data)) + else: + self.lbl_sync_current_source.setText("Auto-match Sync source") + + def _on_sync_source_list_selected(self) -> None: + if not hasattr(self, "list_sync_sources") or not hasattr(self, "combo_sync_behavior_file"): + return + item = self.list_sync_sources.currentItem() + if item is None: + return + data = item.data(QtCore.Qt.ItemDataRole.UserRole) + if not isinstance(data, str) or not data: + return + idx = self.combo_sync_behavior_file.findData(data) + if idx >= 0 and idx != self.combo_sync_behavior_file.currentIndex(): + self.combo_sync_behavior_file.setCurrentIndex(idx) + if hasattr(self, "lbl_sync_current_source"): + self.lbl_sync_current_source.setText(item.text().replace("\n", " | ")) + self._refresh_sync_camera_columns() + self._refresh_sync_signal_preview() + + def _sync_browse_led_video(self) -> None: + path, _ = QtWidgets.QFileDialog.getOpenFileName( + self, + "Open sync video", + self._export_start_dir(), + "Video files (*.mp4 *.avi *.mkv *.mov *.m4v *.wmv);;All files (*.*)", + ) + if path: + self._sync_load_led_video(path) + + def _sync_load_led_video(self, path: str) -> None: + try: + import cv2 + except Exception as exc: + QtWidgets.QMessageBox.warning(self, "Sync", f"OpenCV is required for LED extraction:\n{exc}") + return + cap = cv2.VideoCapture(path) + if not cap.isOpened(): + QtWidgets.QMessageBox.warning(self, "Sync", "Could not open the selected video.") + return + fps = float(cap.get(cv2.CAP_PROP_FPS) or 30.0) + n_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT) or 0) + ok, frame = cap.read() + cap.release() + if not ok or frame is None: + QtWidgets.QMessageBox.warning(self, "Sync", "Could not read a preview frame from the selected video.") + return + + self._sync_led_video_path = path + self._sync_led_fps = fps if np.isfinite(fps) and fps > 0 else 30.0 + self._sync_led_n_frames = max(0, n_frames) + self.edit_sync_video_path.setText(path) + self.sync_led_preview.set_frame_bgr(frame) + h, w = frame.shape[:2] + default_w = max(8, min(80, w // 6)) + default_h = max(8, min(80, h // 6)) + self._sync_led_roi_to_spins( + (max(0, w // 2 - default_w // 2), max(0, h // 2 - default_h // 2), default_w, default_h) + ) + max_frame = max(0, self._sync_led_n_frames - 1) + for spin in (self.spin_sync_led_start, self.spin_sync_led_frame): + spin.blockSignals(True) + spin.setMaximum(max_frame) + spin.setValue(0) + spin.blockSignals(False) + self.spin_sync_led_end.blockSignals(True) + self.spin_sync_led_end.setMaximum(max(1, self._sync_led_n_frames)) + self.spin_sync_led_end.setValue(max(1, self._sync_led_n_frames)) + self.spin_sync_led_end.blockSignals(False) + self.lbl_sync_led_info.setText(f"{Path(path).name}: {self._sync_led_n_frames} frame(s), {self._sync_led_fps:.3f} fps") + + def _sync_seek_led_frame(self, *_args) -> None: + path = str(getattr(self, "_sync_led_video_path", "") or "") + if not path: + return + try: + import cv2 + except Exception: + return + frame_idx = int(self.spin_sync_led_frame.value()) + cap = cv2.VideoCapture(path) + if not cap.isOpened(): + return + try: + cap.set(cv2.CAP_PROP_POS_FRAMES, max(0, frame_idx)) + ok, frame = cap.read() + finally: + cap.release() + if ok and frame is not None: + self.sync_led_preview.set_frame_bgr(frame) + self.lbl_sync_led_info.setText( + f"{Path(path).name}: frame {frame_idx} / {max(0, int(getattr(self, '_sync_led_n_frames', 0)) - 1)}" + ) + + def _sync_led_fit_roi(self) -> None: + size = getattr(self.sync_led_preview, "_image_size", QtCore.QSize(0, 0)) + if size.width() <= 0 or size.height() <= 0: + return + w = max(8, min(80, size.width() // 6)) + h = max(8, min(80, size.height() // 6)) + self._sync_led_roi_to_spins((max(0, size.width() // 2 - w // 2), max(0, size.height() // 2 - h // 2), w, h)) + + def _sync_led_roi_to_spins(self, roi: Tuple[int, int, int, int]) -> None: + if getattr(self, "_sync_led_updating_roi", False): + return + self._sync_led_updating_roi = True + try: + x, y, w, h = [int(v) for v in roi] + self.spin_sync_led_x.setValue(max(0, x)) + self.spin_sync_led_y.setValue(max(0, y)) + self.spin_sync_led_w.setValue(max(1, w)) + self.spin_sync_led_h.setValue(max(1, h)) + self.sync_led_preview.set_roi(x, y, w, h, emit_signal=False) + finally: + self._sync_led_updating_roi = False + + def _sync_led_spins_to_roi(self, *_args) -> None: + if getattr(self, "_sync_led_updating_roi", False): + return + self.sync_led_preview.set_roi( + self.spin_sync_led_x.value(), + self.spin_sync_led_y.value(), + self.spin_sync_led_w.value(), + self.spin_sync_led_h.value(), + emit_signal=False, + ) + + def _sync_led_config_from_workspace(self) -> Dict[str, object]: + path = str(getattr(self, "_sync_led_video_path", "") or "") + if not path: + raise ValueError("Open a video first.") + x, y, w, h = self.sync_led_preview.roi() + return { + "video_path": path, + "fps": float(getattr(self, "_sync_led_fps", 30.0) or 30.0), + "n_frames": int(getattr(self, "_sync_led_n_frames", 0) or 0), + "roi": (int(x), int(y), int(w), int(h)), + "channel": self.combo_sync_led_channel.currentText(), + "start_frame": int(self.spin_sync_led_start.value()), + "end_frame": int(self.spin_sync_led_end.value()), + } + + def _open_sync_led_extract_dialog(self) -> None: + if not str(getattr(self, "_sync_led_video_path", "") or ""): + self._sync_browse_led_video() + if not str(getattr(self, "_sync_led_video_path", "") or ""): + return + try: + key = self._extract_led_sync_source(self._sync_led_config_from_workspace()) + except Exception as exc: + QtWidgets.QMessageBox.warning(self, "Sync", f"LED extraction failed:\n{exc}") + return + self.sync_status.set(f"Extracted LED Sync source: {key}", "ok") + self.statusUpdate.emit(f"Extracted LED Sync source: {key}", 5000) + self._refresh_sync_source_list() + self._refresh_sync_signal_preview() + + def _refresh_sync_signal_preview(self) -> None: + if not hasattr(self, "plot_sync_signal"): + return + self.plot_sync_signal.clear() + if not hasattr(self, "combo_sync_camera_column"): + return + + try: + self.plot_sync_signal.setTitle("Reference and photometry sync signals") + self.plot_sync_signal.setLabel("left", "Normalized signal") + threshold = None if self.cb_sync_auto_threshold.isChecked() else float(self.spin_sync_threshold.value()) + min_interval = float(self.spin_sync_min_interval.value()) + quality_lines: List[str] = [] + plotted_rows: List[Tuple[float, str]] = [] + + def _clean_trace(t_in: np.ndarray, x_in: np.ndarray) -> Tuple[np.ndarray, np.ndarray]: + t_arr = np.asarray(t_in, float) + x_arr = np.asarray(x_in, float) + n = min(t_arr.size, x_arr.size) + if n < 2: + return np.array([], float), np.array([], float) + t_arr = t_arr[:n] + x_arr = x_arr[:n] + good = np.isfinite(t_arr) & np.isfinite(x_arr) + t_arr = t_arr[good] + x_arr = x_arr[good] + if t_arr.size < 2: + return np.array([], float), np.array([], float) + order = np.argsort(t_arr) + return t_arr[order], x_arr[order] + + def _normalize_trace(x_arr: np.ndarray) -> Tuple[np.ndarray, float, float]: + finite = np.asarray(x_arr[np.isfinite(x_arr)], float) + if finite.size == 0: + return np.zeros_like(x_arr, dtype=float), 0.0, 1.0 + lo = float(np.nanmin(finite)) + hi = float(np.nanmax(finite)) + span = hi - lo + if not np.isfinite(span) or span <= 0: + return np.zeros_like(x_arr, dtype=float), lo, hi + return (x_arr - lo) / span, lo, hi + + def _event_guides(events: np.ndarray, offset: float, color: Tuple[int, int, int, int]) -> None: + shown = np.asarray(events[:400], float) + if shown.size == 0: + return + xs = np.empty(shown.size * 3, dtype=float) + ys = np.empty(shown.size * 3, dtype=float) + xs[0::3] = shown + xs[1::3] = shown + xs[2::3] = np.nan + ys[0::3] = offset + ys[1::3] = offset + 1.0 + ys[2::3] = np.nan + self.plot_sync_signal.plot(xs, ys, pen=pg.mkPen(color, width=0.8, style=QtCore.Qt.PenStyle.DashLine)) + + def _stats_text(role: str, label: str, sample_text: str, events: np.ndarray, values: Optional[np.ndarray]) -> None: + intervals = np.diff(events) if events.size > 1 else np.array([], float) + mean_interval = float(np.nanmean(intervals)) if intervals.size else np.nan + cv = float(np.nanstd(intervals) / mean_interval) if intervals.size and mean_interval > 0 else np.nan + cv_text = f"{cv:.3g}" if np.isfinite(cv) else "-" + sat_text = "-" + if values is not None: + finite = np.asarray(values[np.isfinite(values)], float) + if finite.size: + lo = float(np.nanmin(finite)) + hi = float(np.nanmax(finite)) + saturated = float(np.mean((finite <= lo) | (finite >= hi))) if hi > lo else 0.0 + sat_text = f"{saturated:.3g}" if np.isfinite(saturated) else "-" + quality_lines.append( + f"{role}: {label}\n" + f"{role} samples: {sample_text} | edges: {events.size} | interval CV: {cv_text} | saturated frac: {sat_text}" + ) + + def _plot_trace( + role: str, + label: str, + t_arr: np.ndarray, + x_arr: np.ndarray, + mode: str, + offset: float, + color: Tuple[int, int, int], + event_color: Tuple[int, int, int, int], + ) -> bool: + t_arr, x_arr = _clean_trace(t_arr, x_arr) + if t_arr.size < 2: + quality_lines.append(f"{role}: {label}\n{role} samples: {t_arr.size} | edges: -") + return False + x_norm, lo, hi = _normalize_trace(x_arr) + events = extract_sync_events(t_arr, x_arr, mode=mode, threshold=threshold, min_interval_s=min_interval) + stride = max(1, int(np.ceil(t_arr.size / 30000))) + self.plot_sync_signal.plot( + t_arr[::stride], + offset + x_norm[::stride], + pen=pg.mkPen(color, width=1.1), + name=role, + ) + _event_guides(events, offset, event_color) + if events.size: + marker_events = events[::max(1, int(np.ceil(events.size / 1200)))] + self.plot_sync_signal.plot( + marker_events, + np.full_like(marker_events, offset + 1.06), + pen=None, + symbol="t", + symbolSize=7, + symbolBrush=pg.mkBrush(event_color), + symbolPen=pg.mkPen(event_color[:3], width=0.6), + ) + if threshold is not None and np.isfinite(threshold) and np.isfinite(hi - lo) and hi > lo: + y_thr = offset + ((float(threshold) - lo) / (hi - lo)) + if offset - 0.05 <= y_thr <= offset + 1.05: + self.plot_sync_signal.plot( + [float(t_arr[0]), float(t_arr[-1])], + [y_thr, y_thr], + pen=pg.mkPen(event_color[:3], width=1.0, style=QtCore.Qt.PenStyle.DotLine), + ) + _stats_text(role, label, str(t_arr.size), events, x_arr) + plotted_rows.append((offset + 0.5, role)) + return True + + data = self.combo_sync_camera_column.currentData() + if isinstance(data, str) and "::" in data: + try: + kind, name = data.split("::", 1) + ref_label = self.combo_sync_camera_column.currentText() + ref_t = np.array([], float) + ref_x: Optional[np.ndarray] = None + timestamp_events: Optional[np.ndarray] = None + if kind == "external": + source_key, column_name = name.split("::", 1) + ext = self._sync_external_sources.get(source_key) or {} + columns = ext.get("columns") or {} + ref_t = np.asarray(ext.get("time", np.array([], float)), float) + ref_x = np.asarray(columns.get(column_name, np.array([], float)), float) + else: + proc = self._selected_proc_for_sync() + info = self._sync_behavior_source_for_proc(proc) if proc is not None else None + if info is None and self._behavior_sources: + info = next(iter(self._behavior_sources.values())) + if not info: + raise ValueError("Load a behavior, Sync file, or LED video for the reference signal.") + if kind == "behavior": + behaviors = info.get("behaviors") or {} + if name not in behaviors: + raise ValueError(f"Reference sync column not found: {name}") + if str(info.get("kind", _BEHAVIOR_PARSE_BINARY)) == _BEHAVIOR_PARSE_TIMESTAMPS: + timestamp_events = np.asarray(behaviors[name], float) + else: + ref_t = np.asarray(info.get("time", np.array([], float)), float) + ref_x = np.asarray(behaviors[name], float) + elif kind == "trajectory": + traj = info.get("trajectory") or {} + if name not in traj: + raise ValueError(f"Reference sync column not found: {name}") + ref_t = np.asarray(info.get("trajectory_time", np.array([], float)), float) + if ref_t.size == 0: + ref_t = np.asarray(info.get("time", np.array([], float)), float) + ref_x = np.asarray(traj[name], float) + if timestamp_events is not None: + events = extract_sync_events(np.asarray(timestamp_events, float), None, min_interval_s=min_interval) + if events.size: + _event_guides(events, 0.0, (255, 138, 128, 135)) + self.plot_sync_signal.plot( + events, + np.full_like(events, 0.5), + pen=None, + symbol="t", + symbolSize=8, + symbolBrush=pg.mkBrush(127, 209, 230, 220), + ) + _stats_text("Reference", ref_label, "timestamp list", events, None) + plotted_rows.append((0.5, "Reference")) + elif ref_x is not None: + ref_mode = self._sync_mode_key(self.combo_sync_camera_mode.currentText()) + _plot_trace("Reference", ref_label, ref_t, ref_x, ref_mode, 0.0, (127, 209, 230), (255, 138, 128, 135)) + except Exception as exc: + quality_lines.append(f"Reference: {exc}") + else: + quality_lines.append("Reference: choose a camera, behavior, Sync file, or LED signal.") + + try: + proc = self._selected_proc_for_sync() + target_t, target_x = self._sync_fiber_trace_for_proc(proc) + target_mode = self._sync_mode_key(self.combo_sync_fiber_mode.currentText()) + target_label = self.combo_sync_fiber_source.currentText() if hasattr(self, "combo_sync_fiber_source") else "Photometry" + _plot_trace( + "Photometry", + target_label, + target_t, + target_x, + target_mode, + 1.35, + (255, 214, 102), + (91, 255, 197, 150), + ) + except Exception as exc: + quality_lines.append(f"Photometry: {exc}") + + if plotted_rows: + self.plot_sync_signal.getAxis("left").setTicks([[(float(pos), str(label)) for pos, label in plotted_rows]]) + self.plot_sync_signal.setYRange(-0.12, max(pos for pos, _label in plotted_rows) + 0.72, padding=0.02) + else: + self.plot_sync_signal.getAxis("left").setTicks([[]]) + self.lbl_sync_quality.setText("\n".join(quality_lines) if quality_lines else "No Sync signal available to preview.") + except Exception as exc: + self.lbl_sync_quality.setText(f"Signal preview failed: {exc}") + + def _extract_led_sync_source(self, config: Dict[str, object]) -> str: + try: + import cv2 + except Exception as exc: + raise RuntimeError(f"OpenCV is required for LED extraction: {exc}") from exc + path = str(config.get("video_path") or "") + if not path or not os.path.isfile(path): + raise ValueError("Choose a video file.") + fps = float(config.get("fps") or 30.0) + if not np.isfinite(fps) or fps <= 0: + fps = 30.0 + n_frames = max(0, int(config.get("n_frames") or 0)) + start = max(0, int(config.get("start_frame") or 0)) + end = int(config.get("end_frame") or n_frames or 0) + if n_frames: + end = min(max(start + 1, end), n_frames) + if end <= start: + raise ValueError("End frame must be greater than start frame.") + roi = tuple(int(v) for v in (config.get("roi") or (0, 0, 1, 1))) + x, y, w, h = roi + if w <= 0 or h <= 0: + raise ValueError("Choose a non-empty LED ROI.") + channel = str(config.get("channel") or "Grayscale").lower() + cap = cv2.VideoCapture(path) + if not cap.isOpened(): + raise ValueError("Could not open video.") + cap.set(cv2.CAP_PROP_POS_FRAMES, start) + total = end - start + values = np.empty(total, dtype=float) + progress = QtWidgets.QProgressDialog("Extracting LED Sync signal...", "Cancel", 0, total, self) + progress.setWindowModality(QtCore.Qt.WindowModality.WindowModal) + progress.setMinimumDuration(300) + count = 0 + try: + for idx in range(total): + if progress.wasCanceled(): + break + ok, frame = cap.read() + if not ok or frame is None: + break + crop = frame[max(0, y): y + h, max(0, x): x + w] + if crop.size == 0: + values[count] = np.nan + elif channel.startswith("red"): + values[count] = float(np.nanmean(crop[:, :, 2])) + elif channel.startswith("green"): + values[count] = float(np.nanmean(crop[:, :, 1])) + elif channel.startswith("blue"): + values[count] = float(np.nanmean(crop[:, :, 0])) + else: + means = crop.mean(axis=(0, 1)) + values[count] = float(0.114 * means[0] + 0.587 * means[1] + 0.299 * means[2]) + count += 1 + if idx % 30 == 0: + progress.setValue(idx) + QtWidgets.QApplication.processEvents() + finally: + cap.release() + progress.setValue(total) + if count < 2: + raise ValueError("No usable LED samples were extracted.") + values = values[:count] + time = (np.arange(count, dtype=float) + start) / fps + key = f"{os.path.splitext(os.path.basename(path))[0]}_LED" + return self._add_sync_external_source( + key, + time, + {"LED signal": values}, + source_path=path, + kind="led_video", + time_col="video_time_s", + ) + def _sync_temporal_modeling_context(self) -> None: if not hasattr(self, "section_temporal"): return @@ -5031,7 +6302,8 @@ def _update_data_availability(self) -> None: and np.asarray(getattr(proc, "sync_aligned_time"), float).size == np.asarray(getattr(proc, "time", []), float).size for proc in (self._processed or []) ) - sync_ready = has_processed and has_behavior + has_sync_source = has_behavior or bool(getattr(self, "_sync_external_sources", {}) or {}) + sync_ready = has_processed and has_sync_source for w in ( self.combo_sync_behavior_file, self.combo_sync_camera_column, @@ -5049,6 +6321,17 @@ def _update_data_availability(self) -> None: self.btn_sync_apply_batch, ): w.setEnabled(sync_ready) + self.btn_sync_load_file.setEnabled(True) + self.btn_sync_load_file_side.setEnabled(True) + self.btn_sync_load_processed.setEnabled(True) + self.btn_sync_load_processed_side.setEnabled(True) + self.btn_sync_open_video.setEnabled(True) + self.btn_sync_open_video_inline.setEnabled(True) + self.btn_sync_open_video_side.setEnabled(True) + self.btn_sync_refresh_side.setEnabled(True) + self.btn_sync_remove_source.setEnabled(bool(self._behavior_sources or self._sync_external_sources)) + self.btn_sync_extract_led.setEnabled(True) + self.btn_open_sync_dialog.setEnabled(True) self.cb_sync_use_aligned.setEnabled(has_aligned_time) self.combo_sync_export_format.setEnabled(has_aligned_time) self.btn_sync_export.setEnabled(has_aligned_time) @@ -5370,6 +6653,19 @@ def _project_dirty_fingerprint(self) -> str: } ) + sync_sources: List[Dict[str, object]] = [] + for name, info in sorted((getattr(self, "_sync_external_sources", {}) or {}).items()): + data = info if isinstance(info, dict) else {} + sync_sources.append( + { + "name": str(name), + "kind": str(data.get("kind", "") or ""), + "source_path": str(data.get("source_path", "") or ""), + "time_col": str(data.get("time_col", "") or ""), + "columns": sorted(str(k) for k in (data.get("columns", {}) or {}).keys()), + } + ) + signal = getattr(self, "last_signal_events", None) signal_summary: Dict[str, object] = {} if isinstance(signal, dict) and signal: @@ -5418,6 +6714,7 @@ def _project_dirty_fingerprint(self) -> str: "tab_sources_index": int(self.tab_sources.currentIndex()) if hasattr(self, "tab_sources") else 0, "processed": processed, "behavior": behavior, + "sync_sources": sync_sources, "signal_events": signal_summary, "behavior_analysis": behavior_summary, "temporal": temporal_summary, @@ -7021,7 +8318,9 @@ def _sync_mode_key(self, combo_text: str) -> str: return "ttl_rising" def _on_sync_auto_threshold_toggled(self, checked: bool) -> None: - ready = bool(getattr(self, "_processed", None)) and bool(getattr(self, "_behavior_sources", None)) + ready = bool(getattr(self, "_processed", None)) and ( + bool(getattr(self, "_behavior_sources", None)) or bool(getattr(self, "_sync_external_sources", None)) + ) self.spin_sync_threshold.setEnabled(ready and not bool(checked)) def _on_sync_use_aligned_changed(self) -> None: @@ -7041,9 +8340,14 @@ def _refresh_sync_sources(self) -> None: self.combo_sync_behavior_file.blockSignals(True) try: self.combo_sync_behavior_file.clear() - self.combo_sync_behavior_file.addItem("Auto-match behavior file", "") + self.combo_sync_behavior_file.addItem("Auto-match Sync source", "") for stem in (self._behavior_sources or {}).keys(): self.combo_sync_behavior_file.addItem(str(stem), str(stem)) + for key, info in (self._sync_external_sources or {}).items(): + label = f"Sync file: {key}" + if str((info or {}).get("kind", "")) == "led_video": + label = f"LED video: {key}" + self.combo_sync_behavior_file.addItem(label, f"external::{key}") if isinstance(prev_beh, str): idx = self.combo_sync_behavior_file.findData(prev_beh) if idx >= 0: @@ -7053,6 +8357,8 @@ def _refresh_sync_sources(self) -> None: self._refresh_sync_camera_columns() self._refresh_sync_fiber_sources() self._refresh_sync_results_table() + self._refresh_sync_source_list() + self._refresh_sync_signal_preview() def _refresh_sync_camera_columns(self) -> None: if not hasattr(self, "combo_sync_camera_column"): @@ -7063,10 +8369,16 @@ def _refresh_sync_camera_columns(self) -> None: self.combo_sync_camera_column.clear() selected_stem = self.combo_sync_behavior_file.currentData() if hasattr(self, "combo_sync_behavior_file") else "" sources: Dict[str, Dict[str, Any]] = {} + external_sources: Dict[str, Dict[str, Any]] = {} if isinstance(selected_stem, str) and selected_stem and selected_stem in self._behavior_sources: sources[selected_stem] = self._behavior_sources[selected_stem] + elif isinstance(selected_stem, str) and selected_stem.startswith("external::"): + ext_key = selected_stem.split("::", 1)[1] + if ext_key in self._sync_external_sources: + external_sources[ext_key] = self._sync_external_sources[ext_key] else: sources = dict(self._behavior_sources or {}) + external_sources = dict(self._sync_external_sources or {}) names: List[Tuple[str, str, str]] = [] seen = set() for info in sources.values(): @@ -7080,8 +8392,14 @@ def _refresh_sync_camera_columns(self) -> None: if key not in seen: seen.add(key) names.append((f"Column: {name}", "trajectory", str(name))) + for source_key, info in external_sources.items(): + for name in (info.get("columns") or {}).keys(): + key = ("external", str(source_key), str(name)) + if key not in seen: + seen.add(key) + names.append((f"{source_key}: {name}", "external", f"{source_key}::{name}")) if not names: - self.combo_sync_camera_column.addItem("Load behavior or camera CSV/XLSX first", "") + self.combo_sync_camera_column.addItem("Load behavior, Sync file, or LED video first", "") else: for label, kind, name in names: self.combo_sync_camera_column.addItem(label, f"{kind}::{name}") @@ -7108,15 +8426,19 @@ def _refresh_sync_fiber_sources(self) -> None: entries: List[Tuple[str, str]] = [] seen_data: set[str] = set() + def _add_data_entry(label: str, data: str) -> None: + data = str(data or "").strip() + if not data or data in seen_data: + return + seen_data.add(data) + entries.append((str(label), data)) + def _add_entry(label: str, name: str) -> None: name = str(name or "").strip() if not name: return data = f"dio::{name}" - if data in seen_data: - return - seen_data.add(data) - entries.append((label, data)) + _add_data_entry(label, data) names = list(getattr(self, "_known_dio_channels", []) or []) for name in names: @@ -7127,6 +8449,14 @@ def _add_entry(label: str, name: str) -> None: _add_entry(f"Embedded DIO: {name}", name) for trig_name in (getattr(proc, "triggers", {}) or {}).keys(): _add_entry(f"Processed column: {trig_name}", str(trig_name)) + for source_key, info in (self._sync_external_sources or {}).items(): + kind = str((info or {}).get("kind", "") or "") + prefix = "LED video" if kind == "led_video" else "Sync file" + for col_name in (info.get("columns") or {}).keys(): + _add_data_entry( + f"{prefix}: {source_key} / {col_name}", + f"external::{source_key}::{col_name}", + ) for label, data in entries: self.combo_sync_fiber_source.addItem(label, data) if isinstance(prev, str): @@ -7154,8 +8484,6 @@ def _sync_behavior_source_for_proc(self, proc: ProcessedTrial) -> Optional[Dict[ def _sync_camera_events_for_proc(self, proc: ProcessedTrial) -> np.ndarray: info = self._sync_behavior_source_for_proc(proc) - if not info: - raise ValueError("No behavior/camera file matched this recording.") data = self.combo_sync_camera_column.currentData() if hasattr(self, "combo_sync_camera_column") else "" if not isinstance(data, str) or "::" not in data: raise ValueError("Choose a camera sync behavior or column.") @@ -7163,6 +8491,21 @@ def _sync_camera_events_for_proc(self, proc: ProcessedTrial) -> np.ndarray: min_interval = float(self.spin_sync_min_interval.value()) mode = self._sync_mode_key(self.combo_sync_camera_mode.currentText()) threshold = None if self.cb_sync_auto_threshold.isChecked() else float(self.spin_sync_threshold.value()) + if kind == "external": + if "::" not in name: + raise ValueError("Choose a loaded Sync signal column.") + source_key, column_name = name.split("::", 1) + ext = self._sync_external_sources.get(source_key) + if not ext: + raise ValueError(f"Loaded Sync source not found: {source_key}") + columns = ext.get("columns") or {} + if column_name not in columns: + raise ValueError(f"Sync signal column not found: {column_name}") + t = np.asarray(ext.get("time", np.array([], float)), float) + x = np.asarray(columns[column_name], float) + return extract_sync_events(t, x, mode=mode, threshold=threshold, min_interval_s=min_interval) + if not info: + raise ValueError("No behavior/camera file matched this recording.") if kind == "behavior": behaviors = info.get("behaviors") or {} if name not in behaviors: @@ -7187,6 +8530,20 @@ def _sync_fiber_trace_for_proc(self, proc: ProcessedTrial) -> Tuple[np.ndarray, source = self.combo_sync_fiber_source.currentData() if hasattr(self, "combo_sync_fiber_source") else "__embedded__" if not isinstance(source, str): source = "__embedded__" + if source.startswith("external::"): + parts = source.split("::", 2) + if len(parts) != 3: + raise ValueError("Choose a loaded target Sync column.") + _kind, source_key, column_name = parts + ext = self._sync_external_sources.get(source_key) + if not ext: + raise ValueError(f"Loaded target Sync source not found: {source_key}") + columns = ext.get("columns") or {} + if column_name not in columns: + raise ValueError(f"Target Sync column not found: {column_name}") + t = np.asarray(ext.get("time", np.array([], float)), float) + x = np.asarray(columns[column_name], float) + return t, x if source == "__embedded__": t = np.asarray(getattr(proc, "time", np.array([], float)), float) x = getattr(proc, "dio", None) @@ -7235,7 +8592,7 @@ def _compute_time_sync_for_proc(self, proc: ProcessedTrial) -> SyncResult: ) method = "interpolation" if "interp" in self.combo_sync_method.currentText().lower() else "linear" return align_timebase( - np.asarray(getattr(proc, "time", np.array([], float)), float), + np.asarray(fiber_time, float), camera_events, fiber_events, method=method, @@ -7322,6 +8679,7 @@ def _store_sync_result(self, proc: ProcessedTrial, result: SyncResult) -> None: report = result.summary_dict() report["camera_source"] = self.combo_sync_camera_column.currentText() report["fiber_source"] = self.combo_sync_fiber_source.currentText() + report["sync_tool"] = str(_SYNC_TOOL_ROOT) report["created_utc"] = datetime.now(timezone.utc).isoformat() proc.sync_report = report self._sync_results_by_file[self._file_id_for_proc(proc)] = report @@ -7370,14 +8728,14 @@ def _apply_time_sync_batch(self) -> None: if not self._processed: self.sync_status.set("Load processed recordings first.", "warn") return - progress = QtWidgets.QProgressDialog("Synchronizing files...", "Cancel", 0, len(self._processed), self) + progress = QtWidgets.QProgressDialog("Running Sync...", "Cancel", 0, len(self._processed), self) progress.setWindowModality(QtCore.Qt.WindowModality.WindowModal) progress.setMinimumDuration(300) ok_count = 0 errors: List[str] = [] for idx, proc in enumerate(self._processed, start=1): progress.setValue(idx - 1) - progress.setLabelText(f"Synchronizing {self._file_id_for_proc(proc)}...") + progress.setLabelText(f"Sync: {self._file_id_for_proc(proc)}") QtWidgets.QApplication.processEvents() if progress.wasCanceled(): break @@ -7403,10 +8761,10 @@ def _apply_time_sync_batch(self) -> None: self._compute_spatial_heatmap() if errors: self.sync_status.set(f"Aligned {ok_count} file(s); {len(errors)} file(s) need attention.", "warn") - self.txt_sync_report.setPlainText("Batch synchronization warnings:\n" + "\n".join(errors[:30])) + self.txt_sync_report.setPlainText("Batch Sync warnings:\n" + "\n".join(errors[:30])) else: self.sync_status.set(f"Aligned {ok_count} file(s).", "ok") - self.statusUpdate.emit(f"Time synchronization batch finished: {ok_count} file(s).", 5000) + self.statusUpdate.emit(f"Sync batch finished: {ok_count} file(s).", 5000) def _refresh_sync_results_table(self) -> None: if not hasattr(self, "tbl_sync_results"): @@ -7449,7 +8807,7 @@ def _export_sync_aligned_files(self) -> None: and np.asarray(getattr(proc, "sync_aligned_time"), float).size == np.asarray(getattr(proc, "time", []), float).size ] if not ready: - self.sync_status.set("No aligned time columns to export. Apply synchronization first.", "warn") + self.sync_status.set("No aligned time columns to export. Apply Sync first.", "warn") return out_dir = QtWidgets.QFileDialog.getExistingDirectory(self, "Select aligned export folder", self._export_start_dir()) if not out_dir: @@ -10462,6 +11820,25 @@ def _save_project_h5(self, path: str) -> None: ds = trajectory_group.create_dataset(f"item_{t_idx:04d}", data=data, **kwargs) ds.attrs["name"] = str(name) + sync_sources_group = f.create_group("sync_sources") + sync_sources_group.attrs["count"] = int(len(self._sync_external_sources)) + for idx, (key, info) in enumerate((self._sync_external_sources or {}).items()): + source = info or {} + entry = sync_sources_group.create_group(f"item_{idx:04d}") + entry.attrs["key"] = str(key) + entry.attrs["kind"] = str(source.get("kind", "signal_file") or "signal_file") + entry.attrs["source_path"] = str(source.get("source_path", "") or "") + entry.attrs["time_col"] = str(source.get("time_col", "") or "") + self._write_h5_numeric(entry, "time", np.asarray(source.get("time", np.array([], float)), float)) + columns_group = entry.create_group("columns") + for col_idx, (name, values) in enumerate((source.get("columns") or {}).items()): + data = np.asarray(values, float) + kwargs: Dict[str, object] = {} + if data.size > 0: + kwargs["compression"] = "gzip" + ds = columns_group.create_dataset(f"item_{col_idx:04d}", data=data, **kwargs) + ds.attrs["name"] = str(name) + analysis_group = f.create_group("analysis") self._save_signal_events_h5(analysis_group) self._save_behavior_analysis_h5(analysis_group) @@ -10479,6 +11856,7 @@ def _load_project_h5(self, path: str) -> Dict[str, object]: "tab_sources_index": 0, "processed": [], "behavior_sources": {}, + "sync_sources": {}, "recent_paths": {}, "signal_events": None, "behavior_analysis": None, @@ -10665,6 +12043,34 @@ def _aligned(values: Optional[np.ndarray], fill_nan: bool = True) -> np.ndarray: loaded_behavior[stem] = info + loaded_sync_sources: Dict[str, Dict[str, Any]] = {} + sync_sources_group = f.get("sync_sources") + if isinstance(sync_sources_group, h5py.Group): + for key in sorted(sync_sources_group.keys()): + entry = sync_sources_group.get(key) + if not isinstance(entry, h5py.Group): + continue + source_key = self._h5_text(entry.attrs.get("key", key), key) + info: Dict[str, Any] = { + "kind": self._h5_text(entry.attrs.get("kind", "signal_file"), "signal_file"), + "source_path": self._h5_text(entry.attrs.get("source_path", ""), ""), + "time_col": self._h5_text(entry.attrs.get("time_col", ""), ""), + "time": np.asarray( + self._read_h5_numeric(entry, "time") if "time" in entry else np.array([], float), + float, + ), + "columns": {}, + } + columns_group = entry.get("columns") + if isinstance(columns_group, h5py.Group): + for c_key in sorted(columns_group.keys()): + ds = columns_group.get(c_key) + if ds is None: + continue + name = self._h5_text(ds.attrs.get("name", c_key), c_key) + info["columns"][name] = np.asarray(ds[()], float) + loaded_sync_sources[source_key] = info + analysis_group = f.get("analysis") if isinstance(analysis_group, h5py.Group): payload["signal_events"] = self._load_signal_events_h5(analysis_group) @@ -10676,6 +12082,7 @@ def _aligned(values: Optional[np.ndarray], fill_nan: bool = True) -> np.ndarray: payload["processed"] = loaded_processed payload["behavior_sources"] = loaded_behavior + payload["sync_sources"] = loaded_sync_sources return payload def _autosave_project_cache_path(self) -> str: @@ -10689,7 +12096,7 @@ def _autosave_project_cache_path(self) -> str: return os.path.join(cache_dir, "autosave_project.h5") def _has_project_state_for_autosave(self) -> bool: - if self._processed or self._behavior_sources: + if self._processed or self._behavior_sources or self._sync_external_sources: return True if isinstance(self.last_signal_events, dict) and bool(self.last_signal_events): return True @@ -10822,6 +12229,7 @@ def _reset_project_state(self) -> None: self._clear_cached_analysis_outputs() self._processed = [] self._behavior_sources = {} + self._sync_external_sources = {} self._continuous_align_rules = {} self._sync_results_by_file = {} self._last_sync_preview = None @@ -10941,6 +12349,7 @@ def _load_project_from_path(self, path: str, from_autosave: bool = False) -> boo settings_data = payload.get("settings", {}) processed = payload.get("processed", []) behavior_sources = payload.get("behavior_sources", {}) + sync_sources = payload.get("sync_sources", {}) tab_idx = payload.get("tab_sources_index", 0) recent_paths = payload.get("recent_paths", {}) if isinstance(payload.get("recent_paths", {}), dict) else {} @@ -10950,6 +12359,7 @@ def _load_project_from_path(self, path: str, from_autosave: bool = False) -> boo self._clear_cached_analysis_outputs() self._processed = list(processed) if isinstance(processed, list) else [] self._behavior_sources = dict(behavior_sources) if isinstance(behavior_sources, dict) else {} + self._sync_external_sources = dict(sync_sources) if isinstance(sync_sources, dict) else {} self.lbl_group.setText(f"{len(self._processed)} file(s) loaded") kinds = { diff --git a/pyBer/styles.py b/pyBer/styles.py index 3ad38b1..c284606 100644 --- a/pyBer/styles.py +++ b/pyBer/styles.py @@ -948,6 +948,39 @@ def _paint_temporal(p, r, c): border-radius: 12px; } +QFrame#SidePanel, QFrame#syncPanel { + background: #1f242e; + border: 1px solid #343c4d; + border-radius: 8px; +} + +QFrame#SidePanel QListWidget { + background: #202631; + border: 1px solid #3a4050; + border-radius: 8px; + padding: 4px; +} + +QFrame#SidePanel QListWidget::item { + padding: 8px; + border-radius: 6px; + margin: 2px 0; +} + +QFrame#SidePanel QListWidget::item:selected { + background: #4a3678; + color: #ffffff; +} + +QLabel#BadgeLabel { + background: #243342; + border: 1px solid #3d5669; + border-radius: 10px; + color: #9ee6f4; + padding: 4px 8px; + font-weight: 700; +} + QToolButton#toolbarIconButton { background: #1f242e; border: 1px solid #343c4d; @@ -1421,6 +1454,18 @@ def _build_light_qss(dark_qss: str) -> str: /* Drawer panels */ QFrame#drawerPanel, QFrame#centerPanel { background: #ffffff; border: 1px solid #d6dde9; } QFrame#transportBar { background: #eef1f7; border: 1px solid #d6dde9; } +QFrame#SidePanel, QFrame#syncPanel { background: #ffffff; border: 1px solid #d6dde9; border-radius: 8px; } +QFrame#SidePanel QListWidget { background: #f7f9fc; border: 1px solid #d6dde9; border-radius: 8px; padding: 4px; } +QFrame#SidePanel QListWidget::item { padding: 8px; border-radius: 6px; margin: 2px 0; } +QFrame#SidePanel QListWidget::item:selected { background: #2563eb; color: #ffffff; } +QLabel#BadgeLabel { + background: #dbeafe; + border: 1px solid #93c5fd; + border-radius: 10px; + color: #172033; + padding: 4px 8px; + font-weight: 700; +} QToolButton#toolbarIconButton { background: #ffffff; border: 1px solid #c2ccda; diff --git a/pyBer/temporal_modeling.py b/pyBer/temporal_modeling.py index e84c44c..627eec1 100644 --- a/pyBer/temporal_modeling.py +++ b/pyBer/temporal_modeling.py @@ -1216,6 +1216,23 @@ def fit( background: #1f6db1; border: 1px solid #35a4e8; } +QToolButton#paramHelpButton { + color: #9bd8ff; + background: #101b2b; + border: 1px solid #38597a; + border-radius: 9px; + padding: 0; + min-width: 18px; + max-width: 18px; + min-height: 18px; + max-height: 18px; + font-weight: 800; +} +QToolButton#paramHelpButton:hover { + color: #ffffff; + background: #1f6db1; + border: 1px solid #35a4e8; +} QTabWidget::pane { border: 1px solid #263a52; border-radius: 6px; @@ -1342,6 +1359,23 @@ def fit( background: #2563eb; border: 1px solid #1d4ed8; } +QToolButton#paramHelpButton { + color: #1d4ed8; + background: #eef6ff; + border: 1px solid #93c5fd; + border-radius: 9px; + padding: 0; + min-width: 18px; + max-width: 18px; + min-height: 18px; + max-height: 18px; + font-weight: 800; +} +QToolButton#paramHelpButton:hover { + color: #ffffff; + background: #2563eb; + border: 1px solid #1d4ed8; +} QTabWidget::pane { border: 1px solid #d6dde9; background: #ffffff; @@ -1786,7 +1820,8 @@ def _build_compact_ui(self): self.combo_fit_scope.setToolTip( "Active: fit the selected animal only.\n" "All: concatenate every loaded recording into one GLM.\n" - "Per-file batch: fit each animal independently, then aggregate for the Group tab." + "Per-file batch: snapshot the current predictor list, fit each animal independently with that same set, " + "then aggregate for the Group tab." ) sb.addWidget(self.combo_fit_scope) @@ -1904,50 +1939,89 @@ def _build_model_page(self): self.combo_basis = QtWidgets.QComboBox() self.combo_basis.addItems(["Raised cosine", "B-spline", "FIR"]) - gl.addRow("Basis", self.combo_basis) + gl.addRow(self._help_label( + "Basis", + "Controls the temporal shape used to model each predictor kernel. " + "Raised cosine is smooth and stable, B-spline is flexible, FIR is least constrained and needs more data.", + self.combo_basis, + ), self.combo_basis) self.spin_n_basis = QtWidgets.QSpinBox() self.spin_n_basis.setRange(2, 50) self.spin_n_basis.setValue(8) - gl.addRow("Basis count", self.spin_n_basis) + gl.addRow(self._help_label( + "Basis count", + "Number of basis functions per predictor. Higher values capture faster or more complex kernels, " + "but increase parameters and overfitting risk.", + self.spin_n_basis, + ), self.spin_n_basis) self.combo_reg = QtWidgets.QComboBox() self.combo_reg.addItems(["Ridge", "Lasso", "OLS"]) - gl.addRow("Regularization", self.combo_reg) + gl.addRow(self._help_label( + "Regularization", + "Penalty used during fitting. Ridge shrinks correlated predictors smoothly, Lasso can zero weak predictors, " + "OLS is unpenalized and can be unstable with many or correlated predictors.", + self.combo_reg, + ), self.combo_reg) self.spin_alpha = QtWidgets.QDoubleSpinBox() self.spin_alpha.setRange(0.001, 1000.0) self.spin_alpha.setValue(1.0) self.spin_alpha.setDecimals(3) - gl.addRow("Alpha", self.spin_alpha) + gl.addRow(self._help_label( + "Alpha", + "Regularization strength for Ridge or Lasso. Larger alpha gives smoother, smaller kernels and less variance, " + "but can hide real effects. OLS ignores this value.", + self.spin_alpha, + ), self.spin_alpha) self.spin_kernel_pre = QtWidgets.QDoubleSpinBox() self.spin_kernel_pre.setRange(-30.0, 0.0) self.spin_kernel_pre.setValue(-1.0) self.spin_kernel_pre.setDecimals(1) self.spin_kernel_pre.setSuffix(" s") - gl.addRow("Kernel pre", self.spin_kernel_pre) + gl.addRow(self._help_label( + "Kernel pre", + "Seconds before each event included in the kernel. Use this to detect pre-event ramps or anticipatory effects. " + "Longer windows cost more parameters.", + self.spin_kernel_pre, + ), self.spin_kernel_pre) self.spin_kernel_post = QtWidgets.QDoubleSpinBox() self.spin_kernel_post.setRange(0.1, 60.0) self.spin_kernel_post.setValue(3.0) self.spin_kernel_post.setDecimals(1) self.spin_kernel_post.setSuffix(" s") - gl.addRow("Kernel post", self.spin_kernel_post) + gl.addRow(self._help_label( + "Kernel post", + "Seconds after each event included in the kernel. Increase it for slow photometry responses, " + "but avoid overly long windows that mix unrelated events.", + self.spin_kernel_post, + ), self.spin_kernel_post) self.spin_glm_bootstrap = QtWidgets.QSpinBox() self.spin_glm_bootstrap.setRange(0, 2000) self.spin_glm_bootstrap.setValue(100) self.spin_glm_bootstrap.setSpecialValueText("off") self.spin_glm_bootstrap.setToolTip("Circular-shift bootstraps for leave-one-out contribution p-values.") - gl.addRow("Shift bootstraps", self.spin_glm_bootstrap) + gl.addRow(self._help_label( + "Shift bootstraps", + "Number of circular-shift null fits for predictor contribution p-values and kernel confidence intervals. " + "More iterations improve p-value resolution but take longer.", + self.spin_glm_bootstrap, + ), self.spin_glm_bootstrap) self.spin_glm_jobs = QtWidgets.QSpinBox() max_jobs = max(1, os.cpu_count() or 1) self.spin_glm_jobs.setRange(1, max_jobs) self.spin_glm_jobs.setValue(min(4, max_jobs)) self.spin_glm_jobs.setToolTip("Parallel jobs used for circular-shift bootstrap fits.") - gl.addRow("Bootstrap jobs", self.spin_glm_jobs) + gl.addRow(self._help_label( + "Bootstrap jobs", + "Parallel worker count for bootstrap fits. More jobs can speed up analysis, but also use more CPU and memory.", + self.spin_glm_jobs, + ), self.spin_glm_jobs) self.spin_glm_cv_folds = QtWidgets.QSpinBox() self.spin_glm_cv_folds.setRange(0, 20) @@ -1958,7 +2032,12 @@ def _build_model_page(self): "Set 0/auto to let pyBer pick: one fold per file when multiple files " "are loaded, otherwise 5 contiguous time blocks for a single file." ) - gl.addRow("CV folds", self.spin_glm_cv_folds) + gl.addRow(self._help_label( + "CV folds", + "Number of held-out folds for out-of-sample R^2. More folds use data efficiently but take longer. " + "Auto uses files as folds when possible.", + self.spin_glm_cv_folds, + ), self.spin_glm_cv_folds) lay.addWidget(self.grp_glm) self.grp_flmm = QtWidgets.QGroupBox("FLMM Settings") @@ -1974,22 +2053,47 @@ def _build_model_page(self): self.edit_formula = QtWidgets.QLineEdit("Y.obs ~ 1") self.edit_formula.setPlaceholderText("Leave as Y.obs ~ 1 to auto-use selected predictors") - fl.addRow("Fixed formula", self.edit_formula) + fl.addRow(self._help_label( + "Fixed formula", + "R fixed-effect formula. Leave as Y.obs ~ 1 to automatically include selected predictors. " + "Custom formulas let you test a specific hypothesis, but omitted predictors are not estimated.", + self.edit_formula, + ), self.edit_formula) self.edit_random = QtWidgets.QLineEdit("~1") self.edit_random.setPlaceholderText("e.g. ~1 or ~time") - fl.addRow("Random", self.edit_random) + fl.addRow(self._help_label( + "Random", + "Random-effect structure for repeated trials or subjects. ~1 estimates baseline differences. " + "More complex terms can model subject-specific trends but need more repeated observations.", + self.edit_random, + ), self.edit_random) self.edit_group_var = QtWidgets.QLineEdit("subject") - fl.addRow("Group var", self.edit_group_var) + fl.addRow(self._help_label( + "Group var", + "Column used as the random-effect grouping variable, usually subject or animal. " + "Wrong grouping can make uncertainty estimates misleading.", + self.edit_group_var, + ), self.edit_group_var) self.spin_nknots = QtWidgets.QSpinBox() self.spin_nknots.setRange(0, 100) self.spin_nknots.setValue(0) self.spin_nknots.setSpecialValueText("auto") - fl.addRow("Min knots", self.spin_nknots) + fl.addRow(self._help_label( + "Min knots", + "Minimum spline knots for fastFMM. More knots allow sharper time-varying coefficient curves, " + "but can overfit noisy peri-event responses.", + self.spin_nknots, + ), self.spin_nknots) self.spin_boots = QtWidgets.QSpinBox() self.spin_boots.setRange(0, 5000) self.spin_boots.setValue(0) self.spin_boots.setSpecialValueText("analytic") - fl.addRow("Bootstrap iter", self.spin_boots) + fl.addRow(self._help_label( + "Bootstrap iter", + "Bootstrap iterations for FLMM uncertainty when supported. Analytic is faster. " + "Bootstrap can be more robust but increases runtime heavily.", + self.spin_boots, + ), self.spin_boots) self.combo_flmm_importance = QtWidgets.QComboBox() self.combo_flmm_importance.addItem("Fast coefficient ranking", "fast") self.combo_flmm_importance.addItem("Permutation contribution test", "perm") @@ -2001,7 +2105,12 @@ def _build_model_page(self): "refit, derive a p-value for variable contribution (slower but rigorous).\n" "Leave-one-out: refit FLMM without each predictor and compare AICs (very slow)." ) - fl.addRow("Contribution", self.combo_flmm_importance) + fl.addRow(self._help_label( + "Contribution", + "How variable importance is estimated. Fast is descriptive, permutation gives a null test, " + "leave-one-out AIC is rigorous but very slow.", + self.combo_flmm_importance, + ), self.combo_flmm_importance) self.spin_flmm_perm = QtWidgets.QSpinBox() self.spin_flmm_perm.setRange(50, 5000) @@ -2012,7 +2121,12 @@ def _build_model_page(self): "(only used when Contribution = Permutation). The smallest " "achievable p-value is 1/(N+1)." ) - fl.addRow("Permutations", self.spin_flmm_perm) + fl.addRow(self._help_label( + "Permutations", + "Number of shuffled-label FLMM refits for contribution testing. Higher values give finer p-values " + "and more stable results, but scale runtime linearly.", + self.spin_flmm_perm, + ), self.spin_flmm_perm) lay.addWidget(self.grp_flmm) lay.addStretch(1) self.stack_controls.addWidget(page) @@ -2038,8 +2152,13 @@ def _build_predictor_page(self): row = QtWidgets.QHBoxLayout() self.btn_add_predictor = QtWidgets.QPushButton("+ Add") + self.btn_add_all_predictors = QtWidgets.QPushButton("Add all variables") + self.btn_add_all_predictors.setToolTip( + "Add every currently available predictor to the model. Remove unwanted variables from the selected list before fitting." + ) self.btn_remove_predictor = QtWidgets.QPushButton("- Remove") row.addWidget(self.btn_add_predictor) + row.addWidget(self.btn_add_all_predictors) row.addWidget(self.btn_remove_predictor) row.addStretch(1) pl.addLayout(row) @@ -2066,7 +2185,7 @@ def _build_files_page(self): hint = QtWidgets.QLabel( "Select an animal to make it the active recording. The Group tab aggregates " - "per-file fits once you run a Per-file batch." + "per-file fits once you run a Per-file batch. Batch fitting uses the current predictor list for every file." ) hint.setProperty("class", "muted") hint.setWordWrap(True) @@ -2077,8 +2196,12 @@ def _build_files_page(self): self.list_files.setMinimumHeight(220) gl.addWidget(self.list_files, 1) - self.btn_fit_all_files = QtWidgets.QPushButton("Fit each file (per-file batch)") + self.btn_fit_all_files = QtWidgets.QPushButton("Fit all files with current predictors") self.btn_fit_all_files.setProperty("class", "primary") + self.btn_fit_all_files.setToolTip( + "Run one GLM per loaded file using the exact predictor list currently shown in the Predictors panel. " + "The same set is attempted for every file, then cached and aggregated in the Group tab." + ) gl.addWidget(self.btn_fit_all_files) self.lbl_batch_status = QtWidgets.QLabel("") @@ -2429,6 +2552,28 @@ def _make_nav_button(self, text: str) -> QtWidgets.QToolButton: btn.setToolButtonStyle(QtCore.Qt.ToolButtonStyle.ToolButtonTextOnly) return btn + def _make_param_help(self, text: str) -> QtWidgets.QToolButton: + btn = QtWidgets.QToolButton() + btn.setObjectName("paramHelpButton") + btn.setText("?") + btn.setAutoRaise(True) + btn.setCursor(QtCore.Qt.CursorShape.WhatsThisCursor) + btn.setToolTip(str(text or "")) + return btn + + def _help_label(self, label: str, help_text: str, field: Optional[QtWidgets.QWidget] = None) -> QtWidgets.QWidget: + if field is not None and not field.toolTip(): + field.setToolTip(help_text) + wrap = QtWidgets.QWidget() + lay = QtWidgets.QHBoxLayout(wrap) + lay.setContentsMargins(0, 0, 0, 0) + lay.setSpacing(5) + lbl = QtWidgets.QLabel(label) + lay.addWidget(lbl) + lay.addWidget(self._make_param_help(help_text)) + lay.addStretch(1) + return wrap + def _select_control_page(self, index: int) -> None: self.stack_controls.setCurrentIndex(index) buttons = (self.btn_nav_model, self.btn_nav_predictors, self.btn_nav_files, self.btn_nav_fit) @@ -2540,6 +2685,8 @@ def _connect_signals(self): self.combo_model_type.currentIndexChanged.connect(self._on_model_type_changed) self.btn_fit.clicked.connect(self._on_fit_clicked) self.btn_add_predictor.clicked.connect(self._on_add_predictor) + if hasattr(self, "btn_add_all_predictors"): + self.btn_add_all_predictors.clicked.connect(self._on_add_all_predictors) self.btn_remove_predictor.clicked.connect(self._on_remove_predictor) for widget in ( self.combo_basis, @@ -2551,8 +2698,10 @@ def _connect_signals(self): self.spin_kernel_post, self.spin_glm_bootstrap, self.spin_glm_jobs, + self.spin_glm_cv_folds, self.spin_nknots, self.spin_boots, + self.spin_flmm_perm, ): signal = getattr(widget, "currentIndexChanged", None) or getattr(widget, "valueChanged", None) if signal is not None: @@ -2604,11 +2753,15 @@ def _load_temporal_settings(self) -> None: self.spin_kernel_post.setValue(float(self._settings.value(prefix + "kernel_post", self.spin_kernel_post.value()))) self.spin_glm_bootstrap.setValue(int(self._settings.value(prefix + "glm_shift_bootstraps", self.spin_glm_bootstrap.value()))) self.spin_glm_jobs.setValue(int(self._settings.value(prefix + "glm_bootstrap_jobs", self.spin_glm_jobs.value()))) + if hasattr(self, "spin_glm_cv_folds"): + self.spin_glm_cv_folds.setValue(int(self._settings.value(prefix + "glm_cv_folds", self.spin_glm_cv_folds.value()))) self.edit_formula.setText(str(self._settings.value(prefix + "flmm_formula", self.edit_formula.text()) or "Y.obs ~ 1")) self.edit_random.setText(str(self._settings.value(prefix + "flmm_random", self.edit_random.text()) or "~1")) self.edit_group_var.setText(str(self._settings.value(prefix + "flmm_group_var", self.edit_group_var.text()) or "subject")) self.spin_nknots.setValue(int(self._settings.value(prefix + "flmm_nknots", self.spin_nknots.value()))) self.spin_boots.setValue(int(self._settings.value(prefix + "flmm_boots", self.spin_boots.value()))) + if hasattr(self, "spin_flmm_perm"): + self.spin_flmm_perm.setValue(int(self._settings.value(prefix + "flmm_permutations", self.spin_flmm_perm.value()))) mode = str(self._settings.value(prefix + "flmm_importance_mode", "fast") or "fast") idx = self.combo_flmm_importance.findData(mode, QtCore.Qt.ItemDataRole.UserRole) if idx < 0: @@ -2651,12 +2804,16 @@ def _save_temporal_settings(self) -> None: self._settings.setValue(prefix + "kernel_post", self.spin_kernel_post.value()) self._settings.setValue(prefix + "glm_shift_bootstraps", self.spin_glm_bootstrap.value()) self._settings.setValue(prefix + "glm_bootstrap_jobs", self.spin_glm_jobs.value()) + if hasattr(self, "spin_glm_cv_folds"): + self._settings.setValue(prefix + "glm_cv_folds", self.spin_glm_cv_folds.value()) self._settings.setValue(prefix + "flmm_formula", self.edit_formula.text().strip()) self._settings.setValue(prefix + "flmm_random", self.edit_random.text().strip()) self._settings.setValue(prefix + "flmm_group_var", self.edit_group_var.text().strip()) self._settings.setValue(prefix + "flmm_nknots", self.spin_nknots.value()) self._settings.setValue(prefix + "flmm_boots", self.spin_boots.value()) self._settings.setValue(prefix + "flmm_importance_mode", self.combo_flmm_importance.currentData(QtCore.Qt.ItemDataRole.UserRole) or "fast") + if hasattr(self, "spin_flmm_perm"): + self._settings.setValue(prefix + "flmm_permutations", self.spin_flmm_perm.value()) predictors = self._selected_predictor_keys() if hasattr(self, "list_predictors") else self._saved_predictor_keys predictors = list(predictors) self._saved_predictor_keys = list(predictors) @@ -3850,9 +4007,12 @@ def _predictor_vector_for_proc(self, key: str, proc: Any, time: np.ndarray) -> T return np.zeros(time.size, float), str(entry.get("kind", "event")) def _build_glm_dataset_from_selected_predictors( - self, file_filter: Optional[str] = None + self, + file_filter: Optional[str] = None, + predictor_keys: Optional[List[str]] = None, ) -> Dict[str, Any]: - selected = self._selected_predictor_keys() + selected = [str(key).strip() for key in (predictor_keys if predictor_keys is not None else self._selected_predictor_keys())] + selected = [key for key in selected if key in self._predictor_catalog] if not selected: return {"error": "Choose at least one predictor before fitting."} if not self._processed_trials: @@ -6232,6 +6392,20 @@ def _on_add_predictor(self): self._save_temporal_settings() self.statusMessage.emit(f"Added predictor: {self._predictor_label(key)}", 3000) + def _on_add_all_predictors(self): + if not self._predictor_catalog: + self.statusMessage.emit("No predictors are available yet. Load or compute behavior/events first.", 5000) + return + added = 0 + for key in self._predictor_catalog.keys(): + if self._add_predictor_item(key): + added += 1 + self._save_temporal_settings() + if added: + self.statusMessage.emit(f"Added {added} predictor(s). Remove unwanted variables before fitting.", 4000) + else: + self.statusMessage.emit("All available predictors are already selected.", 3000) + def _on_remove_predictor(self): sel = self.list_predictors.currentRow() if sel >= 0: @@ -6943,9 +7117,16 @@ def _on_fit_clicked(self): # GLM fit # ------------------------------------------------------------------ - def _fit_glm_catalog(self, file_filter: Optional[str] = None) -> Optional[GLMResult]: + def _fit_glm_catalog( + self, + file_filter: Optional[str] = None, + predictor_keys: Optional[List[str]] = None, + ) -> Optional[GLMResult]: self._save_temporal_settings() - dataset = self._build_glm_dataset_from_selected_predictors(file_filter=file_filter) + dataset = self._build_glm_dataset_from_selected_predictors( + file_filter=file_filter, + predictor_keys=predictor_keys, + ) if "error" in dataset: msg = str(dataset.get("error", "Could not build GLM dataset.")) dropped = dataset.get("dropped_predictors", []) or [] @@ -7164,10 +7345,20 @@ def _fit_glm_catalog(self, file_filter: Optional[str] = None) -> Optional[GLMRes return result def _fit_glm_per_file_batch(self) -> None: - """Fit each loaded file independently and populate the Group tab.""" + """Fit each loaded file independently with one fixed predictor set.""" if not self._processed_trials: self.statusMessage.emit("No recordings loaded.", 5000) return + batch_predictors = [ + key for key in self._selected_predictor_keys() + if key in self._predictor_catalog + ] + if not batch_predictors: + self.statusMessage.emit("Choose predictors before running the batch.", 5000) + self._select_control_page(1) + return + self._saved_predictor_keys = list(batch_predictors) + self._save_temporal_settings() file_ids = [ self._proc_file_id(p, fallback=f"file_{i + 1}") for i, p in enumerate(self._processed_trials) @@ -7185,7 +7376,7 @@ def _fit_glm_per_file_batch(self) -> None: self.lbl_batch_status.setText(f"Fitting {idx}/{n}: {fid}") QtWidgets.QApplication.processEvents() try: - result = self._fit_glm_catalog(file_filter=fid) + result = self._fit_glm_catalog(file_filter=fid, predictor_keys=batch_predictors) except Exception as exc: _LOG.warning("Per-file fit failed for %s: %s", fid, exc) result = None @@ -7195,7 +7386,7 @@ def _fit_glm_per_file_batch(self) -> None: self._progress_update(idx, f"Per-file batch ({idx}/{n})") if hasattr(self, "lbl_batch_status"): self.lbl_batch_status.setText( - f"Batch complete: {len(ok_results)}/{n} files fit successfully." + f"Batch complete: {len(ok_results)}/{n} files fit successfully with {len(batch_predictors)} shared predictor(s)." ) self._aggregate_group_results() if hasattr(self, "tabs_workspace") and ok_results: From 41b0aa8b6372c3ae12621c8c99c2085cee18d55e Mon Sep 17 00:00:00 2001 From: andrianj Date: Mon, 1 Jun 2026 17:25:04 +0200 Subject: [PATCH 5/7] Fix barcode sync event matching --- pyBer/time_sync.py | 64 +++++++++++++++++++++++++++++++++++++++-- tests/test_time_sync.py | 12 ++++++++ 2 files changed, 74 insertions(+), 2 deletions(-) diff --git a/pyBer/time_sync.py b/pyBer/time_sync.py index 9e3842f..7c25e73 100644 --- a/pyBer/time_sync.py +++ b/pyBer/time_sync.py @@ -164,6 +164,59 @@ def _paired_by_offset( return cam[start_cam:start_cam + n], fib[:n] +def _overlap_candidate_offsets( + camera_events: np.ndarray, + fiber_events: np.ndarray, + *, + max_candidates: int = 200, +) -> List[int]: + """Infer plausible pulse offsets when both event vectors share a time range.""" + cam = np.asarray(camera_events, float).reshape(-1) + fib = np.asarray(fiber_events, float).reshape(-1) + if cam.size < 2 or fib.size < 2: + return [] + + overlap_start = max(float(cam[0]), float(fib[0])) + overlap_end = min(float(cam[-1]), float(fib[-1])) + cam_span = max(0.0, float(cam[-1] - cam[0])) + fib_span = max(0.0, float(fib[-1] - fib[0])) + overlap_span = overlap_end - overlap_start + min_span = min(cam_span, fib_span) + if overlap_span <= 0.0 or min_span <= 0.0: + return [] + if overlap_span < max(1.0, 0.02 * min_span): + return [] + + offsets: set[int] = set() + + def _sample_indices(indices: np.ndarray) -> np.ndarray: + idx = np.asarray(indices, int) + if idx.size <= 64: + return idx + picks = np.linspace(0, idx.size - 1, 64) + return np.unique(idx[np.round(picks).astype(int)]) + + cam_idx = _sample_indices(np.flatnonzero((cam >= overlap_start) & (cam <= overlap_end))) + fib_idx = _sample_indices(np.flatnonzero((fib >= overlap_start) & (fib <= overlap_end))) + + for i_raw in cam_idx: + i = int(i_raw) + j = int(np.searchsorted(fib, cam[i], side="left")) + for jj in (j - 1, j, j + 1): + if 0 <= jj < fib.size: + offsets.add(int(jj) - i) + + for j_raw in fib_idx: + j = int(j_raw) + i = int(np.searchsorted(cam, fib[j], side="left")) + for ii in (i - 1, i, i + 1): + if 0 <= ii < cam.size: + offsets.add(j - int(ii)) + + ranked = sorted(offsets, key=lambda val: (abs(int(val)), int(val))) + return [int(val) for val in ranked[:max(1, int(max_candidates))]] + + def _fit_linear(camera: np.ndarray, fiber: np.ndarray) -> Tuple[float, float, np.ndarray, np.ndarray]: cam = np.asarray(camera, float) fib = np.asarray(fiber, float) @@ -200,7 +253,10 @@ def match_sync_events( best: Optional[Tuple[float, int, int, np.ndarray, np.ndarray]] = None max_off = max(0, int(max_offset)) max_pairs = min(cam.size, fib.size) - for offset in range(-max_off, max_off + 1): + inferred_offsets = set(_overlap_candidate_offsets(cam, fib)) + candidate_offsets = set(range(-max_off, max_off + 1)) + candidate_offsets.update(inferred_offsets) + for offset in sorted(candidate_offsets, key=lambda val: (abs(int(val)), int(val))): c, f = _paired_by_offset(cam, fib, offset) if c.size < min_pairs: continue @@ -221,7 +277,11 @@ def match_sync_events( warnings.append("Could not robustly offset-match sync pulses; paired by order.") return cam[:n], fib[:n], 0, warnings _, neg_n, offset, c_best, f_best = best - if int(-neg_n) < min(cam.size, fib.size): + if offset in inferred_offsets and abs(int(offset)) > max_off: + warnings.append( + f"Inferred pulse offset {offset} from overlapping timestamps; check for unmatched leading sync pulses." + ) + elif int(-neg_n) < min(cam.size, fib.size): warnings.append(f"Matched with pulse offset {offset}; check for dropped leading sync pulses.") return c_best, f_best, int(offset), warnings diff --git a/tests/test_time_sync.py b/tests/test_time_sync.py index a9bc01f..f44eb6a 100644 --- a/tests/test_time_sync.py +++ b/tests/test_time_sync.py @@ -37,6 +37,18 @@ def test_interpolation_alignment_uses_pulse_pairs(self): np.testing.assert_allclose(result.aligned_time[[0, 2, 3]], [0.2, 10.1, 30.8], atol=1e-9) self.assertGreater(float(np.nanmax(np.abs(result.residuals))), 0.0) + def test_overlap_time_infers_unmatched_leading_camera_pulses(self): + camera_lead = np.array([4.0, 8.0, 8.5, 12.0, 12.5, 16.0, 16.5]) + shared = np.array([21.0, 21.5, 25.0, 25.5, 29.0, 29.5, 34.0, 34.5]) + camera_events = np.r_[camera_lead, shared] + fiber_events = shared + 0.12 + fiber_time = np.linspace(20.0, 36.0, 50) + result = align_timebase(fiber_time, camera_events, fiber_events, method="linear", max_offset=0) + self.assertEqual(result.pair_offset, -7) + self.assertEqual(result.status, "ok") + self.assertLess(result.rms_error_s, 1e-9) + np.testing.assert_allclose(result.aligned_time, fiber_time - 0.12, atol=1e-9) + if __name__ == "__main__": unittest.main() From cc4461ff6dbe3b990386d9b5e522c30eae563bd2 Mon Sep 17 00:00:00 2001 From: andrianj Date: Mon, 1 Jun 2026 17:41:28 +0200 Subject: [PATCH 6/7] Add barcode packet sync matching --- pyBer/gui_postprocessing.py | 53 +++-- pyBer/time_sync.py | 374 +++++++++++++++++++++++++++++++++--- tests/test_time_sync.py | 62 +++++- 3 files changed, 451 insertions(+), 38 deletions(-) diff --git a/pyBer/gui_postprocessing.py b/pyBer/gui_postprocessing.py index c4958a1..4e05387 100644 --- a/pyBer/gui_postprocessing.py +++ b/pyBer/gui_postprocessing.py @@ -20,7 +20,7 @@ from analysis_core import ProcessedTrial, coerce_time_value from ethovision_process_gui import clean_sheet -from time_sync import SyncResult, align_timebase, extract_sync_events +from time_sync import SyncResult, align_sync_traces, align_timebase, extract_sync_events from temporal_modeling import TemporalModelingWidget from onboarding import ( PanelHeader as _PyberPanelHeader, @@ -8482,15 +8482,12 @@ def _sync_behavior_source_for_proc(self, proc: ProcessedTrial) -> Optional[Dict[ return self._behavior_sources.get(stem) return self._match_behavior_source(proc) - def _sync_camera_events_for_proc(self, proc: ProcessedTrial) -> np.ndarray: + def _sync_camera_trace_for_proc(self, proc: ProcessedTrial) -> Tuple[np.ndarray, Optional[np.ndarray]]: info = self._sync_behavior_source_for_proc(proc) data = self.combo_sync_camera_column.currentData() if hasattr(self, "combo_sync_camera_column") else "" if not isinstance(data, str) or "::" not in data: raise ValueError("Choose a camera sync behavior or column.") kind, name = data.split("::", 1) - min_interval = float(self.spin_sync_min_interval.value()) - mode = self._sync_mode_key(self.combo_sync_camera_mode.currentText()) - threshold = None if self.cb_sync_auto_threshold.isChecked() else float(self.spin_sync_threshold.value()) if kind == "external": if "::" not in name: raise ValueError("Choose a loaded Sync signal column.") @@ -8503,7 +8500,7 @@ def _sync_camera_events_for_proc(self, proc: ProcessedTrial) -> np.ndarray: raise ValueError(f"Sync signal column not found: {column_name}") t = np.asarray(ext.get("time", np.array([], float)), float) x = np.asarray(columns[column_name], float) - return extract_sync_events(t, x, mode=mode, threshold=threshold, min_interval_s=min_interval) + return t, x if not info: raise ValueError("No behavior/camera file matched this recording.") if kind == "behavior": @@ -8511,10 +8508,10 @@ def _sync_camera_events_for_proc(self, proc: ProcessedTrial) -> np.ndarray: if name not in behaviors: raise ValueError(f"Behavior sync column not found: {name}") if str(info.get("kind", _BEHAVIOR_PARSE_BINARY)) == _BEHAVIOR_PARSE_TIMESTAMPS: - return extract_sync_events(np.asarray(behaviors[name], float), None, min_interval_s=min_interval) + return np.asarray(behaviors[name], float), None t = np.asarray(info.get("time", np.array([], float)), float) x = np.asarray(behaviors[name], float) - return extract_sync_events(t, x, mode=mode, threshold=threshold, min_interval_s=min_interval) + return t, x if kind == "trajectory": traj = info.get("trajectory") or {} if name not in traj: @@ -8523,9 +8520,18 @@ def _sync_camera_events_for_proc(self, proc: ProcessedTrial) -> np.ndarray: if t.size == 0: t = np.asarray(info.get("time", np.array([], float)), float) x = np.asarray(traj[name], float) - return extract_sync_events(t, x, mode=mode, threshold=threshold, min_interval_s=min_interval) + return t, x raise ValueError(f"Unsupported camera sync source: {kind}") + def _sync_camera_events_for_proc(self, proc: ProcessedTrial) -> np.ndarray: + t, x = self._sync_camera_trace_for_proc(proc) + min_interval = float(self.spin_sync_min_interval.value()) + if x is None: + return extract_sync_events(t, None, min_interval_s=min_interval) + mode = self._sync_mode_key(self.combo_sync_camera_mode.currentText()) + threshold = None if self.cb_sync_auto_threshold.isChecked() else float(self.spin_sync_threshold.value()) + return extract_sync_events(t, x, mode=mode, threshold=threshold, min_interval_s=min_interval) + def _sync_fiber_trace_for_proc(self, proc: ProcessedTrial) -> Tuple[np.ndarray, np.ndarray]: source = self.combo_sync_fiber_source.currentData() if hasattr(self, "combo_sync_fiber_source") else "__embedded__" if not isinstance(source, str): @@ -8579,18 +8585,35 @@ def _sync_fiber_trace_for_proc(self, proc: ProcessedTrial) -> Tuple[np.ndarray, raise ValueError(f"Unsupported photometry sync source: {source}") def _compute_time_sync_for_proc(self, proc: ProcessedTrial) -> SyncResult: - camera_events = self._sync_camera_events_for_proc(proc) + camera_time, camera_sync = self._sync_camera_trace_for_proc(proc) fiber_time, fiber_sync = self._sync_fiber_trace_for_proc(proc) - mode = self._sync_mode_key(self.combo_sync_fiber_mode.currentText()) + camera_mode = self._sync_mode_key(self.combo_sync_camera_mode.currentText()) + fiber_mode = self._sync_mode_key(self.combo_sync_fiber_mode.currentText()) threshold = None if self.cb_sync_auto_threshold.isChecked() else float(self.spin_sync_threshold.value()) + min_interval = float(self.spin_sync_min_interval.value()) + method = "interpolation" if "interp" in self.combo_sync_method.currentText().lower() else "linear" + if camera_sync is not None: + return align_sync_traces( + np.asarray(fiber_time, float), + np.asarray(camera_time, float), + np.asarray(camera_sync, float), + np.asarray(fiber_sync, float), + camera_mode=camera_mode, + fiber_mode=fiber_mode, + threshold=threshold, + min_interval_s=min_interval, + method=method, + max_offset=int(self.spin_sync_max_offset.value()), + min_pairs=2, + ) + camera_events = extract_sync_events(np.asarray(camera_time, float), None, min_interval_s=min_interval) fiber_events = extract_sync_events( fiber_time, fiber_sync, - mode=mode, + mode=fiber_mode, threshold=threshold, - min_interval_s=float(self.spin_sync_min_interval.value()), + min_interval_s=min_interval, ) - method = "interpolation" if "interp" in self.combo_sync_method.currentText().lower() else "linear" return align_timebase( np.asarray(fiber_time, float), camera_events, @@ -8618,7 +8641,7 @@ def _sync_report_text(self, proc: ProcessedTrial, result: SyncResult) -> str: lines = [ f"File: {file_id}", f"Status: {rep.get('status')} | method: {rep.get('method')}", - f"Matched pulses: {rep.get('n_matched')} / camera {rep.get('n_camera_events')} / photometry {rep.get('n_fiber_events')}", + f"Matched sync anchors: {rep.get('n_matched')} / camera {rep.get('n_camera_events')} / photometry {rep.get('n_fiber_events')}", f"Clock mapping: camera_time = {float(rep.get('slope', np.nan)):.9g} * photometry_time + {float(rep.get('intercept', np.nan)):.9g}", f"Median lag: {float(rep.get('median_lag_s', np.nan)) * 1000:.3f} ms", f"Residual RMS: {float(rep.get('rms_error_s', np.nan)) * 1000:.3f} ms", diff --git a/pyBer/time_sync.py b/pyBer/time_sync.py index 7c25e73..7a33403 100644 --- a/pyBer/time_sync.py +++ b/pyBer/time_sync.py @@ -58,6 +58,15 @@ def summary_dict(self) -> Dict[str, object]: } +@dataclass(frozen=True) +class BarcodePacket: + start_time: float + end_time: float + anchor_time: float + code: Tuple[int, ...] + n_transitions: int + + def _finite_sorted(time: np.ndarray, signal: np.ndarray) -> Tuple[np.ndarray, np.ndarray]: t = np.asarray(time, float).reshape(-1) x = np.asarray(signal, float).reshape(-1) @@ -101,6 +110,11 @@ def _deduplicate_events(events: np.ndarray, min_interval_s: float) -> np.ndarray return np.asarray(keep, float) +def _is_barcode_mode(mode: str) -> bool: + mode_l = str(mode or "").strip().lower() + return "barcode" in mode_l or "change" in mode_l or "value" in mode_l + + def extract_sync_events( time: np.ndarray, signal: Optional[np.ndarray] = None, @@ -120,7 +134,7 @@ def extract_sync_events( mode_l = str(mode or "").strip().lower() polarity_l = str(polarity or "").strip().lower() - if "barcode" in mode_l or "change" in mode_l or "value" in mode_l: + if _is_barcode_mode(mode_l): rounded = np.asarray(x, float) finite = rounded[np.isfinite(rounded)] if finite.size == 0: @@ -145,6 +159,87 @@ def extract_sync_events( return _deduplicate_events(t[idx], min_interval_s) +def decode_barcode_packets( + time: np.ndarray, + signal: np.ndarray, + *, + threshold: Optional[float] = None, + min_transitions: int = 4, +) -> List[BarcodePacket]: + """Decode binary barcode-like bursts into packet anchors and bit patterns.""" + t, x = _finite_sorted(time, signal) + if t.size < 4: + return [] + finite = x[np.isfinite(x)] + if finite.size < 4: + return [] + value_span = float(np.nanmax(finite) - np.nanmin(finite)) + if not np.isfinite(value_span) or value_span <= 0.0: + return [] + if value_span > 2.0: + return [] + + thr = _auto_threshold(x) if threshold is None or not np.isfinite(float(threshold)) else float(threshold) + values = (x > thr).astype(int) + change_idx = np.flatnonzero(values[1:] != values[:-1]) + 1 + if change_idx.size < max(2, int(min_transitions)): + return [] + + change_t = t[change_idx] + intervals = np.diff(change_t) + intervals = intervals[np.isfinite(intervals) & (intervals > 0.0)] + if intervals.size < 2: + return [] + short_dt = float(np.nanpercentile(intervals, 35)) + med_dt = float(np.nanmedian(intervals)) + if not np.isfinite(short_dt) or short_dt <= 0.0: + short_dt = med_dt + if not np.isfinite(med_dt) or med_dt <= 0.0: + return [] + gap_threshold = max(0.35, 4.0 * short_dt, 2.5 * med_dt) + split_after = np.flatnonzero(np.diff(change_t) > gap_threshold) + if split_after.size == 0: + return [] + + starts = np.r_[0, split_after + 1] + stops = np.r_[split_after + 1, change_idx.size] + packets: List[BarcodePacket] = [] + for start, stop in zip(starts, stops): + idx = change_idx[int(start):int(stop)] + if idx.size < max(2, int(min_transitions)): + continue + trans_t = t[idx] + trans_values = values[idx] + local_dt = np.diff(trans_t) + local_dt = local_dt[np.isfinite(local_dt) & (local_dt > 0.0)] + if local_dt.size < 2: + continue + unit = float(np.nanpercentile(local_dt, 35)) + if not np.isfinite(unit) or unit <= 0.0: + unit = float(np.nanmedian(local_dt)) + if not np.isfinite(unit) or unit <= 0.0: + continue + units = np.clip(np.rint(np.diff(trans_t) / unit).astype(int), 1, 16) + if units.size == 0: + continue + bits: List[int] = [] + for val, reps in zip(trans_values[:-1], units): + bits.extend([int(val)] * int(reps)) + bits.append(int(trans_values[-1])) + if len(bits) < 3 or len(set(bits)) < 2: + continue + packets.append( + BarcodePacket( + start_time=float(trans_t[0]), + end_time=float(trans_t[-1]), + anchor_time=float(trans_t[0]), + code=tuple(bits), + n_transitions=int(idx.size), + ) + ) + return packets + + def _paired_by_offset( camera_events: np.ndarray, fiber_events: np.ndarray, @@ -217,6 +312,149 @@ def _sample_indices(indices: np.ndarray) -> np.ndarray: return [int(val) for val in ranked[:max(1, int(max_candidates))]] +def _paired_packets_by_offset( + camera_packets: List[BarcodePacket], + fiber_packets: List[BarcodePacket], + offset: int, +) -> Tuple[List[BarcodePacket], List[BarcodePacket]]: + if offset >= 0: + n = min(len(camera_packets), len(fiber_packets) - offset) + if n <= 0: + return [], [] + return camera_packets[:n], fiber_packets[offset:offset + n] + start_cam = -offset + n = min(len(camera_packets) - start_cam, len(fiber_packets)) + if n <= 0: + return [], [] + return camera_packets[start_cam:start_cam + n], fiber_packets[:n] + + +def _barcode_code_distance(a: Tuple[int, ...], b: Tuple[int, ...]) -> float: + if not a and not b: + return 0.0 + if not a or not b: + return 1.0 + prev = list(range(len(b) + 1)) + for i, aval in enumerate(a, start=1): + cur = [i] + [0] * len(b) + for j, bval in enumerate(b, start=1): + cur[j] = min( + prev[j] + 1, + cur[j - 1] + 1, + prev[j - 1] + (0 if aval == bval else 1), + ) + prev = cur + return float(prev[-1]) / float(max(len(a), len(b), 1)) + + +def match_barcode_packets( + camera_packets: List[BarcodePacket], + fiber_packets: List[BarcodePacket], + *, + max_offset: int = 5, + min_pairs: int = 2, +) -> Tuple[np.ndarray, np.ndarray, int, List[str]]: + """Match barcode packets by decoded code, then return their anchor times.""" + warnings: List[str] = [] + if len(camera_packets) < max(1, int(min_pairs)) or len(fiber_packets) < max(1, int(min_pairs)): + return np.array([], float), np.array([], float), 0, ["Not enough decoded barcode packets."] + + cam_anchors = np.asarray([pkt.anchor_time for pkt in camera_packets], float) + fib_anchors = np.asarray([pkt.anchor_time for pkt in fiber_packets], float) + + cam_step = float(np.nanmedian(np.diff(cam_anchors))) if cam_anchors.size > 1 else 1.0 + fib_step = float(np.nanmedian(np.diff(fib_anchors))) if fib_anchors.size > 1 else 1.0 + typical_step = float(np.nanmedian([cam_step, fib_step])) + if not np.isfinite(typical_step) or typical_step <= 0.0: + typical_step = 1.0 + time_tolerance = max(0.35, min(2.0, 0.45 * typical_step)) + bin_width = max(0.05, min(0.5, 0.12 * typical_step)) + max_code_distance = 0.35 + + pair_rows: List[Tuple[int, int, float, float]] = [] + for ci, c_pkt in enumerate(camera_packets): + for fi, f_pkt in enumerate(fiber_packets): + dist = _barcode_code_distance(c_pkt.code, f_pkt.code) + if dist <= max_code_distance: + pair_rows.append((ci, fi, float(c_pkt.anchor_time - f_pkt.anchor_time), float(dist))) + if not pair_rows: + return np.array([], float), np.array([], float), 0, ["No decoded barcode packet identities matched."] + + bins: Dict[int, List[Tuple[int, int, float, float]]] = {} + for row in pair_rows: + key = int(np.round(row[2] / bin_width)) + bins.setdefault(key, []).append(row) + ranked_bins = sorted( + bins.items(), + key=lambda item: (-len(item[1]), float(np.nanmean([row[3] for row in item[1]])), abs(item[0])), + )[:30] + + best: Optional[Tuple[float, int, int, np.ndarray, np.ndarray, float, float]] = None + for _bin_key, rows in ranked_bins: + lag0 = float(np.nanmedian([row[2] for row in rows])) + used_camera: set[int] = set() + matched: List[Tuple[int, int, float, float, float]] = [] + for fi, f_pkt in enumerate(fiber_packets): + predicted = float(f_pkt.anchor_time + lag0) + left = int(np.searchsorted(cam_anchors, predicted - time_tolerance, side="left")) + right = int(np.searchsorted(cam_anchors, predicted + time_tolerance, side="right")) + best_local: Optional[Tuple[float, int, float, float]] = None + for ci in range(left, right): + if ci in used_camera: + continue + dist = _barcode_code_distance(camera_packets[ci].code, f_pkt.code) + if dist > max_code_distance: + continue + dt = abs(float(cam_anchors[ci] - predicted)) + local_score = dt + time_tolerance * dist + candidate = (local_score, ci, dt, dist) + if best_local is None or candidate < best_local: + best_local = candidate + if best_local is None: + continue + _local_score, ci, dt, dist = best_local + used_camera.add(int(ci)) + matched.append((int(ci), int(fi), float(cam_anchors[ci]), float(f_pkt.anchor_time), float(dist))) + if len(matched) < min_pairs: + continue + c = np.asarray([row[2] for row in matched], float) + f = np.asarray([row[3] for row in matched], float) + dists = np.asarray([row[4] for row in matched], float) + slope, intercept, fitted, resid = _fit_linear(c, f) + finite = resid[np.isfinite(resid)] + if finite.size == 0: + continue + rms = float(np.sqrt(np.nanmean(finite ** 2))) + if finite.size >= 4: + keep = np.abs(finite - float(np.nanmedian(finite))) <= max(0.15, 0.5 * time_tolerance) + if int(np.sum(keep)) >= min_pairs and int(np.sum(keep)) < finite.size: + c = c[keep] + f = f[keep] + dists = dists[keep] + slope, intercept, fitted, resid = _fit_linear(c, f) + finite = resid[np.isfinite(resid)] + if finite.size == 0: + continue + rms = float(np.sqrt(np.nanmean(finite ** 2))) + mean_dist = float(np.nanmean(dists)) if dists.size else 1.0 + median_offset = int(np.round(np.nanmedian([row[1] - row[0] for row in matched]))) + score = rms + 0.05 * mean_dist + 0.5 / np.sqrt(max(1, int(c.size))) + 0.0001 * abs(lag0) + candidate = (score, -int(c.size), median_offset, c, f, rms, mean_dist) + if best is None or candidate[:3] < best[:3]: + best = candidate + + if best is None: + return np.array([], float), np.array([], float), 0, ["No decoded barcode packet identities matched."] + + _score, neg_n, offset, c_best, f_best, _rms, mean_dist = best + if int(-neg_n) < min(len(camera_packets), len(fiber_packets)): + warnings.append(f"Matched barcode packets with median packet offset {offset}.") + if mean_dist > 0.0: + warnings.append(f"Mean barcode identity edit distance: {mean_dist:.3g}.") + warnings.append(f"Decoded barcode packets: camera {len(camera_packets)}, photometry {len(fiber_packets)}.") + return np.asarray(c_best, float), np.asarray(f_best, float), int(offset), warnings + + def _fit_linear(camera: np.ndarray, fiber: np.ndarray) -> Tuple[float, float, np.ndarray, np.ndarray]: cam = np.asarray(camera, float) fib = np.asarray(fiber, float) @@ -307,26 +545,24 @@ def _interp_with_linear_extrapolation(x: np.ndarray, xp: np.ndarray, fp: np.ndar return out -def align_timebase( +def _sync_result_from_matches( fiber_time: np.ndarray, camera_events: np.ndarray, fiber_events: np.ndarray, + matched_camera_events: np.ndarray, + matched_fiber_events: np.ndarray, *, - method: str = "linear", - max_offset: int = 5, - min_pairs: int = 2, + method: str, + pair_offset: int, + warnings: List[str], + method_prefix: str = "", ) -> SyncResult: - """Return a camera-time vector for each photometry sample.""" ft = np.asarray(fiber_time, float).reshape(-1) - cam, fib, offset, warnings = match_sync_events( - camera_events, - fiber_events, - max_offset=max_offset, - min_pairs=min_pairs, - ) + cam = np.asarray(matched_camera_events, float).reshape(-1) + fib = np.asarray(matched_fiber_events, float).reshape(-1) if cam.size == 0 or fib.size == 0: return SyncResult( - method=str(method), + method=str(method_prefix or method), status="failed", aligned_time=np.full(ft.shape, np.nan, dtype=float), camera_events=np.asarray(camera_events, float), @@ -335,21 +571,19 @@ def align_timebase( matched_fiber_events=fib, fitted_camera_events=np.array([], float), residuals=np.array([], float), - pair_offset=offset, - warnings=warnings, + pair_offset=int(pair_offset), + warnings=list(warnings), ) method_l = str(method or "").strip().lower() slope, intercept, fitted, residuals = _fit_linear(cam, fib) if "interp" in method_l and cam.size >= 2: aligned = _interp_with_linear_extrapolation(ft, fib, cam) - method_out = "interpolation" - # Keep the linear fit and its residuals for diagnostics. Interpolation - # passes exactly through matched pulses, so interpolation residuals would - # hide dropped-pulse or non-linear clock problems in the QC report. + base_method = "interpolation" else: aligned = slope * ft + intercept if np.isfinite(slope) and np.isfinite(intercept) else np.full(ft.shape, np.nan) - method_out = "linear_regression" + base_method = "linear_regression" + method_out = f"{method_prefix}_{base_method}" if method_prefix else base_method finite_resid = residuals[np.isfinite(residuals)] if finite_resid.size: @@ -366,12 +600,13 @@ def align_timebase( drift_ppm = float((slope - 1.0) * 1e6) if np.isfinite(slope) else float("nan") status = "ok" + result_warnings = list(warnings) if cam.size < 3: status = "warning" - warnings.append("Fewer than 3 matched sync events; inspect the residual plot.") + result_warnings.append("Fewer than 3 matched sync events; inspect the residual plot.") if np.isfinite(rms) and rms > 0.2: status = "warning" - warnings.append(f"High sync residual RMS ({rms * 1000:.1f} ms).") + result_warnings.append(f"High sync residual RMS ({rms * 1000:.1f} ms).") return SyncResult( method=method_out, @@ -390,6 +625,101 @@ def align_timebase( max_abs_error_s=max_abs, median_lag_s=median_lag, drift_ppm=drift_ppm, + pair_offset=int(pair_offset), + warnings=result_warnings, + ) + + +def align_timebase( + fiber_time: np.ndarray, + camera_events: np.ndarray, + fiber_events: np.ndarray, + *, + method: str = "linear", + max_offset: int = 5, + min_pairs: int = 2, +) -> SyncResult: + """Return a camera-time vector for each photometry sample.""" + cam, fib, offset, warnings = match_sync_events( + camera_events, + fiber_events, + max_offset=max_offset, + min_pairs=min_pairs, + ) + return _sync_result_from_matches( + fiber_time, + camera_events, + fiber_events, + cam, + fib, + method=method, pair_offset=offset, warnings=warnings, ) + + +def align_sync_traces( + fiber_time: np.ndarray, + camera_time: np.ndarray, + camera_signal: np.ndarray, + fiber_signal: np.ndarray, + *, + camera_mode: str = "ttl_rising", + fiber_mode: str = "ttl_rising", + threshold: Optional[float] = None, + min_interval_s: float = 0.2, + method: str = "linear", + max_offset: int = 5, + min_pairs: int = 2, +) -> SyncResult: + """Align raw sync traces, using decoded barcode packets when available.""" + barcode_requested = _is_barcode_mode(camera_mode) or _is_barcode_mode(fiber_mode) + camera_packets = decode_barcode_packets(camera_time, camera_signal, threshold=threshold) + fiber_packets = decode_barcode_packets(fiber_time, fiber_signal, threshold=threshold) + if len(camera_packets) >= min_pairs and len(fiber_packets) >= min_pairs: + cam, fib, offset, warnings = match_barcode_packets( + camera_packets, + fiber_packets, + max_offset=max_offset, + min_pairs=min_pairs, + ) + if cam.size >= min_pairs and fib.size >= min_pairs: + if not barcode_requested: + warnings.insert(0, "Auto-detected barcode packets in both sync traces.") + return _sync_result_from_matches( + fiber_time, + np.asarray([pkt.anchor_time for pkt in camera_packets], float), + np.asarray([pkt.anchor_time for pkt in fiber_packets], float), + cam, + fib, + method=method, + pair_offset=offset, + warnings=warnings, + method_prefix="barcode_packets", + ) + + camera_events = extract_sync_events( + camera_time, + camera_signal, + mode=camera_mode, + threshold=threshold, + min_interval_s=min_interval_s, + ) + fiber_events = extract_sync_events( + fiber_time, + fiber_signal, + mode=fiber_mode, + threshold=threshold, + min_interval_s=min_interval_s, + ) + result = align_timebase( + fiber_time, + camera_events, + fiber_events, + method=method, + max_offset=max_offset, + min_pairs=min_pairs, + ) + if barcode_requested: + result.warnings.insert(0, "No reliable barcode packets decoded; fell back to edge-train matching.") + return result diff --git a/tests/test_time_sync.py b/tests/test_time_sync.py index f44eb6a..09c94a4 100644 --- a/tests/test_time_sync.py +++ b/tests/test_time_sync.py @@ -7,10 +7,21 @@ ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), "..")) sys.path.insert(0, os.path.join(ROOT, "pyBer")) -from time_sync import align_timebase, extract_sync_events # noqa: E402 +from time_sync import align_sync_traces, align_timebase, decode_barcode_packets, extract_sync_events # noqa: E402 class TimeSyncTests(unittest.TestCase): + def _barcode_trace(self, packet_starts, packet_codes, dt=0.05, fs=200.0): + end = float(max(packet_starts) + max(len(code) for code in packet_codes) * dt + 1.0) + t = np.arange(0.0, end, 1.0 / fs) + x = np.zeros_like(t) + for start, code in zip(packet_starts, packet_codes): + for bit_idx, bit in enumerate(code): + lo = float(start + bit_idx * dt) + hi = float(start + (bit_idx + 1) * dt) + x[(t >= lo) & (t < hi)] = float(bit) + return t, x + def test_extract_ttl_rising_edges(self): t = np.arange(0, 10, 0.1) x = np.zeros_like(t) @@ -49,6 +60,55 @@ def test_overlap_time_infers_unmatched_leading_camera_pulses(self): self.assertLess(result.rms_error_s, 1e-9) np.testing.assert_allclose(result.aligned_time, fiber_time - 0.12, atol=1e-9) + def test_barcode_packets_match_by_identity_before_fitting(self): + codes = [ + (1, 0, 1, 1, 0, 0, 1), + (1, 1, 0, 1, 0, 1, 0), + (1, 0, 0, 1, 1, 0, 1), + (1, 1, 1, 0, 0, 1, 0), + ] + camera_starts = np.array([2.0, 6.0, 10.0, 14.0, 18.0, 22.0]) + camera_codes = [codes[0], codes[1], codes[2], codes[3], codes[1], codes[2]] + fiber_starts = np.array([10.15, 14.15, 18.15, 22.15]) + fiber_codes = [codes[2], codes[3], codes[1], codes[2]] + camera_t, camera_x = self._barcode_trace(camera_starts, camera_codes) + fiber_t, fiber_x = self._barcode_trace(fiber_starts, fiber_codes) + result = align_sync_traces( + fiber_t, + camera_t, + camera_x, + fiber_x, + camera_mode="ttl_rising", + fiber_mode="ttl_rising", + method="linear", + max_offset=0, + ) + self.assertEqual(result.method, "barcode_packets_linear_regression") + self.assertEqual(result.status, "ok") + self.assertGreaterEqual(result.matched_camera_events.size, 4) + self.assertLess(result.rms_error_s, 0.01) + np.testing.assert_allclose(result.aligned_time, fiber_t - 0.15, atol=0.02) + + def test_regular_10hz_ttl_does_not_require_barcode_packets(self): + t = np.arange(0.0, 10.0, 0.002) + camera_x = ((t % 0.1) < 0.03).astype(float) + fiber_t = t + 0.2 + fiber_x = ((fiber_t % 0.1) < 0.03).astype(float) + self.assertEqual(decode_barcode_packets(t, camera_x), []) + result = align_sync_traces( + fiber_t, + t, + camera_x, + fiber_x, + camera_mode="ttl_rising", + fiber_mode="ttl_rising", + method="linear", + max_offset=5, + ) + self.assertNotIn("barcode_packets", result.method) + self.assertEqual(result.status, "ok") + self.assertGreater(result.matched_camera_events.size, 20) + if __name__ == "__main__": unittest.main() From 5b04d2bb285a418c1fe8366f49410541168ecdd5 Mon Sep 17 00:00:00 2001 From: andrianj Date: Mon, 1 Jun 2026 17:53:42 +0200 Subject: [PATCH 7/7] Move sync after behavior in postprocessing rail --- pyBer/gui_postprocessing.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pyBer/gui_postprocessing.py b/pyBer/gui_postprocessing.py index 4e05387..980b809 100644 --- a/pyBer/gui_postprocessing.py +++ b/pyBer/gui_postprocessing.py @@ -3020,7 +3020,7 @@ def _sync_form(panel: QtWidgets.QWidget) -> QtWidgets.QFormLayout: rail_layout = QtWidgets.QVBoxLayout(self._post_side_rail) rail_layout.setContentsMargins(8, 10, 8, 10) rail_layout.setSpacing(6) - for key in ("setup", "sync", "psth", "spatial", "temporal", "signal", "behavior", "export"): + for key in ("setup", "psth", "spatial", "temporal", "signal", "behavior", "sync", "export"): rail_layout.addWidget(self._section_buttons[key], 0, QtCore.Qt.AlignmentFlag.AlignHCenter) rail_layout.addStretch(1)