From 5941c7f0ef9b281299ef42dfbdcff924d2fda79c Mon Sep 17 00:00:00 2001 From: Guillaume Favelier Date: Thu, 8 Oct 2020 12:43:30 +0200 Subject: [PATCH 1/4] Move time viewer code into Brain --- mne/icons/README.rst | 2 +- mne/viz/_3d.py | 14 +- mne/viz/_brain/__init__.py | 1 - mne/viz/_brain/_brain.py | 1136 +++++++++++++++++++++++++++- mne/viz/_brain/_linkviewer.py | 15 +- mne/viz/_brain/_notebook.py | 11 +- mne/viz/_brain/_timeviewer.py | 1049 ------------------------- mne/viz/_brain/mplcanvas.py | 4 +- mne/viz/_brain/tests/test_brain.py | 101 ++- 9 files changed, 1176 insertions(+), 1157 deletions(-) delete mode 100644 mne/viz/_brain/_timeviewer.py diff --git a/mne/icons/README.rst b/mne/icons/README.rst index 7fb5052d4a6..70499a98fa2 100644 --- a/mne/icons/README.rst +++ b/mne/icons/README.rst @@ -4,7 +4,7 @@ Documentation ============= -The icons are used in ``mne/viz/_brain`` for the toolbar of ``_TimeViewer``. +The icons are used in ``mne/viz/_brain/_Brain.py`` for the toolbar. It is necessary to compile those icons into a resource file for proper use by the application. diff --git a/mne/viz/_3d.py b/mne/viz/_3d.py index a717194d53c..adcb2e36cfe 100644 --- a/mne/viz/_3d.py +++ b/mne/viz/_3d.py @@ -1557,19 +1557,17 @@ def link_brains(brains, time=True, camera=False, colorbar=True, if _get_3d_backend() != 'pyvista': raise NotImplementedError("Expected 3d backend is pyvista but" " {} was given.".format(_get_3d_backend())) - from ._brain import Brain, _TimeViewer, _LinkViewer + from ._brain import Brain, _LinkViewer if not isinstance(brains, Iterable): brains = [brains] if len(brains) == 0: raise ValueError("The collection of brains is empty.") for brain in brains: - if isinstance(brain, Brain): - # check if the _TimeViewer wrapping is not already applied - if not hasattr(brain, 'time_viewer') or brain.time_viewer is None: - brain = _TimeViewer(brain) - else: + if not isinstance(brain, Brain): raise TypeError("Expected type is Brain but" " {} was given.".format(type(brain))) + # enable time viewer if necessary + brain.setup_time_viewer() subjects = [brain._subject_id for brain in brains] if subjects.count(subjects[0]) != len(subjects): raise RuntimeError("Cannot link brains from different subjects.") @@ -1949,8 +1947,8 @@ def _plot_stc(stc, subject, surface, hemi, colormap, time_label, from surfer import TimeViewer TimeViewer(brain) else: # PyVista - from ._brain import _TimeViewer as TimeViewer - TimeViewer(brain, show_traces=show_traces) + brain.setup_time_viewer(time_viewer=time_viewer, + show_traces=show_traces) return brain diff --git a/mne/viz/_brain/__init__.py b/mne/viz/_brain/__init__.py index 4fdca2ab96f..525cde1b9d9 100644 --- a/mne/viz/_brain/__init__.py +++ b/mne/viz/_brain/__init__.py @@ -11,7 +11,6 @@ from ._brain import Brain from ._scraper import _BrainScraper -from ._timeviewer import _TimeViewer from ._linkviewer import _LinkViewer __all__ = ['Brain'] diff --git a/mne/viz/_brain/_brain.py b/mne/viz/_brain/_brain.py index 22befe99d8a..14adbbe48ff 100644 --- a/mne/viz/_brain/_brain.py +++ b/mne/viz/_brain/_brain.py @@ -7,25 +7,44 @@ # # License: Simplified BSD +import contextlib from functools import partial import os import os.path as op +import sys +import time +import traceback +import warnings import numpy as np from scipy import sparse from .colormap import calculate_lut from .surface import Surface -from .view import views_dicts +from .view import views_dicts, _lh_views_dict +from .mplcanvas import MplCanvas +from .callback import (ShowView, IntSlider, TimeSlider, SmartSlider, + BumpColorbarPoints, UpdateColorbarScale) +from ..utils import _show_help, _get_color_list from .._3d import _process_clim, _handle_time, _check_views +from ...externals.decorator import decorator from ...defaults import _handle_default from ...surface import mesh_edges -from ...source_space import SourceSpaces +from ...source_space import SourceSpaces, vertex_to_mni, _read_talxfm from ...transforms import apply_trans from ...utils import (_check_option, logger, verbose, fill_doc, _validate_type, - use_log_level, Bunch) + use_log_level, Bunch, _ReuseCycle, warn) + + +@decorator +def safe_event(fun, *args, **kwargs): + """Protect against PyQt5 exiting on event-handling errors.""" + try: + return fun(*args, **kwargs) + except Exception: + traceback.print_exc(file=sys.stderr) @fill_doc @@ -213,6 +232,7 @@ def __init__(self, subject_id, hemi, surf, title=None, 'sequence of ints.') self._size = size if len(size) == 2 else size * 2 # 1-tuple to 2-tuple + self.time_viewer = False self._notebook = (_get_3d_backend() == "notebook") self._hemi = hemi self._units = units @@ -293,8 +313,6 @@ def __init__(self, subject_id, hemi, surf, title=None, self.interaction = interaction self._closed = False - if show: - self._renderer.show() # update the views once the geometry is all set for h in self._hemis: for ri, ci, v in self._iter_views(h): @@ -306,6 +324,980 @@ def __init__(self, subject_id, hemi, surf, title=None, if hemi == 'rh' and hasattr(self._renderer, "_orient_lights"): self._renderer._orient_lights() + if show: + self.show() + + def setup_time_viewer(self, time_viewer=True, show_traces=True): + """Configure the time viewer parameters. + + time_viewer : bool + If True, enable widgets interaction. Defaults to True. + + show_traces : bool + If True, enable visualization of time traces. Defaults to True. + """ + if self.time_viewer: + return + self.time_viewer = time_viewer + self.orientation = list(_lh_views_dict.keys()) + self.default_smoothing_range = [0, 15] + + # detect notebook + if self._notebook: + self.notebook = True + self._configure_notebook() + return + else: + self.notebook = False + + # Default configuration + self.playback = False + self.visibility = False + self.refresh_rate_ms = max(int(round(1000. / 60.)), 1) + self.default_scaling_range = [0.2, 2.0] + self.default_playback_speed_range = [0.01, 1] + self.default_playback_speed_value = 0.05 + self.default_status_bar_msg = "Press ? for help" + all_keys = ('lh', 'rh', 'vol') + self.act_data_smooth = {key: (None, None) for key in all_keys} + self.color_cycle = None + self.mpl_canvas = None + self.picked_points = {key: list() for key in all_keys} + self.pick_table = dict() + self._mouse_no_mvt = -1 + self.icons = dict() + self.actions = dict() + self.callbacks = dict() + self.sliders = dict() + self.keys = ('fmin', 'fmid', 'fmax') + self.slider_length = 0.02 + self.slider_width = 0.04 + self.slider_color = (0.43137255, 0.44313725, 0.45882353) + self.slider_tube_width = 0.04 + self.slider_tube_color = (0.69803922, 0.70196078, 0.70980392) + + # Direct access parameters: + self.plotter = self._renderer.plotter + self.main_menu = self.plotter.main_menu + self.window = self.plotter.app_window + self.tool_bar = self.window.addToolBar("toolbar") + self.status_bar = self.window.statusBar() + self.interactor = self.plotter.interactor + self.window.signal_close.connect(self._clean) + + # Derived parameters: + self.playback_speed = self.default_playback_speed_value + _validate_type(show_traces, (bool, str, 'numeric'), 'show_traces') + self.interactor_fraction = 0.25 + if isinstance(show_traces, str): + assert 'show_traces' == 'separate' # should be guaranteed earlier + self.show_traces = True + self.separate_canvas = True + else: + if isinstance(show_traces, bool): + self.show_traces = show_traces + else: + show_traces = float(show_traces) + if not 0 < show_traces < 1: + raise ValueError( + 'show traces, if numeric, must be between 0 and 1, ' + f'got {show_traces}') + self.show_traces = True + self.interactor_fraction = show_traces + self.separate_canvas = False + del show_traces + + self._spheres = list() + self._load_icons() + self._configure_time_label() + self._configure_sliders() + self._configure_scalar_bar() + self._configure_playback() + self._configure_point_picking() + self._configure_menu() + self._configure_tool_bar() + self._configure_status_bar() + + # show everything at the end + self.toggle_interface() + with self.ensure_minimum_sizes(): + self.show() + + @safe_event + def _clean(self): + # resolve the reference cycle + self.clear_points() + self._clear_callbacks() + self.actions.clear() + self.sliders.clear() + self.reps = None + self.plotter = None + self.main_menu = None + self.window = None + self.tool_bar = None + self.status_bar = None + self.interactor = None + if self.mpl_canvas is not None: + self.mpl_canvas.clear() + self.mpl_canvas = None + self.time_actor = None + self.picked_renderer = None + for key in list(self.act_data_smooth.keys()): + self.act_data_smooth[key] = None + + @contextlib.contextmanager + def ensure_minimum_sizes(self): + """Ensure that widgets respect the windows size.""" + from ..backends._pyvista import _process_events + sz = self._size + adjust_mpl = self.show_traces and not self.separate_canvas + if not adjust_mpl: + yield + else: + mpl_h = int(round((sz[1] * self.interactor_fraction) / + (1 - self.interactor_fraction))) + self.mpl_canvas.canvas.setMinimumSize(sz[0], mpl_h) + try: + yield + finally: + self.splitter.setSizes([sz[1], mpl_h]) + _process_events(self.plotter) + _process_events(self.plotter) + self.mpl_canvas.canvas.setMinimumSize(0, 0) + _process_events(self.plotter) + _process_events(self.plotter) + # sizes could change, update views + for hemi in ('lh', 'rh'): + for ri, ci, v in self._iter_views(hemi): + self.show_view(view=v, row=ri, col=ci) + _process_events(self.plotter) + + def toggle_interface(self, value=None): + """Toggle the interface. + + value : bool | None + If True, the widgets are shown and if False, they + are hidden. If None, the state of the widgets is + toggled. Defaults to None. + """ + if value is None: + self.visibility = not self.visibility + else: + self.visibility = value + + # update tool bar icon + if self.visibility: + self.actions["visibility"].setIcon(self.icons["visibility_on"]) + else: + self.actions["visibility"].setIcon(self.icons["visibility_off"]) + + # manage sliders + for slider in self.plotter.slider_widgets: + slider_rep = slider.GetRepresentation() + if self.visibility: + slider_rep.VisibilityOn() + else: + slider_rep.VisibilityOff() + + # manage time label + time_label = self._data['time_label'] + # if we actually have time points, we will show the slider so + # hide the time actor + have_ts = self._times is not None and len(self._times) > 1 + if self.time_actor is not None: + if self.visibility and time_label is not None and not have_ts: + self.time_actor.SetInput(time_label(self._current_time)) + self.time_actor.VisibilityOn() + else: + self.time_actor.VisibilityOff() + + self.plotter.update() + + def apply_auto_scaling(self): + """Detect automatically fitting scaling parameters.""" + self._update_auto_scaling() + for key in ('fmin', 'fmid', 'fmax'): + self.reps[key].SetValue(self._data[key]) + self.plotter.update() + + def restore_user_scaling(self): + """Restore original scaling parameters.""" + self._update_auto_scaling(restore=True) + for key in ('fmin', 'fmid', 'fmax'): + self.reps[key].SetValue(self._data[key]) + self.plotter.update() + + def toggle_playback(self, value=None): + """Toggle time playback. + + value : bool | None + If True, automatic time playback is enabled and if False, + it's disabled. If None, the state of time playback is toggled. + Defaults to None. + """ + if value is None: + self.playback = not self.playback + else: + self.playback = value + + # update tool bar icon + if self.playback: + self.actions["play"].setIcon(self.icons["pause"]) + else: + self.actions["play"].setIcon(self.icons["play"]) + + if self.playback: + time_data = self._data['time'] + max_time = np.max(time_data) + if self._current_time == max_time: # start over + self.set_time_point(0) # first index + self._last_tick = time.time() + + def reset(self): + """Reset view and time step.""" + self.reset_view() + max_time = len(self._data['time']) - 1 + if max_time > 0: + self.callbacks["time"]( + self._data["initial_time_idx"], + update_widget=True, + ) + self.plotter.update() + + def set_playback_speed(self, speed): + """Set the time playback speed.""" + self.playback_speed = speed + + @safe_event + def _play(self): + if self.playback: + try: + self._advance() + except Exception: + self.toggle_playback(value=False) + raise + + def _advance(self): + this_time = time.time() + delta = this_time - self._last_tick + self._last_tick = time.time() + time_data = self._data['time'] + times = np.arange(self._n_times) + time_shift = delta * self.playback_speed + max_time = np.max(time_data) + time_point = min(self._current_time + time_shift, max_time) + # always use linear here -- this does not determine the data + # interpolation mode, it just finds where we are (in time) in + # terms of the time indices + idx = np.interp(time_point, time_data, times) + self.callbacks["time"](idx, update_widget=True) + if time_point == max_time: + self.toggle_playback(value=False) + + def _set_slider_style(self): + for slider in self.sliders.values(): + if slider is not None: + slider_rep = slider.GetRepresentation() + slider_rep.SetSliderLength(self.slider_length) + slider_rep.SetSliderWidth(self.slider_width) + slider_rep.SetTubeWidth(self.slider_tube_width) + slider_rep.GetSliderProperty().SetColor(self.slider_color) + slider_rep.GetTubeProperty().SetColor(self.slider_tube_color) + slider_rep.GetLabelProperty().SetShadow(False) + slider_rep.GetLabelProperty().SetBold(True) + slider_rep.GetLabelProperty().SetColor(self._fg_color) + slider_rep.GetTitleProperty().ShallowCopy( + slider_rep.GetLabelProperty() + ) + slider_rep.GetCapProperty().SetOpacity(0) + + def _configure_notebook(self): + from ._notebook import _NotebookInteractor + self._renderer.figure.display = _NotebookInteractor(self) + + def _configure_time_label(self): + self.time_actor = self._data.get('time_actor') + if self.time_actor is not None: + self.time_actor.SetPosition(0.5, 0.03) + self.time_actor.GetTextProperty().SetJustificationToCentered() + self.time_actor.GetTextProperty().BoldOn() + self.time_actor.VisibilityOff() + + def _configure_scalar_bar(self): + if self._colorbar_added: + scalar_bar = self.plotter.scalar_bar + scalar_bar.SetOrientationToVertical() + scalar_bar.SetHeight(0.6) + scalar_bar.SetWidth(0.05) + scalar_bar.SetPosition(0.02, 0.2) + + def _configure_sliders(self): + # Orientation slider + # Use 'lh' as a reference for orientation for 'both' + if self._hemi == 'both': + hemis_ref = ['lh'] + else: + hemis_ref = self._hemis + for hemi in hemis_ref: + for ri, ci, view in self._iter_views(hemi): + orientation_name = f"orientation_{hemi}_{ri}_{ci}" + self.plotter.subplot(ri, ci) + if view == 'flat': + self.callbacks[orientation_name] = None + continue + self.callbacks[orientation_name] = ShowView( + plotter=self.plotter, + brain=self, + orientation=self.orientation, + hemi=hemi, + row=ri, + col=ci, + ) + self.sliders[orientation_name] = \ + self.plotter.add_text_slider_widget( + self.callbacks[orientation_name], + value=0, + data=self.orientation, + pointa=(0.82, 0.74), + pointb=(0.98, 0.74), + event_type='always' + ) + orientation_rep = \ + self.sliders[orientation_name].GetRepresentation() + orientation_rep.ShowSliderLabelOff() + self.callbacks[orientation_name].slider_rep = orientation_rep + self.callbacks[orientation_name](view, update_widget=True) + + # Put other sliders on the bottom right view + ri, ci = np.array(self._subplot_shape) - 1 + self.plotter.subplot(ri, ci) + + # Smoothing slider + self.callbacks["smoothing"] = IntSlider( + plotter=self.plotter, + callback=self.set_data_smoothing, + first_call=False, + ) + self.sliders["smoothing"] = self.plotter.add_slider_widget( + self.callbacks["smoothing"], + value=self._data['smoothing_steps'], + rng=self.default_smoothing_range, title="smoothing", + pointa=(0.82, 0.90), + pointb=(0.98, 0.90) + ) + self.callbacks["smoothing"].slider_rep = \ + self.sliders["smoothing"].GetRepresentation() + + # Time slider + max_time = len(self._data['time']) - 1 + # VTK on macOS bombs if we create these then hide them, so don't + # even create them + if max_time < 1: + self.callbacks["time"] = None + self.sliders["time"] = None + else: + self.callbacks["time"] = TimeSlider( + plotter=self.plotter, + brain=self, + first_call=False, + callback=self.plot_time_line, + ) + self.sliders["time"] = self.plotter.add_slider_widget( + self.callbacks["time"], + value=self._data['time_idx'], + rng=[0, max_time], + pointa=(0.23, 0.1), + pointb=(0.77, 0.1), + event_type='always' + ) + self.callbacks["time"].slider_rep = \ + self.sliders["time"].GetRepresentation() + # configure properties of the time slider + self.sliders["time"].GetRepresentation().SetLabelFormat( + 'idx=%0.1f') + + current_time = self._current_time + assert current_time is not None # should never be the case, float + time_label = self._data['time_label'] + if callable(time_label): + current_time = time_label(current_time) + else: + current_time = time_label + if self.sliders["time"] is not None: + self.sliders["time"].GetRepresentation().SetTitleText(current_time) + if self.time_actor is not None: + self.time_actor.SetInput(current_time) + del current_time + + # Playback speed slider + if self.sliders["time"] is None: + self.callbacks["playback_speed"] = None + self.sliders["playback_speed"] = None + else: + self.callbacks["playback_speed"] = SmartSlider( + plotter=self.plotter, + callback=self.set_playback_speed, + ) + self.sliders["playback_speed"] = self.plotter.add_slider_widget( + self.callbacks["playback_speed"], + value=self.default_playback_speed_value, + rng=self.default_playback_speed_range, title="speed", + pointa=(0.02, 0.1), + pointb=(0.18, 0.1), + event_type='always' + ) + self.callbacks["playback_speed"].slider_rep = \ + self.sliders["playback_speed"].GetRepresentation() + + # Colormap slider + pointa = np.array((0.82, 0.26)) + pointb = np.array((0.98, 0.26)) + shift = np.array([0, 0.1]) + + for idx, key in enumerate(self.keys): + title = "clim" if not idx else "" + rng = _get_range(self) + self.callbacks[key] = BumpColorbarPoints( + plotter=self.plotter, + brain=self, + name=key + ) + self.sliders[key] = self.plotter.add_slider_widget( + self.callbacks[key], + value=self._data[key], + rng=rng, title=title, + pointa=pointa + idx * shift, + pointb=pointb + idx * shift, + event_type="always", + ) + + # fscale + self.callbacks["fscale"] = UpdateColorbarScale( + plotter=self.plotter, + brain=self, + ) + self.sliders["fscale"] = self.plotter.add_slider_widget( + self.callbacks["fscale"], + value=1.0, + rng=self.default_scaling_range, title="fscale", + pointa=(0.82, 0.10), + pointb=(0.98, 0.10) + ) + self.callbacks["fscale"].slider_rep = \ + self.sliders["fscale"].GetRepresentation() + + # register colorbar slider representations + self.reps = \ + {key: self.sliders[key].GetRepresentation() for key in self.keys} + for name in ("fmin", "fmid", "fmax", "fscale"): + self.callbacks[name].reps = self.reps + + # set the slider style + self._set_slider_style() + + def _configure_playback(self): + self.plotter.add_callback(self._play, self.refresh_rate_ms) + + def _configure_point_picking(self): + if not self.show_traces: + return + from ..backends._pyvista import _update_picking_callback + # use a matplotlib canvas + self.color_cycle = _ReuseCycle(_get_color_list()) + win = self.plotter.app_window + dpi = win.windowHandle().screen().logicalDotsPerInch() + ratio = (1 - self.interactor_fraction) / self.interactor_fraction + w = self.interactor.geometry().width() + h = self.interactor.geometry().height() / ratio + # Get the fractional components for the brain and mpl + self.mpl_canvas = MplCanvas(self, w / dpi, h / dpi, dpi) + xlim = [np.min(self._data['time']), + np.max(self._data['time'])] + with warnings.catch_warnings(): + warnings.filterwarnings("ignore", category=UserWarning) + self.mpl_canvas.axes.set(xlim=xlim) + if not self.separate_canvas: + from PyQt5.QtWidgets import QSplitter + from PyQt5.QtCore import Qt + canvas = self.mpl_canvas.canvas + vlayout = self.plotter.frame.layout() + vlayout.removeWidget(self.interactor) + self.splitter = splitter = QSplitter( + orientation=Qt.Vertical, parent=self.plotter.frame) + vlayout.addWidget(splitter) + splitter.addWidget(self.interactor) + splitter.addWidget(canvas) + self.mpl_canvas.set_color( + bg_color=self._bg_color, + fg_color=self._fg_color, + ) + self.mpl_canvas.show() + + # get data for each hemi + for idx, hemi in enumerate(['vol', 'lh', 'rh']): + hemi_data = self._data.get(hemi) + if hemi_data is not None: + act_data = hemi_data['array'] + if act_data.ndim == 3: + act_data = np.linalg.norm(act_data, axis=1) + smooth_mat = hemi_data.get('smooth_mat') + vertices = hemi_data['vertices'] + if hemi == 'vol': + assert smooth_mat is None + smooth_mat = sparse.csr_matrix( + (np.ones(len(vertices)), + (vertices, np.arange(len(vertices))))) + self.act_data_smooth[hemi] = (act_data, smooth_mat) + + # plot the GFP + y = np.concatenate(list(v[0] for v in self.act_data_smooth.values() + if v[0] is not None)) + y = np.linalg.norm(y, axis=0) / np.sqrt(len(y)) + self.mpl_canvas.axes.plot( + self._data['time'], y, + lw=3, label='GFP', zorder=3, color=self._fg_color, + alpha=0.5, ls=':') + + # now plot the time line + self.plot_time_line() + + # then the picked points + for idx, hemi in enumerate(['lh', 'rh', 'vol']): + act_data = self.act_data_smooth.get(hemi, [None])[0] + if act_data is None: + continue + hemi_data = self._data[hemi] + vertices = hemi_data['vertices'] + + # simulate a picked renderer + if self._hemi in ('both', 'rh') or hemi == 'vol': + idx = 0 + self.picked_renderer = self.plotter.renderers[idx] + + # initialize the default point + if self._data['initial_time'] is not None: + # pick at that time + use_data = act_data[ + :, [np.round(self._data['time_idx']).astype(int)]] + else: + use_data = act_data + ind = np.unravel_index(np.argmax(np.abs(use_data), axis=None), + use_data.shape) + if hemi == 'vol': + mesh = hemi_data['grid'] + else: + mesh = hemi_data['mesh'] + vertex_id = vertices[ind[0]] + self.add_point(hemi, mesh, vertex_id) + + _update_picking_callback( + self.plotter, + self._on_mouse_move, + self._on_button_press, + self._on_button_release, + self._on_pick + ) + + def _load_icons(self): + from PyQt5.QtGui import QIcon + from ..backends._pyvista import _init_resources + _init_resources() + self.icons["help"] = QIcon(":/help.svg") + self.icons["play"] = QIcon(":/play.svg") + self.icons["pause"] = QIcon(":/pause.svg") + self.icons["reset"] = QIcon(":/reset.svg") + self.icons["scale"] = QIcon(":/scale.svg") + self.icons["clear"] = QIcon(":/clear.svg") + self.icons["movie"] = QIcon(":/movie.svg") + self.icons["restore"] = QIcon(":/restore.svg") + self.icons["screenshot"] = QIcon(":/screenshot.svg") + self.icons["visibility_on"] = QIcon(":/visibility_on.svg") + self.icons["visibility_off"] = QIcon(":/visibility_off.svg") + + def _configure_tool_bar(self): + self.actions["screenshot"] = self.tool_bar.addAction( + self.icons["screenshot"], + "Take a screenshot", + self.plotter._qt_screenshot + ) + self.actions["movie"] = self.tool_bar.addAction( + self.icons["movie"], + "Save movie...", + self.save_movie + ) + self.actions["visibility"] = self.tool_bar.addAction( + self.icons["visibility_on"], + "Toggle Visibility", + self.toggle_interface + ) + self.actions["play"] = self.tool_bar.addAction( + self.icons["play"], + "Play/Pause", + self.toggle_playback + ) + self.actions["reset"] = self.tool_bar.addAction( + self.icons["reset"], + "Reset", + self.reset + ) + self.actions["scale"] = self.tool_bar.addAction( + self.icons["scale"], + "Auto-Scale", + self.apply_auto_scaling + ) + self.actions["restore"] = self.tool_bar.addAction( + self.icons["restore"], + "Restore scaling", + self.restore_user_scaling + ) + self.actions["clear"] = self.tool_bar.addAction( + self.icons["clear"], + "Clear traces", + self.clear_points + ) + self.actions["help"] = self.tool_bar.addAction( + self.icons["help"], + "Help", + self.help + ) + + self.actions["movie"].setShortcut("ctrl+shift+s") + self.actions["visibility"].setShortcut("i") + self.actions["play"].setShortcut(" ") + self.actions["scale"].setShortcut("s") + self.actions["restore"].setShortcut("r") + self.actions["clear"].setShortcut("c") + self.actions["help"].setShortcut("?") + + def _configure_menu(self): + # remove default picking menu + to_remove = list() + for action in self.main_menu.actions(): + if action.text() == "Tools": + to_remove.append(action) + for action in to_remove: + self.main_menu.removeAction(action) + + # add help menu + menu = self.main_menu.addMenu('Help') + menu.addAction('Show MNE key bindings\t?', self.help) + + def _configure_status_bar(self): + from PyQt5.QtWidgets import QLabel, QProgressBar + self.status_msg = QLabel(self.default_status_bar_msg) + self.status_progress = QProgressBar() + self.status_bar.layout().addWidget(self.status_msg, 1) + self.status_bar.layout().addWidget(self.status_progress, 0) + self.status_progress.hide() + + def _on_mouse_move(self, vtk_picker, event): + if self._mouse_no_mvt: + self._mouse_no_mvt -= 1 + + def _on_button_press(self, vtk_picker, event): + self._mouse_no_mvt = 2 + + def _on_button_release(self, vtk_picker, event): + if self._mouse_no_mvt > 0: + x, y = vtk_picker.GetEventPosition() + # programmatically detect the picked renderer + self.picked_renderer = self.plotter.iren.FindPokedRenderer(x, y) + # trigger the pick + self.plotter.picker.Pick(x, y, 0, self.picked_renderer) + self._mouse_no_mvt = 0 + + def _on_pick(self, vtk_picker, event): + # vtk_picker is a vtkCellPicker + cell_id = vtk_picker.GetCellId() + mesh = vtk_picker.GetDataSet() + + if mesh is None or cell_id == -1 or not self._mouse_no_mvt: + return # don't pick + + # 1) Check to see if there are any spheres along the ray + if len(self._spheres): + collection = vtk_picker.GetProp3Ds() + found_sphere = None + for ii in range(collection.GetNumberOfItems()): + actor = collection.GetItemAsObject(ii) + for sphere in self._spheres: + if any(a is actor for a in sphere._actors): + found_sphere = sphere + break + if found_sphere is not None: + break + if found_sphere is not None: + assert found_sphere._is_point + mesh = found_sphere + + # 2) Remove sphere if it's what we have + if hasattr(mesh, "_is_point"): + self.remove_point(mesh) + return + + # 3) Otherwise, pick the objects in the scene + try: + hemi = mesh._hemi + except AttributeError: # volume + hemi = 'vol' + else: + assert hemi in ('lh', 'rh') + if self.act_data_smooth[hemi][0] is None: # no data to add for hemi + return + pos = np.array(vtk_picker.GetPickPosition()) + if hemi == 'vol': + # VTK will give us the point closest to the viewer in the vol. + # We want to pick the point with the maximum value along the + # camera-to-click array, which fortunately we can get "just" + # by inspecting the points that are sufficiently close to the + # ray. + grid = mesh = self._data[hemi]['grid'] + vertices = self._data[hemi]['vertices'] + coords = self._data[hemi]['grid_coords'][vertices] + scalars = grid.cell_arrays['values'][vertices] + spacing = np.array(grid.GetSpacing()) + max_dist = np.linalg.norm(spacing) / 2. + origin = vtk_picker.GetRenderer().GetActiveCamera().GetPosition() + ori = pos - origin + ori /= np.linalg.norm(ori) + # the magic formula: distance from a ray to a given point + dists = np.linalg.norm(np.cross(ori, coords - pos), axis=1) + assert dists.shape == (len(coords),) + mask = dists <= max_dist + idx = np.where(mask)[0] + if len(idx) == 0: + return # weird point on edge of volume? + # useful for debugging the ray by mapping it into the volume: + # dists = dists - dists.min() + # dists = (1. - dists / dists.max()) * self._cmap_range[1] + # grid.cell_arrays['values'][vertices] = dists * mask + idx = idx[np.argmax(np.abs(scalars[idx]))] + vertex_id = vertices[idx] + # Naive way: convert pos directly to idx; i.e., apply mri_src_t + # shape = self._data[hemi]['grid_shape'] + # taking into account the cell vs point difference (spacing/2) + # shift = np.array(grid.GetOrigin()) + spacing / 2. + # ijk = np.round((pos - shift) / spacing).astype(int) + # vertex_id = np.ravel_multi_index(ijk, shape, order='F') + else: + vtk_cell = mesh.GetCell(cell_id) + cell = [vtk_cell.GetPointId(point_id) for point_id + in range(vtk_cell.GetNumberOfPoints())] + vertices = mesh.points[cell] + idx = np.argmin(abs(vertices - pos), axis=0) + vertex_id = cell[idx[0]] + + if vertex_id not in self.picked_points[hemi]: + self.add_point(hemi, mesh, vertex_id) + + def add_point(self, hemi, mesh, vertex_id): + """Pick a vertex on the brain. + + hemi : str + The hemisphere id of the vertex. + mesh : vtkPolyData + The mesh where picking is expected. + vertex_id : int + The vertex identifier in the mesh. + + Returns + ------- + The glyph created for the picked point. + """ + # skip if the wrong hemi is selected + if self.act_data_smooth[hemi][0] is None: + return + from ..backends._pyvista import _sphere + color = next(self.color_cycle) + line = self.plot_time_course(hemi, vertex_id, color) + if hemi == 'vol': + ijk = np.unravel_index( + vertex_id, np.array(mesh.GetDimensions()) - 1, order='F') + # should just be GetCentroid(center), but apparently it's VTK9+: + # center = np.empty(3) + # voxel.GetCentroid(center) + voxel = mesh.GetCell(*ijk) + pts = voxel.GetPoints() + n_pts = pts.GetNumberOfPoints() + center = np.empty((n_pts, 3)) + for ii in range(pts.GetNumberOfPoints()): + pts.GetPoint(ii, center[ii]) + center = np.mean(center, axis=0) + else: + center = mesh.GetPoints().GetPoint(vertex_id) + del mesh + + # from the picked renderer to the subplot coords + rindex = self.plotter.renderers.index(self.picked_renderer) + row, col = self.plotter.index_to_loc(rindex) + + actors = list() + spheres = list() + for ri, ci, _ in self._iter_views(hemi): + self.plotter.subplot(ri, ci) + # Using _sphere() instead of renderer.sphere() for 2 reasons: + # 1) renderer.sphere() fails on Windows in a scenario where a lot + # of picking requests are done in a short span of time (could be + # mitigated with synchronization/delay?) + # 2) the glyph filter is used in renderer.sphere() but only one + # sphere is required in this function. + actor, sphere = _sphere( + plotter=self.plotter, + center=np.array(center), + color=color, + radius=4.0, + ) + actors.append(actor) + spheres.append(sphere) + + # add metadata for picking + for sphere in spheres: + sphere._is_point = True + sphere._hemi = hemi + sphere._line = line + sphere._actors = actors + sphere._color = color + sphere._vertex_id = vertex_id + + self.picked_points[hemi].append(vertex_id) + self._spheres.extend(spheres) + self.pick_table[vertex_id] = spheres + return sphere + + def remove_point(self, mesh): + """Remove the picked point from its glyph. + + mesh : vtkPolyData + The mesh associated to the point to remove. + """ + vertex_id = mesh._vertex_id + if vertex_id not in self.pick_table: + return + + hemi = mesh._hemi + color = mesh._color + spheres = self.pick_table[vertex_id] + spheres[0]._line.remove() + self.mpl_canvas.update_plot() + self.picked_points[hemi].remove(vertex_id) + + with warnings.catch_warnings(record=True): + # We intentionally ignore these in case we have traversed the + # entire color cycle + warnings.simplefilter('ignore') + self.color_cycle.restore(color) + for sphere in spheres: + # remove all actors + self.plotter.remove_actor(sphere._actors) + sphere._actors = None + self._spheres.pop(self._spheres.index(sphere)) + self.pick_table.pop(vertex_id) + + def clear_points(self): + """Clear the picked points.""" + for sphere in list(self._spheres): # will remove itself, so copy + self.remove_point(sphere) + assert sum(len(v) for v in self.picked_points.values()) == 0 + assert len(self.pick_table) == 0 + assert len(self._spheres) == 0 + + def plot_time_course(self, hemi, vertex_id, color): + """Plot the vertex time course. + + hemi : str + The hemisphere id of the vertex. + vertex_id : int + The vertex identifier in the mesh. + color : matplotlib color + The color of the time course. + """ + if self.mpl_canvas is None: + return + time = self._data['time'].copy() # avoid circular ref + if hemi == 'vol': + hemi_str = 'V' + xfm = _read_talxfm( + self._subject_id, self._subjects_dir) + if self._units == 'm': + xfm['trans'][:3, 3] /= 1000. + ijk = np.unravel_index( + vertex_id, self._data[hemi]['grid_shape'], order='F') + src_mri_t = self._data[hemi]['grid_src_mri_t'] + mni = apply_trans(np.dot(xfm['trans'], src_mri_t), ijk) + else: + hemi_str = 'L' if hemi == 'lh' else 'R' + mni = vertex_to_mni( + vertices=vertex_id, + hemis=0 if hemi == 'lh' else 1, + subject=self._subject_id, + subjects_dir=self._subjects_dir + ) + label = "{}:{} MNI: {}".format( + hemi_str, str(vertex_id).ljust(6), + ', '.join('%5.1f' % m for m in mni)) + act_data, smooth = self.act_data_smooth[hemi] + if smooth is not None: + act_data = smooth[vertex_id].dot(act_data)[0] + else: + act_data = act_data[vertex_id].copy() + line = self.mpl_canvas.plot( + time, + act_data, + label=label, + lw=1., + color=color, + zorder=4, + ) + return line + + def plot_time_line(self): + """Add the time line to the MPL widget.""" + if self.mpl_canvas is None: + return + if isinstance(self.show_traces, bool) and self.show_traces: + # add time information + current_time = self._current_time + if not hasattr(self, "time_line"): + self.time_line = self.mpl_canvas.plot_time_line( + x=current_time, + label='time', + color=self._fg_color, + lw=1, + ) + self.time_line.set_xdata(current_time) + self.mpl_canvas.update_plot() + + def help(self): + """Display the help window.""" + pairs = [ + ('?', 'Display help window'), + ('i', 'Toggle interface'), + ('s', 'Apply auto-scaling'), + ('r', 'Restore original clim'), + ('c', 'Clear all traces'), + ('Space', 'Start/Pause playback'), + ] + text1, text2 = zip(*pairs) + text1 = '\n'.join(text1) + text2 = '\n'.join(text2) + _show_help( + col1=text1, + col2=text2, + width=5, + height=2, + ) + + def _clear_callbacks(self): + for callback in self.callbacks.values(): + if callback is not None: + if hasattr(callback, "plotter"): + callback.plotter = None + if hasattr(callback, "brain"): + callback.brain = None + if hasattr(callback, "slider_rep"): + callback.slider_rep = None + self.callbacks.clear() + @property def interaction(self): """The interaction style.""" @@ -1188,10 +2180,10 @@ def screenshot(self, mode='rgb', time_viewer=False): Image pixel values. """ img = self._renderer.screenshot(mode) - if time_viewer and getattr(self, 'time_viewer', None) is not None and \ - self.time_viewer.show_traces and \ - not self.time_viewer.separate_canvas: - canvas = self.time_viewer.mpl_canvas.fig.canvas + if time_viewer and self.time_viewer and \ + self.show_traces and \ + not self.separate_canvas: + canvas = self.mpl_canvas.fig.canvas canvas.draw_idle() # In theory, one of these should work: # @@ -1522,6 +2514,23 @@ def views(self): def hemis(self): return self._hemis + def _save_movie(self, filename, time_dilation=4., tmin=None, tmax=None, + framerate=24, interpolation=None, codec=None, + bitrate=None, callback=None, time_viewer=False, **kwargs): + import imageio + images = self._make_movie_frames( + time_dilation, tmin, tmax, framerate, interpolation, callback, + time_viewer) + # find imageio FFMPEG parameters + if 'fps' not in kwargs: + kwargs['fps'] = framerate + if codec is not None: + kwargs['codec'] = codec + if bitrate is not None: + kwargs['bitrate'] = bitrate + + imageio.mimwrite(filename, images) + @fill_doc def save_movie(self, filename, time_dilation=4., tmin=None, tmax=None, framerate=24, interpolation=None, codec=None, @@ -1570,19 +2579,83 @@ def save_movie(self, filename, time_dilation=4., tmin=None, tmax=None, **kwargs : dict Specify additional options for :mod:`imageio`. """ - import imageio - images = self._make_movie_frames( - time_dilation, tmin, tmax, framerate, interpolation, callback, - time_viewer) - # find imageio FFMPEG parameters - if 'fps' not in kwargs: - kwargs['fps'] = framerate - if codec is not None: - kwargs['codec'] = codec - if bitrate is not None: - kwargs['bitrate'] = bitrate + if self.time_viewer: + try: + from pyvista.plotting.qt_plotting import FileDialog + except ImportError: + from pyvistaqt.plotting import FileDialog + + if filename is None: + self.status_msg.setText("Choose movie path ...") + self.status_msg.show() + self.status_progress.setValue(0) + + def _post_setup(unused): + del unused + self.status_msg.hide() + self.status_progress.hide() + + dialog = FileDialog( + self.plotter.app_window, + callback=partial(self._save_movie, **kwargs) + ) + dialog.setDirectory(os.getcwd()) + dialog.finished.connect(_post_setup) + return dialog + else: + from PyQt5.QtCore import Qt + from PyQt5.QtGui import QCursor + + def frame_callback(frame, n_frames): + if frame == n_frames: + # On the ImageIO step + self.status_msg.setText( + "Saving with ImageIO: %s" + % filename + ) + self.status_msg.show() + self.status_progress.hide() + self.status_bar.layout().update() + else: + self.status_msg.setText( + "Rendering images (frame %d / %d) ..." + % (frame + 1, n_frames) + ) + self.status_msg.show() + self.status_progress.show() + self.status_progress.setRange(0, n_frames - 1) + self.status_progress.setValue(frame) + self.status_progress.update() + self.status_progress.repaint() + self.status_msg.update() + self.status_msg.parent().update() + self.status_msg.repaint() + + # temporarily hide interface + default_visibility = self.visibility + self.toggle_interface(value=False) + # set cursor to busy + default_cursor = self.interactor.cursor() + self.interactor.setCursor(QCursor(Qt.WaitCursor)) + + try: + self._save_movie( + filename=filename, + time_dilation=(1. / self.playback_speed), + callback=frame_callback, + **kwargs + ) + except (Exception, KeyboardInterrupt): + warn('Movie saving aborted:\n' + traceback.format_exc()) - imageio.mimwrite(filename, images) + # restore visibility + self.toggle_interface(value=default_visibility) + # restore cursor + self.interactor.setCursor(default_cursor) + else: + self._save_movie(filename, time_dilation, tmin, tmax, + framerate, interpolation, codec, + bitrate, callback, time_viewer, **kwargs) def _make_movie_frames(self, time_dilation, tmin, tmax, framerate, interpolation, callback, time_viewer): @@ -1622,14 +2695,14 @@ def _make_movie_frames(self, time_dilation, tmin, tmax, framerate, try: images = [ self.screenshot(time_viewer=time_viewer) - for _ in self._iter_time(time_idx, callback, time_viewer)] + for _ in self._iter_time(time_idx, callback)] finally: self.set_time_interpolation(old_mode) if callback is not None: callback(frame=len(time_idx), n_frames=len(time_idx)) return images - def _iter_time(self, time_idx, callback, time_viewer=False): + def _iter_time(self, time_idx, callback): """Iterate through time points, then reset to current time. Parameters @@ -1638,8 +2711,6 @@ def _iter_time(self, time_idx, callback, time_viewer=False): Time point indexes through which to iterate. callback : callable | None Callback to call before yielding each frame. - time_viewer : bool - If True, route through self.time_viewer. Yields ------ @@ -1650,8 +2721,8 @@ def _iter_time(self, time_idx, callback, time_viewer=False): ----- Used by movie and image sequence saving functions. """ - if hasattr(self, 'time_viewer'): - func = partial(self.time_viewer.callbacks["time"], + if self.time_viewer: + func = partial(self.callbacks["time"], update_widget=True) else: func = self.set_time_point @@ -1777,7 +2848,7 @@ def get_picked_points(self): The vertices picked by the time viewer. """ if hasattr(self, "time_viewer"): - return self.time_viewer.picked_points + return self.picked_points def __hash__(self): """Hash the object.""" @@ -1817,3 +2888,12 @@ def _update_limits(fmin, fmid, fmax, center, array): % (fmid, fmax)) return fmin, fmid, fmax + + +def _get_range(brain): + val = np.abs(np.concatenate(list(brain._current_act_data.values()))) + return [np.min(val), np.max(val)] + + +def _normalize(point, shape): + return (point[0] / shape[1], point[1] / shape[0]) diff --git a/mne/viz/_brain/_linkviewer.py b/mne/viz/_brain/_linkviewer.py index b77cd696453..00eab9929ea 100644 --- a/mne/viz/_brain/_linkviewer.py +++ b/mne/viz/_brain/_linkviewer.py @@ -8,12 +8,11 @@ class _LinkViewer(object): - """Class to link multiple _TimeViewer objects.""" + """Class to link multiple _Brain objects.""" def __init__(self, brains, time=True, camera=False, colorbar=True, picking=False): - self.brains = brains - self.time_viewers = [brain.time_viewer for brain in brains] + self.time_viewers = brains # check time infos times = [brain._times for brain in brains] @@ -83,16 +82,16 @@ def _func_remove(*args, **kwargs): # link the initial points leader = self.time_viewers[0] # select a time_viewer as leader for hemi in initial_points.keys(): - if hemi in time_viewer.brain._hemi_meshes: - mesh = time_viewer.brain._hemi_meshes[hemi] + if hemi in time_viewer._hemi_meshes: + mesh = time_viewer._hemi_meshes[hemi] for vertex_id in initial_points[hemi]: leader.add_point(hemi, mesh, vertex_id) if colorbar: leader = self.time_viewers[0] # select a time_viewer as leader - fmin = leader.brain._data["fmin"] - fmid = leader.brain._data["fmid"] - fmax = leader.brain._data["fmax"] + fmin = leader._data["fmin"] + fmid = leader._data["fmid"] + fmax = leader._data["fmax"] for time_viewer in self.time_viewers: time_viewer.callbacks["fmin"](fmin) time_viewer.callbacks["fmid"](fmid) diff --git a/mne/viz/_brain/_notebook.py b/mne/viz/_brain/_notebook.py index 7773df581fe..801ba240c07 100644 --- a/mne/viz/_brain/_notebook.py +++ b/mne/viz/_brain/_notebook.py @@ -7,9 +7,8 @@ class _NotebookInteractor(_PyVistaNotebookInteractor): - def __init__(self, time_viewer): - self.time_viewer = time_viewer - self.brain = self.time_viewer.brain + def __init__(self, brain): + self.brain = brain super().__init__(self.brain._renderer) def configure_controllers(self): @@ -19,13 +18,13 @@ def configure_controllers(self): # orientation self.controllers["orientation"] = interactive( self.set_orientation, - orientation=self.time_viewer.orientation, + orientation=self.brain.orientation, ) # smoothing self.sliders["smoothing"] = IntSlider( value=self.brain._data['smoothing_steps'], - min=self.time_viewer.default_smoothing_range[0], - max=self.time_viewer.default_smoothing_range[1], + min=self.brain.default_smoothing_range[0], + max=self.brain.default_smoothing_range[1], continuous_update=False ) self.controllers["smoothing"] = VBox([ diff --git a/mne/viz/_brain/_timeviewer.py b/mne/viz/_brain/_timeviewer.py deleted file mode 100644 index 6ba8d09f05c..00000000000 --- a/mne/viz/_brain/_timeviewer.py +++ /dev/null @@ -1,1049 +0,0 @@ -# Authors: Alexandre Gramfort -# Eric Larson -# Guillaume Favelier -# -# License: Simplified BSD - -import contextlib -from functools import partial -import os -import sys -import time -import traceback -import warnings - -import numpy as np -from scipy import sparse - -from ._brain import Brain -from .callback import (ShowView, IntSlider, TimeSlider, SmartSlider, - BumpColorbarPoints, UpdateColorbarScale) -from .mplcanvas import MplCanvas -from .view import _lh_views_dict - -from ..utils import _show_help, _get_color_list -from ...externals.decorator import decorator -from ...source_space import vertex_to_mni, _read_talxfm -from ...transforms import apply_trans -from ...utils import _ReuseCycle, warn, copy_doc, _validate_type - - -@decorator -def safe_event(fun, *args, **kwargs): - """Protect against PyQt5 exiting on event-handling errors.""" - try: - return fun(*args, **kwargs) - except Exception: - traceback.print_exc(file=sys.stderr) - - -class _TimeViewer(object): - """Class to interact with Brain.""" - - def __init__(self, brain, show_traces=False): - from ..backends._pyvista import _require_minimum_version - _require_minimum_version('0.24') - - # shared configuration - if hasattr(brain, 'time_viewer'): - raise RuntimeError('brain already has a TimeViewer') - self.brain = brain - self.orientation = list(_lh_views_dict.keys()) - self.default_smoothing_range = [0, 15] - - # detect notebook - if brain._notebook: - self.notebook = True - self.configure_notebook() - return - else: - self.notebook = False - - # Default configuration - self.playback = False - self.visibility = False - self.refresh_rate_ms = max(int(round(1000. / 60.)), 1) - self.default_scaling_range = [0.2, 2.0] - self.default_playback_speed_range = [0.01, 1] - self.default_playback_speed_value = 0.05 - self.default_status_bar_msg = "Press ? for help" - all_keys = ('lh', 'rh', 'vol') - self.act_data_smooth = {key: (None, None) for key in all_keys} - self.color_cycle = None - self.mpl_canvas = None - self.picked_points = {key: list() for key in all_keys} - self.pick_table = dict() - self._mouse_no_mvt = -1 - self.icons = dict() - self.actions = dict() - self.callbacks = dict() - self.sliders = dict() - self.keys = ('fmin', 'fmid', 'fmax') - self.slider_length = 0.02 - self.slider_width = 0.04 - self.slider_color = (0.43137255, 0.44313725, 0.45882353) - self.slider_tube_width = 0.04 - self.slider_tube_color = (0.69803922, 0.70196078, 0.70980392) - - # Direct access parameters: - self.brain.time_viewer = self - self.plotter = brain._renderer.plotter - self.main_menu = self.plotter.main_menu - self.window = self.plotter.app_window - self.tool_bar = self.window.addToolBar("toolbar") - self.status_bar = self.window.statusBar() - self.interactor = self.plotter.interactor - self.window.signal_close.connect(self.clean) - - # Derived parameters: - self.playback_speed = self.default_playback_speed_value - _validate_type(show_traces, (bool, str, 'numeric'), 'show_traces') - self.interactor_fraction = 0.25 - if isinstance(show_traces, str): - assert 'show_traces' == 'separate' # should be guaranteed earlier - self.show_traces = True - self.separate_canvas = True - else: - if isinstance(show_traces, bool): - self.show_traces = show_traces - else: - show_traces = float(show_traces) - if not 0 < show_traces < 1: - raise ValueError( - 'show traces, if numeric, must be between 0 and 1, ' - f'got {show_traces}') - self.show_traces = True - self.interactor_fraction = show_traces - self.separate_canvas = False - del show_traces - - self._spheres = list() - self.load_icons() - self.configure_time_label() - self.configure_sliders() - self.configure_scalar_bar() - self.configure_playback() - self.configure_point_picking() - self.configure_menu() - self.configure_tool_bar() - self.configure_status_bar() - - # show everything at the end - self.toggle_interface() - with self.ensure_minimum_sizes(): - self.brain.show() - - @contextlib.contextmanager - def ensure_minimum_sizes(self): - from ..backends._pyvista import _process_events - sz = self.brain._size - adjust_mpl = self.show_traces and not self.separate_canvas - if not adjust_mpl: - yield - else: - mpl_h = int(round((sz[1] * self.interactor_fraction) / - (1 - self.interactor_fraction))) - self.mpl_canvas.canvas.setMinimumSize(sz[0], mpl_h) - try: - yield - finally: - self.splitter.setSizes([sz[1], mpl_h]) - _process_events(self.plotter) - _process_events(self.plotter) - self.mpl_canvas.canvas.setMinimumSize(0, 0) - _process_events(self.plotter) - _process_events(self.plotter) - # sizes could change, update views - for hemi in ('lh', 'rh'): - for ri, ci, v in self.brain._iter_views(hemi): - self.brain.show_view(view=v, row=ri, col=ci) - _process_events(self.plotter) - - def toggle_interface(self, value=None): - if value is None: - self.visibility = not self.visibility - else: - self.visibility = value - - # update tool bar icon - if self.visibility: - self.actions["visibility"].setIcon(self.icons["visibility_on"]) - else: - self.actions["visibility"].setIcon(self.icons["visibility_off"]) - - # manage sliders - for slider in self.plotter.slider_widgets: - slider_rep = slider.GetRepresentation() - if self.visibility: - slider_rep.VisibilityOn() - else: - slider_rep.VisibilityOff() - - # manage time label - time_label = self.brain._data['time_label'] - # if we actually have time points, we will show the slider so - # hide the time actor - have_ts = self.brain._times is not None and len(self.brain._times) > 1 - if self.time_actor is not None: - if self.visibility and time_label is not None and not have_ts: - self.time_actor.SetInput(time_label(self.brain._current_time)) - self.time_actor.VisibilityOn() - else: - self.time_actor.VisibilityOff() - - self.plotter.update() - - def _save_movie(self, filename, **kwargs): - from PyQt5.QtCore import Qt - from PyQt5.QtGui import QCursor - - def frame_callback(frame, n_frames): - if frame == n_frames: - # On the ImageIO step - self.status_msg.setText( - "Saving with ImageIO: %s" - % filename - ) - self.status_msg.show() - self.status_progress.hide() - self.status_bar.layout().update() - else: - self.status_msg.setText( - "Rendering images (frame %d / %d) ..." - % (frame + 1, n_frames) - ) - self.status_msg.show() - self.status_progress.show() - self.status_progress.setRange(0, n_frames - 1) - self.status_progress.setValue(frame) - self.status_progress.update() - self.status_progress.repaint() - self.status_msg.update() - self.status_msg.parent().update() - self.status_msg.repaint() - - # temporarily hide interface - default_visibility = self.visibility - self.toggle_interface(value=False) - # set cursor to busy - default_cursor = self.interactor.cursor() - self.interactor.setCursor(QCursor(Qt.WaitCursor)) - - try: - self.brain.save_movie( - filename=filename, - time_dilation=(1. / self.playback_speed), - callback=frame_callback, - **kwargs - ) - except (Exception, KeyboardInterrupt): - warn('Movie saving aborted:\n' + traceback.format_exc()) - - # restore visibility - self.toggle_interface(value=default_visibility) - # restore cursor - self.interactor.setCursor(default_cursor) - - @copy_doc(Brain.save_movie) - def save_movie(self, filename=None, **kwargs): - try: - from pyvista.plotting.qt_plotting import FileDialog - except ImportError: - from pyvistaqt.plotting import FileDialog - - if filename is None: - self.status_msg.setText("Choose movie path ...") - self.status_msg.show() - self.status_progress.setValue(0) - - def _clean(unused): - del unused - self.status_msg.hide() - self.status_progress.hide() - - dialog = FileDialog( - self.plotter.app_window, - callback=partial(self._save_movie, **kwargs) - ) - dialog.setDirectory(os.getcwd()) - dialog.finished.connect(_clean) - return dialog - else: - self._save_movie(filename=filename, **kwargs) - return - - def apply_auto_scaling(self): - self.brain._update_auto_scaling() - for key in ('fmin', 'fmid', 'fmax'): - self.reps[key].SetValue(self.brain._data[key]) - self.plotter.update() - - def restore_user_scaling(self): - self.brain._update_auto_scaling(restore=True) - for key in ('fmin', 'fmid', 'fmax'): - self.reps[key].SetValue(self.brain._data[key]) - self.plotter.update() - - def toggle_playback(self, value=None): - if value is None: - self.playback = not self.playback - else: - self.playback = value - - # update tool bar icon - if self.playback: - self.actions["play"].setIcon(self.icons["pause"]) - else: - self.actions["play"].setIcon(self.icons["play"]) - - if self.playback: - time_data = self.brain._data['time'] - max_time = np.max(time_data) - if self.brain._current_time == max_time: # start over - self.brain.set_time_point(0) # first index - self._last_tick = time.time() - - def reset(self): - self.brain.reset_view() - max_time = len(self.brain._data['time']) - 1 - if max_time > 0: - self.callbacks["time"]( - self.brain._data["initial_time_idx"], - update_widget=True, - ) - self.plotter.update() - - def set_playback_speed(self, speed): - self.playback_speed = speed - - @safe_event - def play(self): - if self.playback: - try: - self._advance() - except Exception: - self.toggle_playback(value=False) - raise - - def _advance(self): - this_time = time.time() - delta = this_time - self._last_tick - self._last_tick = time.time() - time_data = self.brain._data['time'] - times = np.arange(self.brain._n_times) - time_shift = delta * self.playback_speed - max_time = np.max(time_data) - time_point = min(self.brain._current_time + time_shift, max_time) - # always use linear here -- this does not determine the data - # interpolation mode, it just finds where we are (in time) in - # terms of the time indices - idx = np.interp(time_point, time_data, times) - self.callbacks["time"](idx, update_widget=True) - if time_point == max_time: - self.toggle_playback(value=False) - - def set_slider_style(self): - for slider in self.sliders.values(): - if slider is not None: - slider_rep = slider.GetRepresentation() - slider_rep.SetSliderLength(self.slider_length) - slider_rep.SetSliderWidth(self.slider_width) - slider_rep.SetTubeWidth(self.slider_tube_width) - slider_rep.GetSliderProperty().SetColor(self.slider_color) - slider_rep.GetTubeProperty().SetColor(self.slider_tube_color) - slider_rep.GetLabelProperty().SetShadow(False) - slider_rep.GetLabelProperty().SetBold(True) - slider_rep.GetLabelProperty().SetColor(self.brain._fg_color) - slider_rep.GetTitleProperty().ShallowCopy( - slider_rep.GetLabelProperty() - ) - slider_rep.GetCapProperty().SetOpacity(0) - - def configure_notebook(self): - from ._notebook import _NotebookInteractor - self.brain._renderer.figure.display = _NotebookInteractor(self) - - def configure_time_label(self): - self.time_actor = self.brain._data.get('time_actor') - if self.time_actor is not None: - self.time_actor.SetPosition(0.5, 0.03) - self.time_actor.GetTextProperty().SetJustificationToCentered() - self.time_actor.GetTextProperty().BoldOn() - self.time_actor.VisibilityOff() - - def configure_scalar_bar(self): - if self.brain._colorbar_added: - scalar_bar = self.plotter.scalar_bar - scalar_bar.SetOrientationToVertical() - scalar_bar.SetHeight(0.6) - scalar_bar.SetWidth(0.05) - scalar_bar.SetPosition(0.02, 0.2) - - def configure_sliders(self): - # Orientation slider - # Use 'lh' as a reference for orientation for 'both' - if self.brain._hemi == 'both': - hemis_ref = ['lh'] - else: - hemis_ref = self.brain._hemis - for hemi in hemis_ref: - for ri, ci, view in self.brain._iter_views(hemi): - orientation_name = f"orientation_{hemi}_{ri}_{ci}" - self.plotter.subplot(ri, ci) - if view == 'flat': - self.callbacks[orientation_name] = None - continue - self.callbacks[orientation_name] = ShowView( - plotter=self.plotter, - brain=self.brain, - orientation=self.orientation, - hemi=hemi, - row=ri, - col=ci, - ) - self.sliders[orientation_name] = \ - self.plotter.add_text_slider_widget( - self.callbacks[orientation_name], - value=0, - data=self.orientation, - pointa=(0.82, 0.74), - pointb=(0.98, 0.74), - event_type='always' - ) - orientation_rep = \ - self.sliders[orientation_name].GetRepresentation() - orientation_rep.ShowSliderLabelOff() - self.callbacks[orientation_name].slider_rep = orientation_rep - self.callbacks[orientation_name](view, update_widget=True) - - # Put other sliders on the bottom right view - ri, ci = np.array(self.brain._subplot_shape) - 1 - self.plotter.subplot(ri, ci) - - # Smoothing slider - self.callbacks["smoothing"] = IntSlider( - plotter=self.plotter, - callback=self.brain.set_data_smoothing, - first_call=False, - ) - self.sliders["smoothing"] = self.plotter.add_slider_widget( - self.callbacks["smoothing"], - value=self.brain._data['smoothing_steps'], - rng=self.default_smoothing_range, title="smoothing", - pointa=(0.82, 0.90), - pointb=(0.98, 0.90) - ) - self.callbacks["smoothing"].slider_rep = \ - self.sliders["smoothing"].GetRepresentation() - - # Time slider - max_time = len(self.brain._data['time']) - 1 - # VTK on macOS bombs if we create these then hide them, so don't - # even create them - if max_time < 1: - self.callbacks["time"] = None - self.sliders["time"] = None - else: - self.callbacks["time"] = TimeSlider( - plotter=self.plotter, - brain=self.brain, - first_call=False, - callback=self.plot_time_line, - ) - self.sliders["time"] = self.plotter.add_slider_widget( - self.callbacks["time"], - value=self.brain._data['time_idx'], - rng=[0, max_time], - pointa=(0.23, 0.1), - pointb=(0.77, 0.1), - event_type='always' - ) - self.callbacks["time"].slider_rep = \ - self.sliders["time"].GetRepresentation() - # configure properties of the time slider - self.sliders["time"].GetRepresentation().SetLabelFormat( - 'idx=%0.1f') - - current_time = self.brain._current_time - assert current_time is not None # should never be the case, float - time_label = self.brain._data['time_label'] - if callable(time_label): - current_time = time_label(current_time) - else: - current_time = time_label - if self.sliders["time"] is not None: - self.sliders["time"].GetRepresentation().SetTitleText(current_time) - if self.time_actor is not None: - self.time_actor.SetInput(current_time) - del current_time - - # Playback speed slider - if self.sliders["time"] is None: - self.callbacks["playback_speed"] = None - self.sliders["playback_speed"] = None - else: - self.callbacks["playback_speed"] = SmartSlider( - plotter=self.plotter, - callback=self.set_playback_speed, - ) - self.sliders["playback_speed"] = self.plotter.add_slider_widget( - self.callbacks["playback_speed"], - value=self.default_playback_speed_value, - rng=self.default_playback_speed_range, title="speed", - pointa=(0.02, 0.1), - pointb=(0.18, 0.1), - event_type='always' - ) - self.callbacks["playback_speed"].slider_rep = \ - self.sliders["playback_speed"].GetRepresentation() - - # Colormap slider - pointa = np.array((0.82, 0.26)) - pointb = np.array((0.98, 0.26)) - shift = np.array([0, 0.1]) - - for idx, key in enumerate(self.keys): - title = "clim" if not idx else "" - rng = _get_range(self.brain) - self.callbacks[key] = BumpColorbarPoints( - plotter=self.plotter, - brain=self.brain, - name=key - ) - self.sliders[key] = self.plotter.add_slider_widget( - self.callbacks[key], - value=self.brain._data[key], - rng=rng, title=title, - pointa=pointa + idx * shift, - pointb=pointb + idx * shift, - event_type="always", - ) - - # fscale - self.callbacks["fscale"] = UpdateColorbarScale( - plotter=self.plotter, - brain=self.brain, - ) - self.sliders["fscale"] = self.plotter.add_slider_widget( - self.callbacks["fscale"], - value=1.0, - rng=self.default_scaling_range, title="fscale", - pointa=(0.82, 0.10), - pointb=(0.98, 0.10) - ) - self.callbacks["fscale"].slider_rep = \ - self.sliders["fscale"].GetRepresentation() - - # register colorbar slider representations - self.reps = \ - {key: self.sliders[key].GetRepresentation() for key in self.keys} - for name in ("fmin", "fmid", "fmax", "fscale"): - self.callbacks[name].reps = self.reps - - # set the slider style - self.set_slider_style() - - def configure_playback(self): - self.plotter.add_callback(self.play, self.refresh_rate_ms) - - def configure_point_picking(self): - if not self.show_traces: - return - from ..backends._pyvista import _update_picking_callback - # use a matplotlib canvas - self.color_cycle = _ReuseCycle(_get_color_list()) - win = self.plotter.app_window - dpi = win.windowHandle().screen().logicalDotsPerInch() - ratio = (1 - self.interactor_fraction) / self.interactor_fraction - w = self.interactor.geometry().width() - h = self.interactor.geometry().height() / ratio - # Get the fractional components for the brain and mpl - self.mpl_canvas = MplCanvas(self, w / dpi, h / dpi, dpi) - xlim = [np.min(self.brain._data['time']), - np.max(self.brain._data['time'])] - with warnings.catch_warnings(): - warnings.filterwarnings("ignore", category=UserWarning) - self.mpl_canvas.axes.set(xlim=xlim) - if not self.separate_canvas: - from PyQt5.QtWidgets import QSplitter - from PyQt5.QtCore import Qt - canvas = self.mpl_canvas.canvas - vlayout = self.plotter.frame.layout() - vlayout.removeWidget(self.interactor) - self.splitter = splitter = QSplitter( - orientation=Qt.Vertical, parent=self.plotter.frame) - vlayout.addWidget(splitter) - splitter.addWidget(self.interactor) - splitter.addWidget(canvas) - self.mpl_canvas.set_color( - bg_color=self.brain._bg_color, - fg_color=self.brain._fg_color, - ) - self.mpl_canvas.show() - - # get data for each hemi - for idx, hemi in enumerate(['vol', 'lh', 'rh']): - hemi_data = self.brain._data.get(hemi) - if hemi_data is not None: - act_data = hemi_data['array'] - if act_data.ndim == 3: - act_data = np.linalg.norm(act_data, axis=1) - smooth_mat = hemi_data.get('smooth_mat') - vertices = hemi_data['vertices'] - if hemi == 'vol': - assert smooth_mat is None - smooth_mat = sparse.csr_matrix( - (np.ones(len(vertices)), - (vertices, np.arange(len(vertices))))) - self.act_data_smooth[hemi] = (act_data, smooth_mat) - - # plot the GFP - y = np.concatenate(list(v[0] for v in self.act_data_smooth.values() - if v[0] is not None)) - y = np.linalg.norm(y, axis=0) / np.sqrt(len(y)) - self.mpl_canvas.axes.plot( - self.brain._data['time'], y, - lw=3, label='GFP', zorder=3, color=self.brain._fg_color, - alpha=0.5, ls=':') - - # now plot the time line - self.plot_time_line() - - # then the picked points - for idx, hemi in enumerate(['lh', 'rh', 'vol']): - act_data = self.act_data_smooth.get(hemi, [None])[0] - if act_data is None: - continue - hemi_data = self.brain._data[hemi] - vertices = hemi_data['vertices'] - - # simulate a picked renderer - if self.brain._hemi in ('both', 'rh') or hemi == 'vol': - idx = 0 - self.picked_renderer = self.plotter.renderers[idx] - - # initialize the default point - if self.brain._data['initial_time'] is not None: - # pick at that time - use_data = act_data[ - :, [np.round(self.brain._data['time_idx']).astype(int)]] - else: - use_data = act_data - ind = np.unravel_index(np.argmax(np.abs(use_data), axis=None), - use_data.shape) - if hemi == 'vol': - mesh = hemi_data['grid'] - else: - mesh = hemi_data['mesh'] - vertex_id = vertices[ind[0]] - self.add_point(hemi, mesh, vertex_id) - - _update_picking_callback( - self.plotter, - self.on_mouse_move, - self.on_button_press, - self.on_button_release, - self.on_pick - ) - - def load_icons(self): - from PyQt5.QtGui import QIcon - from ..backends._pyvista import _init_resources - _init_resources() - self.icons["help"] = QIcon(":/help.svg") - self.icons["play"] = QIcon(":/play.svg") - self.icons["pause"] = QIcon(":/pause.svg") - self.icons["reset"] = QIcon(":/reset.svg") - self.icons["scale"] = QIcon(":/scale.svg") - self.icons["clear"] = QIcon(":/clear.svg") - self.icons["movie"] = QIcon(":/movie.svg") - self.icons["restore"] = QIcon(":/restore.svg") - self.icons["screenshot"] = QIcon(":/screenshot.svg") - self.icons["visibility_on"] = QIcon(":/visibility_on.svg") - self.icons["visibility_off"] = QIcon(":/visibility_off.svg") - - def configure_tool_bar(self): - self.actions["screenshot"] = self.tool_bar.addAction( - self.icons["screenshot"], - "Take a screenshot", - self.plotter._qt_screenshot - ) - self.actions["movie"] = self.tool_bar.addAction( - self.icons["movie"], - "Save movie...", - self.save_movie - ) - self.actions["visibility"] = self.tool_bar.addAction( - self.icons["visibility_on"], - "Toggle Visibility", - self.toggle_interface - ) - self.actions["play"] = self.tool_bar.addAction( - self.icons["play"], - "Play/Pause", - self.toggle_playback - ) - self.actions["reset"] = self.tool_bar.addAction( - self.icons["reset"], - "Reset", - self.reset - ) - self.actions["scale"] = self.tool_bar.addAction( - self.icons["scale"], - "Auto-Scale", - self.apply_auto_scaling - ) - self.actions["restore"] = self.tool_bar.addAction( - self.icons["restore"], - "Restore scaling", - self.restore_user_scaling - ) - self.actions["clear"] = self.tool_bar.addAction( - self.icons["clear"], - "Clear traces", - self.clear_points - ) - self.actions["help"] = self.tool_bar.addAction( - self.icons["help"], - "Help", - self.help - ) - - self.actions["movie"].setShortcut("ctrl+shift+s") - self.actions["visibility"].setShortcut("i") - self.actions["play"].setShortcut(" ") - self.actions["scale"].setShortcut("s") - self.actions["restore"].setShortcut("r") - self.actions["clear"].setShortcut("c") - self.actions["help"].setShortcut("?") - - def configure_menu(self): - # remove default picking menu - to_remove = list() - for action in self.main_menu.actions(): - if action.text() == "Tools": - to_remove.append(action) - for action in to_remove: - self.main_menu.removeAction(action) - - # add help menu - menu = self.main_menu.addMenu('Help') - menu.addAction('Show MNE key bindings\t?', self.help) - - def configure_status_bar(self): - from PyQt5.QtWidgets import QLabel, QProgressBar - self.status_msg = QLabel(self.default_status_bar_msg) - self.status_progress = QProgressBar() - self.status_bar.layout().addWidget(self.status_msg, 1) - self.status_bar.layout().addWidget(self.status_progress, 0) - self.status_progress.hide() - - def on_mouse_move(self, vtk_picker, event): - if self._mouse_no_mvt: - self._mouse_no_mvt -= 1 - - def on_button_press(self, vtk_picker, event): - self._mouse_no_mvt = 2 - - def on_button_release(self, vtk_picker, event): - if self._mouse_no_mvt > 0: - x, y = vtk_picker.GetEventPosition() - # programmatically detect the picked renderer - self.picked_renderer = self.plotter.iren.FindPokedRenderer(x, y) - # trigger the pick - self.plotter.picker.Pick(x, y, 0, self.picked_renderer) - self._mouse_no_mvt = 0 - - def on_pick(self, vtk_picker, event): - # vtk_picker is a vtkCellPicker - cell_id = vtk_picker.GetCellId() - mesh = vtk_picker.GetDataSet() - - if mesh is None or cell_id == -1 or not self._mouse_no_mvt: - return # don't pick - - # 1) Check to see if there are any spheres along the ray - if len(self._spheres): - collection = vtk_picker.GetProp3Ds() - found_sphere = None - for ii in range(collection.GetNumberOfItems()): - actor = collection.GetItemAsObject(ii) - for sphere in self._spheres: - if any(a is actor for a in sphere._actors): - found_sphere = sphere - break - if found_sphere is not None: - break - if found_sphere is not None: - assert found_sphere._is_point - mesh = found_sphere - - # 2) Remove sphere if it's what we have - if hasattr(mesh, "_is_point"): - self.remove_point(mesh) - return - - # 3) Otherwise, pick the objects in the scene - try: - hemi = mesh._hemi - except AttributeError: # volume - hemi = 'vol' - else: - assert hemi in ('lh', 'rh') - if self.act_data_smooth[hemi][0] is None: # no data to add for hemi - return - pos = np.array(vtk_picker.GetPickPosition()) - if hemi == 'vol': - # VTK will give us the point closest to the viewer in the vol. - # We want to pick the point with the maximum value along the - # camera-to-click array, which fortunately we can get "just" - # by inspecting the points that are sufficiently close to the - # ray. - grid = mesh = self.brain._data[hemi]['grid'] - vertices = self.brain._data[hemi]['vertices'] - coords = self.brain._data[hemi]['grid_coords'][vertices] - scalars = grid.cell_arrays['values'][vertices] - spacing = np.array(grid.GetSpacing()) - max_dist = np.linalg.norm(spacing) / 2. - origin = vtk_picker.GetRenderer().GetActiveCamera().GetPosition() - ori = pos - origin - ori /= np.linalg.norm(ori) - # the magic formula: distance from a ray to a given point - dists = np.linalg.norm(np.cross(ori, coords - pos), axis=1) - assert dists.shape == (len(coords),) - mask = dists <= max_dist - idx = np.where(mask)[0] - if len(idx) == 0: - return # weird point on edge of volume? - # useful for debugging the ray by mapping it into the volume: - # dists = dists - dists.min() - # dists = (1. - dists / dists.max()) * self.brain._cmap_range[1] - # grid.cell_arrays['values'][vertices] = dists * mask - idx = idx[np.argmax(np.abs(scalars[idx]))] - vertex_id = vertices[idx] - # Naive way: convert pos directly to idx; i.e., apply mri_src_t - # shape = self.brain._data[hemi]['grid_shape'] - # taking into account the cell vs point difference (spacing/2) - # shift = np.array(grid.GetOrigin()) + spacing / 2. - # ijk = np.round((pos - shift) / spacing).astype(int) - # vertex_id = np.ravel_multi_index(ijk, shape, order='F') - else: - vtk_cell = mesh.GetCell(cell_id) - cell = [vtk_cell.GetPointId(point_id) for point_id - in range(vtk_cell.GetNumberOfPoints())] - vertices = mesh.points[cell] - idx = np.argmin(abs(vertices - pos), axis=0) - vertex_id = cell[idx[0]] - - if vertex_id not in self.picked_points[hemi]: - self.add_point(hemi, mesh, vertex_id) - - def add_point(self, hemi, mesh, vertex_id): - # skip if the wrong hemi is selected - if self.act_data_smooth[hemi][0] is None: - return - from ..backends._pyvista import _sphere - color = next(self.color_cycle) - line = self.plot_time_course(hemi, vertex_id, color) - if hemi == 'vol': - ijk = np.unravel_index( - vertex_id, np.array(mesh.GetDimensions()) - 1, order='F') - # should just be GetCentroid(center), but apparently it's VTK9+: - # center = np.empty(3) - # voxel.GetCentroid(center) - voxel = mesh.GetCell(*ijk) - pts = voxel.GetPoints() - n_pts = pts.GetNumberOfPoints() - center = np.empty((n_pts, 3)) - for ii in range(pts.GetNumberOfPoints()): - pts.GetPoint(ii, center[ii]) - center = np.mean(center, axis=0) - else: - center = mesh.GetPoints().GetPoint(vertex_id) - del mesh - - # from the picked renderer to the subplot coords - rindex = self.plotter.renderers.index(self.picked_renderer) - row, col = self.plotter.index_to_loc(rindex) - - actors = list() - spheres = list() - for ri, ci, _ in self.brain._iter_views(hemi): - self.plotter.subplot(ri, ci) - # Using _sphere() instead of renderer.sphere() for 2 reasons: - # 1) renderer.sphere() fails on Windows in a scenario where a lot - # of picking requests are done in a short span of time (could be - # mitigated with synchronization/delay?) - # 2) the glyph filter is used in renderer.sphere() but only one - # sphere is required in this function. - actor, sphere = _sphere( - plotter=self.plotter, - center=np.array(center), - color=color, - radius=4.0, - ) - actors.append(actor) - spheres.append(sphere) - - # add metadata for picking - for sphere in spheres: - sphere._is_point = True - sphere._hemi = hemi - sphere._line = line - sphere._actors = actors - sphere._color = color - sphere._vertex_id = vertex_id - - self.picked_points[hemi].append(vertex_id) - self._spheres.extend(spheres) - self.pick_table[vertex_id] = spheres - - def remove_point(self, mesh): - vertex_id = mesh._vertex_id - if vertex_id not in self.pick_table: - return - - hemi = mesh._hemi - color = mesh._color - spheres = self.pick_table[vertex_id] - spheres[0]._line.remove() - self.mpl_canvas.update_plot() - self.picked_points[hemi].remove(vertex_id) - - with warnings.catch_warnings(record=True): - # We intentionally ignore these in case we have traversed the - # entire color cycle - warnings.simplefilter('ignore') - self.color_cycle.restore(color) - for sphere in spheres: - # remove all actors - self.plotter.remove_actor(sphere._actors) - sphere._actors = None - self._spheres.pop(self._spheres.index(sphere)) - self.pick_table.pop(vertex_id) - - def clear_points(self): - for sphere in list(self._spheres): # will remove itself, so copy - self.remove_point(sphere) - assert sum(len(v) for v in self.picked_points.values()) == 0 - assert len(self.pick_table) == 0 - assert len(self._spheres) == 0 - - def plot_time_course(self, hemi, vertex_id, color): - if self.mpl_canvas is None: - return - time = self.brain._data['time'].copy() # avoid circular ref - if hemi == 'vol': - hemi_str = 'V' - xfm = _read_talxfm( - self.brain._subject_id, self.brain._subjects_dir) - if self.brain._units == 'm': - xfm['trans'][:3, 3] /= 1000. - ijk = np.unravel_index( - vertex_id, self.brain._data[hemi]['grid_shape'], order='F') - src_mri_t = self.brain._data[hemi]['grid_src_mri_t'] - mni = apply_trans(np.dot(xfm['trans'], src_mri_t), ijk) - else: - hemi_str = 'L' if hemi == 'lh' else 'R' - mni = vertex_to_mni( - vertices=vertex_id, - hemis=0 if hemi == 'lh' else 1, - subject=self.brain._subject_id, - subjects_dir=self.brain._subjects_dir - ) - label = "{}:{} MNI: {}".format( - hemi_str, str(vertex_id).ljust(6), - ', '.join('%5.1f' % m for m in mni)) - act_data, smooth = self.act_data_smooth[hemi] - if smooth is not None: - act_data = smooth[vertex_id].dot(act_data)[0] - else: - act_data = act_data[vertex_id].copy() - line = self.mpl_canvas.plot( - time, - act_data, - label=label, - lw=1., - color=color, - zorder=4, - ) - return line - - def plot_time_line(self): - if self.mpl_canvas is None: - return - if isinstance(self.show_traces, bool) and self.show_traces: - # add time information - current_time = self.brain._current_time - if not hasattr(self, "time_line"): - self.time_line = self.mpl_canvas.plot_time_line( - x=current_time, - label='time', - color=self.brain._fg_color, - lw=1, - ) - self.time_line.set_xdata(current_time) - self.mpl_canvas.update_plot() - - def help(self): - pairs = [ - ('?', 'Display help window'), - ('i', 'Toggle interface'), - ('s', 'Apply auto-scaling'), - ('r', 'Restore original clim'), - ('c', 'Clear all traces'), - ('Space', 'Start/Pause playback'), - ] - text1, text2 = zip(*pairs) - text1 = '\n'.join(text1) - text2 = '\n'.join(text2) - _show_help( - col1=text1, - col2=text2, - width=5, - height=2, - ) - - def clear_callbacks(self): - for callback in self.callbacks.values(): - if callback is not None: - if hasattr(callback, "plotter"): - callback.plotter = None - if hasattr(callback, "brain"): - callback.brain = None - if hasattr(callback, "slider_rep"): - callback.slider_rep = None - self.callbacks.clear() - - @safe_event - def clean(self): - # resolve the reference cycle - self.clear_points() - self.clear_callbacks() - self.actions.clear() - self.sliders.clear() - self.reps = None - self.brain.time_viewer = None - self.brain = None - self.plotter = None - self.main_menu = None - self.window = None - self.tool_bar = None - self.status_bar = None - self.interactor = None - if self.mpl_canvas is not None: - self.mpl_canvas.clear() - self.mpl_canvas = None - self.time_actor = None - self.picked_renderer = None - for key in list(self.act_data_smooth.keys()): - self.act_data_smooth[key] = None - - -def _get_range(brain): - val = np.abs(np.concatenate(list(brain._current_act_data.values()))) - return [np.min(val), np.max(val)] - - -def _normalize(point, shape): - return (point[0] / shape[1], point[1] / shape[0]) diff --git a/mne/viz/_brain/mplcanvas.py b/mne/viz/_brain/mplcanvas.py index aff7d3d6b4e..f899b196b19 100644 --- a/mne/viz/_brain/mplcanvas.py +++ b/mne/viz/_brain/mplcanvas.py @@ -64,9 +64,9 @@ def update_plot(self): leg = self.axes.legend( prop={'family': 'monospace', 'size': 'small'}, framealpha=0.5, handlelength=1., - facecolor=self.time_viewer.brain._bg_color) + facecolor=self.time_viewer._bg_color) for text in leg.get_texts(): - text.set_color(self.time_viewer.brain._fg_color) + text.set_color(self.time_viewer._fg_color) with warnings.catch_warnings(record=True): warnings.filterwarnings('ignore', 'constrained_layout') self.canvas.draw() diff --git a/mne/viz/_brain/tests/test_brain.py b/mne/viz/_brain/tests/test_brain.py index dbc785b423e..5c10efd74be 100644 --- a/mne/viz/_brain/tests/test_brain.py +++ b/mne/viz/_brain/tests/test_brain.py @@ -20,7 +20,7 @@ setup_volume_source_space) from mne.datasets import testing from mne.utils import check_version -from mne.viz._brain import Brain, _TimeViewer, _LinkViewer, _BrainScraper +from mne.viz._brain import Brain, _LinkViewer, _BrainScraper from mne.viz._brain.colormap import calculate_lut from matplotlib import cm, image @@ -243,50 +243,46 @@ def test_brain_save_movie(tmpdir, renderer): @testing.requires_testing_data @pytest.mark.slowtest -def test_brain_timeviewer(renderer_interactive, pixel_ratio): - """Test _TimeViewer primitives.""" +def test_brain_time_viewer(renderer_interactive, pixel_ratio): + """Test time viewer primitives.""" if renderer_interactive._get_3d_backend() != 'pyvista': pytest.skip('TimeViewer tests only supported on PyVista') - brain_data = _create_testing_brain(hemi='both', show_traces=False) - - with pytest.raises(RuntimeError, match='already'): - _TimeViewer(brain_data) - time_viewer = brain_data.time_viewer - time_viewer.callbacks["time"](value=0) - time_viewer.callbacks["orientation_lh_0_0"]( + brain = _create_testing_brain(hemi='both', show_traces=False) + brain.callbacks["time"](value=0) + brain.callbacks["orientation_lh_0_0"]( value='lat', update_widget=True ) - time_viewer.callbacks["orientation_lh_0_0"]( + brain.callbacks["orientation_lh_0_0"]( value='medial', update_widget=True ) - time_viewer.callbacks["time"]( + brain.callbacks["time"]( value=0.0, time_as_index=False, ) - time_viewer.callbacks["smoothing"](value=1) - time_viewer.callbacks["fmin"](value=12.0) - time_viewer.callbacks["fmax"](value=4.0) - time_viewer.callbacks["fmid"](value=6.0) - time_viewer.callbacks["fmid"](value=4.0) - time_viewer.callbacks["fscale"](value=1.1) - time_viewer.callbacks["fmin"](value=12.0) - time_viewer.callbacks["fmid"](value=4.0) - time_viewer.toggle_interface() - time_viewer.callbacks["playback_speed"](value=0.1) - time_viewer.toggle_playback() - time_viewer.apply_auto_scaling() - time_viewer.restore_user_scaling() - time_viewer.reset() + brain.callbacks["smoothing"](value=1) + brain.callbacks["fmin"](value=12.0) + brain.callbacks["fmax"](value=4.0) + brain.callbacks["fmid"](value=6.0) + brain.callbacks["fmid"](value=4.0) + brain.callbacks["fscale"](value=1.1) + brain.callbacks["fmin"](value=12.0) + brain.callbacks["fmid"](value=4.0) + brain.toggle_interface() + brain.callbacks["playback_speed"](value=0.1) + brain.toggle_playback() + brain.apply_auto_scaling() + brain.restore_user_scaling() + brain.reset() plt.close('all') - time_viewer.help() + brain.help() assert len(plt.get_fignums()) == 1 plt.close('all') # screenshot - brain_data.show_view(view=dict(azimuth=180., elevation=90.)) - img = brain_data.screenshot(mode='rgb') + brain.show_view(view=dict(azimuth=180., elevation=90.)) + img = brain.screenshot(mode='rgb') want_shape = np.array([300 * pixel_ratio, 300 * pixel_ratio, 3]) assert_allclose(img.shape, want_shape) @@ -304,25 +300,22 @@ def test_brain_timeviewer(renderer_interactive, pixel_ratio): pytest.param('mixed', marks=pytest.mark.slowtest), ]) @pytest.mark.slowtest -def test_brain_timeviewer_traces(renderer_interactive, hemi, src, tmpdir): - """Test _TimeViewer traces.""" +def test_brain_traces(renderer_interactive, hemi, src, tmpdir): + """Test brain traces.""" if renderer_interactive._get_3d_backend() != 'pyvista': pytest.skip('Only PyVista supports traces') - brain_data = _create_testing_brain( + brain = _create_testing_brain( hemi=hemi, surf='white', src=src, show_traces=0.5, initial_time=0, volume_options=None, # for speed, don't upsample n_time=1 if src == 'mixed' else 5, ) - with pytest.raises(RuntimeError, match='already'): - _TimeViewer(brain_data) - time_viewer = brain_data.time_viewer - assert time_viewer.show_traces - assert hasattr(time_viewer, "picked_points") - assert hasattr(time_viewer, "_spheres") + assert brain.show_traces + assert hasattr(brain, "picked_points") + assert hasattr(brain, "_spheres") # test points picked by default - picked_points = brain_data.get_picked_points() - spheres = time_viewer._spheres + picked_points = brain.get_picked_points() + spheres = brain._spheres hemi_str = list() if src in ('surface', 'mixed'): hemi_str.extend([hemi] if hemi in ('lh', 'rh') else ['lh', 'rh']) @@ -336,7 +329,7 @@ def test_brain_timeviewer_traces(renderer_interactive, hemi, src, tmpdir): assert len(spheres) == n_spheres # test removing points - time_viewer.clear_points() + brain.clear_points() assert len(spheres) == 0 for key in ('lh', 'rh', 'vol'): assert len(picked_points[key]) == 0 @@ -346,20 +339,20 @@ def test_brain_timeviewer_traces(renderer_interactive, hemi, src, tmpdir): for idx, current_hemi in enumerate(hemi_str): assert len(spheres) == 0 if current_hemi == 'vol': - current_mesh = brain_data._data['vol']['grid'] - vertices = brain_data._data['vol']['vertices'] + current_mesh = brain._data['vol']['grid'] + vertices = brain._data['vol']['vertices'] values = current_mesh.cell_arrays['values'][vertices] cell_id = vertices[np.argmax(np.abs(values))] else: - current_mesh = brain_data._hemi_meshes[current_hemi] + current_mesh = brain._hemi_meshes[current_hemi] cell_id = rng.randint(0, current_mesh.n_cells) - test_picker = TstVTKPicker(None, None, current_hemi, brain_data) - assert time_viewer.on_pick(test_picker, None) is None + test_picker = TstVTKPicker(None, None, current_hemi, brain) + assert brain._on_pick(test_picker, None) is None test_picker = TstVTKPicker( - current_mesh, cell_id, current_hemi, brain_data) + current_mesh, cell_id, current_hemi, brain) assert cell_id == test_picker.cell_id assert test_picker.point_id is None - time_viewer.on_pick(test_picker, None) + brain._on_pick(test_picker, None) assert test_picker.point_id is not None assert len(picked_points[current_hemi]) == 1 assert picked_points[current_hemi][0] == test_picker.point_id @@ -378,8 +371,8 @@ def test_brain_timeviewer_traces(renderer_interactive, hemi, src, tmpdir): mni = vertex_to_mni( vertices=vertex_id, hemis=hemi_int, - subject=brain_data._subject_id, - subjects_dir=brain_data._subjects_dir + subject=brain._subject_id, + subjects_dir=brain._subjects_dir ) label = "{}:{} MNI: {}".format( hemi_prefix, str(vertex_id).ljust(6), @@ -390,11 +383,11 @@ def test_brain_timeviewer_traces(renderer_interactive, hemi, src, tmpdir): # remove the sphere by clicking in its vicinity old_len = len(spheres) test_picker._actors = sum((s._actors for s in spheres), []) - time_viewer.on_pick(test_picker, None) + brain._on_pick(test_picker, None) assert len(spheres) < old_len - screenshot = brain_data.screenshot() - screenshot_all = brain_data.screenshot(time_viewer=True) + screenshot = brain.screenshot() + screenshot_all = brain.screenshot(time_viewer=True) assert screenshot.shape[0] < screenshot_all.shape[0] # and the scraper for it (will close the instance) # only test one condition to save time @@ -403,7 +396,7 @@ def test_brain_timeviewer_traces(renderer_interactive, hemi, src, tmpdir): return fnames = [str(tmpdir.join(f'temp_{ii}.png')) for ii in range(2)] block_vars = dict(image_path_iterator=iter(fnames), - example_globals=dict(brain=brain_data)) + example_globals=dict(brain=brain)) block = ('code', """ something # brain.save_movie(time_dilation=1, framerate=1, From 1cf1559030d0ffc8c9a18c2b0776793883409ce1 Mon Sep 17 00:00:00 2001 From: Guillaume Favelier Date: Thu, 8 Oct 2020 13:05:27 +0200 Subject: [PATCH 2/4] Fix docstring --- mne/viz/_brain/_brain.py | 20 +++++- mne/viz/_brain/_linkviewer.py | 111 ++++++++++++++--------------- mne/viz/_brain/_scraper.py | 2 +- mne/viz/_brain/mplcanvas.py | 16 ++--- mne/viz/_brain/tests/test_brain.py | 2 +- 5 files changed, 83 insertions(+), 68 deletions(-) diff --git a/mne/viz/_brain/_brain.py b/mne/viz/_brain/_brain.py index 14adbbe48ff..067516122a8 100644 --- a/mne/viz/_brain/_brain.py +++ b/mne/viz/_brain/_brain.py @@ -330,6 +330,8 @@ def __init__(self, subject_id, hemi, surf, title=None, def setup_time_viewer(self, time_viewer=True, show_traces=True): """Configure the time viewer parameters. + Parameters + ---------- time_viewer : bool If True, enable widgets interaction. Defaults to True. @@ -475,6 +477,8 @@ def ensure_minimum_sizes(self): def toggle_interface(self, value=None): """Toggle the interface. + Parameters + ---------- value : bool | None If True, the widgets are shown and if False, they are hidden. If None, the state of the widgets is @@ -530,6 +534,8 @@ def restore_user_scaling(self): def toggle_playback(self, value=None): """Toggle time playback. + Parameters + ---------- value : bool | None If True, automatic time playback is enabled and if False, it's disabled. If None, the state of time playback is toggled. @@ -565,7 +571,13 @@ def reset(self): self.plotter.update() def set_playback_speed(self, speed): - """Set the time playback speed.""" + """Set the time playback speed. + + Parameters + ---------- + speed : float + The speed of the playback. + """ self.playback_speed = speed @safe_event @@ -1093,6 +1105,8 @@ def _on_pick(self, vtk_picker, event): def add_point(self, hemi, mesh, vertex_id): """Pick a vertex on the brain. + Parameters + ---------- hemi : str The hemisphere id of the vertex. mesh : vtkPolyData @@ -1167,6 +1181,8 @@ def add_point(self, hemi, mesh, vertex_id): def remove_point(self, mesh): """Remove the picked point from its glyph. + Parameters + ---------- mesh : vtkPolyData The mesh associated to the point to remove. """ @@ -1204,6 +1220,8 @@ def clear_points(self): def plot_time_course(self, hemi, vertex_id, color): """Plot the vertex time course. + Parameters + ---------- hemi : str The hemisphere id of the vertex. vertex_id : int diff --git a/mne/viz/_brain/_linkviewer.py b/mne/viz/_brain/_linkviewer.py index 00eab9929ea..d6b48d03f80 100644 --- a/mne/viz/_brain/_linkviewer.py +++ b/mne/viz/_brain/_linkviewer.py @@ -8,11 +8,12 @@ class _LinkViewer(object): - """Class to link multiple _Brain objects.""" + """Class to link multiple Brain objects.""" def __init__(self, brains, time=True, camera=False, colorbar=True, picking=False): - self.time_viewers = brains + self.brains = brains + self.leader = self.brains[0] # select a brain as leader # check time infos times = [brain._times for brain in brains] @@ -39,63 +40,61 @@ def __init__(self, brains, time=True, camera=False, colorbar=True, ) # link toggle to start/pause playback - for time_viewer in self.time_viewers: - time_viewer.actions["play"].triggered.disconnect() - time_viewer.actions["play"].triggered.connect( + for brain in self.brains: + brain.actions["play"].triggered.disconnect() + brain.actions["play"].triggered.connect( self.toggle_playback) # link time course canvas def _time_func(*args, **kwargs): - for time_viewer in self.time_viewers: - time_viewer.callbacks["time"](*args, **kwargs) + for brain in self.brains: + brain.callbacks["time"](*args, **kwargs) - for time_viewer in self.time_viewers: - if time_viewer.show_traces: - time_viewer.mpl_canvas.time_func = _time_func + for brain in self.brains: + if brain.show_traces: + brain.mpl_canvas.time_func = _time_func if picking: def _func_add(*args, **kwargs): - for time_viewer in self.time_viewers: - time_viewer._add_point(*args, **kwargs) - time_viewer.plotter.update() + for brain in self.brains: + brain._add_point(*args, **kwargs) + brain.plotter.update() def _func_remove(*args, **kwargs): - for time_viewer in self.time_viewers: - time_viewer._remove_point(*args, **kwargs) + for brain in self.brains: + brain._remove_point(*args, **kwargs) # save initial picked points initial_points = dict() for hemi in ('lh', 'rh'): initial_points[hemi] = set() - for time_viewer in self.time_viewers: + for brain in self.brains: initial_points[hemi] |= \ - set(time_viewer.picked_points[hemi]) + set(brain.picked_points[hemi]) # link the viewers - for time_viewer in self.time_viewers: - time_viewer.clear_points() - time_viewer._add_point = time_viewer.add_point - time_viewer.add_point = _func_add - time_viewer._remove_point = time_viewer.remove_point - time_viewer.remove_point = _func_remove + for brain in self.brains: + brain.clear_points() + brain._add_point = brain.add_point + brain.add_point = _func_add + brain._remove_point = brain.remove_point + brain.remove_point = _func_remove # link the initial points - leader = self.time_viewers[0] # select a time_viewer as leader for hemi in initial_points.keys(): - if hemi in time_viewer._hemi_meshes: - mesh = time_viewer._hemi_meshes[hemi] + if hemi in brain._hemi_meshes: + mesh = brain._hemi_meshes[hemi] for vertex_id in initial_points[hemi]: - leader.add_point(hemi, mesh, vertex_id) + self.leader.add_point(hemi, mesh, vertex_id) if colorbar: - leader = self.time_viewers[0] # select a time_viewer as leader - fmin = leader._data["fmin"] - fmid = leader._data["fmid"] - fmax = leader._data["fmax"] - for time_viewer in self.time_viewers: - time_viewer.callbacks["fmin"](fmin) - time_viewer.callbacks["fmid"](fmid) - time_viewer.callbacks["fmax"](fmax) + fmin = self.leader._data["fmin"] + fmid = self.leader._data["fmid"] + fmax = self.leader._data["fmax"] + for brain in self.brains: + brain.callbacks["fmin"](fmin) + brain.callbacks["fmid"](fmid) + brain.callbacks["fmax"](fmax) for slider_name in ('fmin', 'fmid', 'fmax'): func = getattr(self, "set_" + slider_name) @@ -106,37 +105,36 @@ def _func_remove(*args, **kwargs): ) def set_fmin(self, value): - for time_viewer in self.time_viewers: - time_viewer.callbacks["fmin"](value) + for brain in self.brains: + brain.callbacks["fmin"](value) def set_fmid(self, value): - for time_viewer in self.time_viewers: - time_viewer.callbacks["fmid"](value) + for brain in self.brains: + brain.callbacks["fmid"](value) def set_fmax(self, value): - for time_viewer in self.time_viewers: - time_viewer.callbacks["fmax"](value) + for brain in self.brains: + brain.callbacks["fmax"](value) def set_time_point(self, value): - for time_viewer in self.time_viewers: - time_viewer.callbacks["time"](value, update_widget=True) + for brain in self.brains: + brain.callbacks["time"](value, update_widget=True) def set_playback_speed(self, value): - for time_viewer in self.time_viewers: - time_viewer.callbacks["playback_speed"](value, update_widget=True) + for brain in self.brains: + brain.callbacks["playback_speed"](value, update_widget=True) def toggle_playback(self): - leader = self.time_viewers[0] # select a time_viewer as leader - value = leader.callbacks["time"].slider_rep.GetValue() + value = self.leader.callbacks["time"].slider_rep.GetValue() # synchronize starting points before playback self.set_time_point(value) - for time_viewer in self.time_viewers: - time_viewer.toggle_playback() + for brain in self.brains: + brain.toggle_playback() def link_sliders(self, name, callback, event_type): from ..backends._pyvista import _update_slider_callback - for time_viewer in self.time_viewers: - slider = time_viewer.sliders[name] + for brain in self.brains: + slider = brain.sliders[name] if slider is not None: _update_slider_callback( slider=slider, @@ -148,12 +146,11 @@ def link_cameras(self): from ..backends._pyvista import _add_camera_callback def _update_camera(vtk_picker, event): - for time_viewer in self.time_viewers: - time_viewer.plotter.update() + for brain in self.brains: + brain.plotter.update() - leader = self.time_viewers[0] # select a time_viewer as leader - camera = leader.plotter.camera + camera = self.leader.plotter.camera _add_camera_callback(camera, _update_camera) - for time_viewer in self.time_viewers: - for renderer in time_viewer.plotter.renderers: + for brain in self.brains: + for renderer in brain.plotter.renderers: renderer.camera = camera diff --git a/mne/viz/_brain/_scraper.py b/mne/viz/_brain/_scraper.py index b79b91d58fb..88d9088d6c8 100644 --- a/mne/viz/_brain/_scraper.py +++ b/mne/viz/_brain/_scraper.py @@ -51,7 +51,7 @@ def __call__(self, block, block_vars, gallery_conf): ('time_viewer', False)]: if key not in kwargs: kwargs[key] = default - if hasattr(brain, 'time_viewer'): + if brain.time_viewer: assert kwargs['time_viewer'], 'Must use time_viewer=True' frames = brain._make_movie_frames(callback=None, **kwargs) diff --git a/mne/viz/_brain/mplcanvas.py b/mne/viz/_brain/mplcanvas.py index f899b196b19..23b9f4d7295 100644 --- a/mne/viz/_brain/mplcanvas.py +++ b/mne/viz/_brain/mplcanvas.py @@ -11,15 +11,15 @@ class MplCanvas(object): """Ultimately, this is a QWidget (as well as a FigureCanvasAgg, etc.).""" - def __init__(self, time_viewer, width, height, dpi): + def __init__(self, brain, width, height, dpi): from PyQt5 import QtWidgets from matplotlib import rc_context from matplotlib.figure import Figure from matplotlib.backends.backend_qt5agg import FigureCanvasQTAgg - if time_viewer.separate_canvas: + if brain.separate_canvas: parent = None else: - parent = time_viewer.window + parent = brain.window # prefer constrained layout here but live with tight_layout otherwise context = nullcontext extra_events = ('resize',) @@ -40,8 +40,8 @@ def __init__(self, time_viewer, width, height, dpi): QtWidgets.QSizePolicy.Expanding ) FigureCanvasQTAgg.updateGeometry(self.canvas) - self.time_viewer = time_viewer - self.time_func = time_viewer.callbacks["time"] + self.brain = brain + self.time_func = brain.callbacks["time"] for event in ('button_press', 'motion_notify') + extra_events: self.canvas.mpl_connect( event + '_event', getattr(self, 'on_' + event)) @@ -64,9 +64,9 @@ def update_plot(self): leg = self.axes.legend( prop={'family': 'monospace', 'size': 'small'}, framealpha=0.5, handlelength=1., - facecolor=self.time_viewer._bg_color) + facecolor=self.brain._bg_color) for text in leg.get_texts(): - text.set_color(self.time_viewer._fg_color) + text.set_color(self.brain._fg_color) with warnings.catch_warnings(record=True): warnings.filterwarnings('ignore', 'constrained_layout') self.canvas.draw() @@ -106,7 +106,7 @@ def clear(self): self.close() self.axes.clear() self.fig.clear() - self.time_viewer = None + self.brain = None self.canvas = None on_motion_notify = on_button_press # for now they can be the same diff --git a/mne/viz/_brain/tests/test_brain.py b/mne/viz/_brain/tests/test_brain.py index 5c10efd74be..373fcc4cb08 100644 --- a/mne/viz/_brain/tests/test_brain.py +++ b/mne/viz/_brain/tests/test_brain.py @@ -443,7 +443,7 @@ def test_brain_linkviewer(renderer_interactive): picking=True, ) link_viewer.set_time_point(value=0) - link_viewer.time_viewers[0].mpl_canvas.time_func(0) + link_viewer.brains[0].mpl_canvas.time_func(0) link_viewer.set_fmin(0) link_viewer.set_fmid(0.5) link_viewer.set_fmax(1) From c5de127473299a3a00090d2b85fc637d20a57371 Mon Sep 17 00:00:00 2001 From: Guillaume Favelier Date: Thu, 8 Oct 2020 13:21:25 +0200 Subject: [PATCH 3/4] Fix docstring --- mne/viz/_brain/_brain.py | 29 ++++++++++++++++++----------- 1 file changed, 18 insertions(+), 11 deletions(-) diff --git a/mne/viz/_brain/_brain.py b/mne/viz/_brain/_brain.py index 067516122a8..3e590d86680 100644 --- a/mne/viz/_brain/_brain.py +++ b/mne/viz/_brain/_brain.py @@ -233,7 +233,7 @@ def __init__(self, subject_id, hemi, surf, title=None, self._size = size if len(size) == 2 else size * 2 # 1-tuple to 2-tuple self.time_viewer = False - self._notebook = (_get_3d_backend() == "notebook") + self.notebook = (_get_3d_backend() == "notebook") self._hemi = hemi self._units = units self._alpha = float(alpha) @@ -313,6 +313,8 @@ def __init__(self, subject_id, hemi, surf, title=None, self.interaction = interaction self._closed = False + if show: + self.show() # update the views once the geometry is all set for h in self._hemis: for ri, ci, v in self._iter_views(h): @@ -324,9 +326,6 @@ def __init__(self, subject_id, hemi, surf, title=None, if hemi == 'rh' and hasattr(self._renderer, "_orient_lights"): self._renderer._orient_lights() - if show: - self.show() - def setup_time_viewer(self, time_viewer=True, show_traces=True): """Configure the time viewer parameters. @@ -344,13 +343,10 @@ def setup_time_viewer(self, time_viewer=True, show_traces=True): self.orientation = list(_lh_views_dict.keys()) self.default_smoothing_range = [0, 15] - # detect notebook - if self._notebook: - self.notebook = True + # setup notebook + if self.notebook: self._configure_notebook() return - else: - self.notebook = False # Default configuration self.playback = False @@ -1116,7 +1112,8 @@ def add_point(self, hemi, mesh, vertex_id): Returns ------- - The glyph created for the picked point. + sphere : vtkPolyData + The glyph created for the picked point. """ # skip if the wrong hemi is selected if self.act_data_smooth[hemi][0] is None: @@ -1228,6 +1225,11 @@ def plot_time_course(self, hemi, vertex_id, color): The vertex identifier in the mesh. color : matplotlib color The color of the time course. + + Returns + ------- + line : matplotlib line + The time line object. """ if self.mpl_canvas is None: return @@ -2596,6 +2598,11 @@ def save_movie(self, filename, time_dilation=4., tmin=None, tmax=None, %(brain_screenshot_time_viewer)s **kwargs : dict Specify additional options for :mod:`imageio`. + + Returns + ------- + dialog : QDialog + The opened dialog is returned for testing purpose only. """ if self.time_viewer: try: @@ -2854,7 +2861,7 @@ def enable_depth_peeling(self): def _update(self): from ..backends import renderer if renderer.get_3d_backend() in ['pyvista', 'notebook']: - if self._notebook and self._renderer.figure.display is not None: + if self.notebook and self._renderer.figure.display is not None: self._renderer.figure.display.update() def get_picked_points(self): From 97694940fb7f35c33d788d06883a4c41cd3b8142 Mon Sep 17 00:00:00 2001 From: Guillaume Favelier Date: Thu, 8 Oct 2020 13:39:16 +0200 Subject: [PATCH 4/4] Fix docstring --- mne/viz/_brain/_brain.py | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/mne/viz/_brain/_brain.py b/mne/viz/_brain/_brain.py index 3e590d86680..8c5c61195e4 100644 --- a/mne/viz/_brain/_brain.py +++ b/mne/viz/_brain/_brain.py @@ -1105,14 +1105,14 @@ def add_point(self, hemi, mesh, vertex_id): ---------- hemi : str The hemisphere id of the vertex. - mesh : vtkPolyData + mesh : object The mesh where picking is expected. vertex_id : int The vertex identifier in the mesh. Returns ------- - sphere : vtkPolyData + sphere : object The glyph created for the picked point. """ # skip if the wrong hemi is selected @@ -1180,7 +1180,7 @@ def remove_point(self, mesh): Parameters ---------- - mesh : vtkPolyData + mesh : object The mesh associated to the point to remove. """ vertex_id = mesh._vertex_id @@ -1228,7 +1228,7 @@ def plot_time_course(self, hemi, vertex_id, color): Returns ------- - line : matplotlib line + line : matplotlib object The time line object. """ if self.mpl_canvas is None: @@ -2601,7 +2601,7 @@ def save_movie(self, filename, time_dilation=4., tmin=None, tmax=None, Returns ------- - dialog : QDialog + dialog : object The opened dialog is returned for testing purpose only. """ if self.time_viewer: