This commit is contained in:
cjw
2026-02-12 23:22:11 +08:00
parent 7b09eb3d89
commit 89660bba4e
5988 changed files with 2517516 additions and 0 deletions
@@ -0,0 +1,28 @@
"""Plotting utilities."""
from __future__ import annotations
from .algorithms import active_scalars_algorithm as active_scalars_algorithm
from .algorithms import add_ids_algorithm as add_ids_algorithm
from .algorithms import algorithm_to_mesh_handler as algorithm_to_mesh_handler
from .algorithms import cell_data_to_point_data_algorithm as cell_data_to_point_data_algorithm
from .algorithms import crinkle_algorithm as crinkle_algorithm
from .algorithms import decimation_algorithm as decimation_algorithm
from .algorithms import extract_surface_algorithm as extract_surface_algorithm
from .algorithms import outline_algorithm as outline_algorithm
from .algorithms import point_data_to_cell_data_algorithm as point_data_to_cell_data_algorithm
from .algorithms import pointset_to_polydata_algorithm as pointset_to_polydata_algorithm
from .algorithms import set_algorithm_input as set_algorithm_input
from .algorithms import triangulate_algorithm as triangulate_algorithm
from .cubemap import cubemap as cubemap
from .cubemap import cubemap_from_filenames as cubemap_from_filenames
from .gl_checks import check_depth_peeling as check_depth_peeling
from .gl_checks import uses_egl as uses_egl
from .regression import compare_images as compare_images
from .regression import image_from_window as image_from_window
from .regression import remove_alpha as remove_alpha
from .regression import run_image_filter as run_image_filter
from .regression import wrap_image_array as wrap_image_array
from .sphinx_gallery import Scraper as Scraper
from .sphinx_gallery import _get_sg_image_scraper as _get_sg_image_scraper
from .xvfb import start_xvfb as start_xvfb
@@ -0,0 +1,633 @@
"""Internal :vtk:`vtkAlgorithm` support helpers."""
from __future__ import annotations
import traceback
from typing import TYPE_CHECKING
import numpy as np
import pyvista
from pyvista._deprecate_positional_args import _deprecate_positional_args
from pyvista.core.errors import PyVistaPipelineError
from pyvista.core.utilities.helpers import wrap
from pyvista.core.utilities.misc import _NoNewAttrMixin
from pyvista.plotting import _vtk
if TYPE_CHECKING:
from pyvista.core.utilities.arrays import CellLiteral
from pyvista.core.utilities.arrays import PointLiteral
def algorithm_to_mesh_handler(
mesh_or_algo, port=0
) -> tuple[pyvista.DataSet, _vtk.vtkAlgorithm | _vtk.vtkAlgorithmOutput | None]:
"""Handle :vtk:`vtkAlgorithms` where mesh objects are expected.
This is a convenience method to handle :vtk:`vtkAlgorithms` when passed to methods
that expect a :class:`~pyvista.DataSet`. This method will check if the passed
object is a :vtk:`vtkAlgorithm` or :vtk:`vtkAlgorithmOutput` and if so,
return that algorithm's output dataset (mesh) as the mesh to be used by the
calling function.
Parameters
----------
mesh_or_algo : DataSet | :vtk:`vtkAlgorithm` | :vtk:`vtkAlgorithmOutput`
The input to be used as a data set (mesh) or :vtk:`vtkAlgorithm` object.
port : int, default: 0
If the input (``mesh_or_algo``) is an algorithm, this specifies which output
port to use on that algorithm for the returned mesh.
Returns
-------
mesh : pyvista.DataSet
The resulting mesh data set from the input.
algorithm : :vtk:`vtkAlgorithm` | :vtk:`vtkAlgorithmOutput` | None
If an algorithm is passed, it will be returned. Otherwise returns ``None``.
"""
if isinstance(mesh_or_algo, (_vtk.vtkAlgorithm, _vtk.vtkAlgorithmOutput)):
if isinstance(mesh_or_algo, _vtk.vtkAlgorithmOutput):
algo = mesh_or_algo.GetProducer()
# If vtkAlgorithmOutput, override port argument
port = mesh_or_algo.GetIndex()
output = mesh_or_algo
else:
algo = mesh_or_algo
output = algo.GetOutputPort(port)
algo.Update() # NOTE: this could be expensive... but we need it to get the mesh
# for legacy implementation. This can be refactored.
mesh = wrap(algo.GetOutputDataObject(port))
if mesh is None:
# This is known to happen with vtkPointSet and VTKPythonAlgorithmBase
# see workaround in PreserveTypeAlgorithmBase.
# This check remains as a fail-safe.
msg = 'The passed algorithm is failing to produce an output.' # type: ignore[unreachable]
raise PyVistaPipelineError(msg)
# NOTE: Return the vtkAlgorithmOutput only if port is non-zero. Segfaults can sometimes
# happen with vtkAlgorithmOutput. This logic will mostly avoid those issues.
# See https://gitlab.kitware.com/vtk/vtk/-/issues/18776
return mesh, output if port != 0 else algo
return mesh_or_algo, None
def set_algorithm_input(alg, inp, port=0):
"""Set the input to a :vtk:`vtkAlgorithm`.
Parameters
----------
alg : :vtk:`vtkAlgorithm`
The algorithm whose input is being set.
inp : :vtk:`vtkAlgorithm` | :vtk:`vtkAlgorithmOutput` | :vtk:`vtkDataObject`
The input to the algorithm.
port : int, default: 0
The input port.
"""
if isinstance(inp, _vtk.vtkAlgorithm):
alg.SetInputConnection(port, inp.GetOutputPort())
elif isinstance(inp, _vtk.vtkAlgorithmOutput):
alg.SetInputConnection(port, inp)
else:
alg.SetInputDataObject(port, inp)
class PreserveTypeAlgorithmBase(
_NoNewAttrMixin, _vtk.DisableVtkSnakeCase, _vtk.VTKPythonAlgorithmBase
):
"""Base algorithm to preserve type.
Parameters
----------
nInputPorts : int, default: 1
Number of input ports for the algorithm.
nOutputPorts : int, default: 1
Number of output ports for the algorithm.
"""
def __init__(self, nInputPorts=1, nOutputPorts=1):
"""Initialize algorithm."""
_vtk.VTKPythonAlgorithmBase.__init__(
self,
nInputPorts=nInputPorts,
nOutputPorts=nOutputPorts,
)
def GetInputData(self, inInfo, port, idx):
"""Get input data object.
This will convert :vtk:`vtkPointSet` to :vtk:`vtkPolyData`.
Parameters
----------
inInfo : :vtk:`vtkInformation`
The information object associated with the input port.
port : int
The index of the input port.
idx : int
The index of the data object within the input port.
Returns
-------
:vtk:`vtkDataObject`
The input data object.
"""
inp = wrap(_vtk.VTKPythonAlgorithmBase.GetInputData(self, inInfo, port, idx))
if isinstance(inp, pyvista.PointSet):
return inp.cast_to_polydata()
return inp
# THIS IS CRUCIAL to preserve data type through filter
def RequestDataObject(self, _request, inInfo, outInfo) -> int:
"""Preserve data type.
Parameters
----------
_request : :vtk:`vtkInformation`
The request object for the filter.
inInfo : :vtk:`vtkInformationVector`
The input information vector for the filter.
outInfo : :vtk:`vtkInformationVector`
The output information vector for the filter.
Returns
-------
int
Returns 1 if successful.
"""
class_name = self.GetInputData(inInfo, 0, 0).GetClassName()
if class_name == 'vtkPointSet':
# See https://gitlab.kitware.com/vtk/vtk/-/issues/18771
self.OutputType = 'vtkPolyData'
else:
self.OutputType = class_name
self.FillOutputPortInformation(0, outInfo.GetInformationObject(0))
return 1
class ActiveScalarsAlgorithm(PreserveTypeAlgorithmBase):
"""Algorithm to control active scalars.
The output of this filter is a shallow copy of the input data
set with the active scalars set as specified.
Parameters
----------
name : str
Name of scalars used to set as active on the output mesh.
Accepts a string name of an array that is present on the mesh.
Array should be sized as a single vector.
preference : str, default: 'point'
When ``mesh.n_points == mesh.n_cells`` and setting
scalars, this parameter sets how the scalars will be
mapped to the mesh. The default, ``'point'``, causes the
scalars to be associated with the mesh points. Can be
either ``'point'`` or ``'cell'``.
"""
def __init__(self, name: str, preference: PointLiteral | CellLiteral = 'point'):
"""Initialize algorithm."""
super().__init__()
self.scalars_name = name
self.preference = preference
def RequestData(self, _request, inInfo, outInfo) -> int:
"""Perform algorithm execution.
Parameters
----------
_request : :vtk:`vtkInformation`
The request object.
inInfo : :vtk:`vtkInformationVector`
Information about the input data.
outInfo : :vtk:`vtkInformationVector`
Information about the output data.
Returns
-------
int
1 on success.
"""
try:
inp = wrap(self.GetInputData(inInfo, 0, 0))
out = self.GetOutputData(outInfo, 0)
output = inp.copy()
if output.n_arrays:
output.set_active_scalars(self.scalars_name, preference=self.preference)
out.ShallowCopy(output)
except Exception: # pragma: no cover
traceback.print_exc()
raise
return 1
class PointSetToPolyDataAlgorithm(
_NoNewAttrMixin, _vtk.DisableVtkSnakeCase, _vtk.VTKPythonAlgorithmBase
):
"""Algorithm to cast PointSet to PolyData.
This is implemented with :func:`pyvista.PointSet.cast_to_polydata`.
"""
def __init__(self):
"""Initialize algorithm."""
_vtk.VTKPythonAlgorithmBase.__init__(
self,
nInputPorts=1,
nOutputPorts=1,
inputType='vtkPointSet',
outputType='vtkPolyData',
)
def RequestData(self, _request, inInfo, outInfo) -> int:
"""Perform algorithm execution.
Parameters
----------
_request : :vtk:`vtkInformation`
Information associated with the request.
inInfo : :vtk:`vtkInformationVector`
Information about the input data.
outInfo : :vtk:`vtkInformationVector`
Information about the output data.
Returns
-------
int
1 when successful.
"""
try:
inp = wrap(self.GetInputData(inInfo, 0, 0))
out = self.GetOutputData(outInfo, 0)
output = inp.cast_to_polydata(deep=False)
out.ShallowCopy(output)
except Exception: # pragma: no cover
traceback.print_exc()
raise
return 1
class AddIDsAlgorithm(PreserveTypeAlgorithmBase):
"""Algorithm to add point or cell IDs.
Output of this filter is a shallow copy of the input with
point and/or cell ID arrays added.
Parameters
----------
point_ids : bool, default: True
Whether to add point IDs.
cell_ids : bool, default: True
Whether to add cell IDs.
Raises
------
ValueError
If neither point IDs nor cell IDs are set.
"""
@_deprecate_positional_args
def __init__(self, point_ids: bool = True, cell_ids: bool = True): # noqa: FBT001, FBT002
"""Initialize algorithm."""
super().__init__()
if not point_ids and not cell_ids: # pragma: no cover
msg = 'IDs must be set for points or cells or both.'
raise ValueError(msg)
self.point_ids = point_ids
self.cell_ids = cell_ids
def RequestData(self, _request, inInfo, outInfo) -> int:
"""Perform algorithm execution.
Parameters
----------
_request : :vtk:`vtkInformation`
Information associated with the request.
inInfo : :vtk:`vtkInformationVector`
Information about the input data.
outInfo : :vtk:`vtkInformationVector`
Information about the output data.
Returns
-------
int
Returns 1 if the algorithm was successful.
Raises
------
Exception
If the algorithm fails to execute properly.
"""
try:
inp = wrap(self.GetInputData(inInfo, 0, 0))
out = self.GetOutputData(outInfo, 0)
output = inp.copy()
if self.point_ids:
output.point_data['point_ids'] = np.arange(0, output.n_points, dtype=int)
if self.cell_ids:
output.cell_data['cell_ids'] = np.arange(0, output.n_cells, dtype=int)
if output.active_scalars_name in ['point_ids', 'cell_ids']:
output.active_scalars_name = inp.active_scalars_name
out.ShallowCopy(output)
except Exception: # pragma: no cover
traceback.print_exc()
raise
return 1
class CrinkleAlgorithm(_NoNewAttrMixin, _vtk.DisableVtkSnakeCase, _vtk.VTKPythonAlgorithmBase):
"""Algorithm to crinkle cell IDs."""
def __init__(self):
"""Initialize algorithm."""
super().__init__(
nInputPorts=2,
outputType='vtkUnstructuredGrid',
)
def RequestData(self, _request, inInfo, outInfo) -> int:
"""Perform algorithm execution based on the input data and produce the output.
Parameters
----------
_request : :vtk:`vtkInformation`
The request information associated with the algorithm.
inInfo : :vtk:`vtkInformationVector`
Information vector describing the input data.
outInfo : :vtk:`vtkInformationVector`
Information vector where the output data should be placed.
Returns
-------
int
Status of the execution. Returns 1 on successful completion.
"""
try:
clipped = wrap(self.GetInputData(inInfo, 0, 0))
source = wrap(self.GetInputData(inInfo, 1, 0))
out = self.GetOutputData(outInfo, 0)
output = source.extract_cells(np.unique(clipped.cell_data['cell_ids']))
out.ShallowCopy(output)
except Exception: # pragma: no cover
traceback.print_exc()
raise
return 1
@_deprecate_positional_args(allowed=['inp'])
def outline_algorithm(inp, generate_faces: bool = False): # noqa: FBT001, FBT002
"""Add :vtk:`vtkOutlineFilter` to pipeline.
Parameters
----------
inp : pyvista.Common
Input data to be filtered.
generate_faces : bool, default: False
Whether to generate faces for the outline.
Returns
-------
:vtk:`vtkOutlineFilter`
Outline filter applied to the input data.
"""
alg = _vtk.vtkOutlineFilter()
set_algorithm_input(alg, inp)
alg.SetGenerateFaces(generate_faces)
return alg
@_deprecate_positional_args(allowed=['inp'])
def extract_surface_algorithm( # noqa: PLR0917
inp,
pass_pointid: bool = False, # noqa: FBT001, FBT002
pass_cellid: bool = False, # noqa: FBT001, FBT002
nonlinear_subdivision=1,
):
"""Add :vtk:`vtkDataSetSurfaceFilter` to pipeline.
Parameters
----------
inp : pyvista.Common
Input data to be filtered.
pass_pointid : bool, default: False
If ``True``, pass point IDs to the output.
pass_cellid : bool, default: False
If ``True``, pass cell IDs to the output.
nonlinear_subdivision : int, default: 1
Level of nonlinear subdivision.
Returns
-------
:vtk:`vtkDataSetSurfaceFilter`
Surface filter applied to the input data.
"""
surf_filter = _vtk.vtkDataSetSurfaceFilter()
surf_filter.SetPassThroughPointIds(pass_pointid)
surf_filter.SetPassThroughCellIds(pass_cellid)
if nonlinear_subdivision != 1:
surf_filter.SetNonlinearSubdivisionLevel(nonlinear_subdivision)
set_algorithm_input(surf_filter, inp)
return surf_filter
def active_scalars_algorithm(inp, name, preference='point'):
"""Add a filter that sets the active scalars.
Parameters
----------
inp : pyvista.Common
Input data to be filtered.
name : str
Name of the scalars to set as active.
preference : str, default: 'point'
Preference for the scalars to be set as active. Options are 'point', 'cell', or 'field'.
Returns
-------
:vtk:`vtkAlgorithm`
Active scalars filter applied to the input data.
"""
alg = ActiveScalarsAlgorithm(
name=name,
preference=preference,
)
set_algorithm_input(alg, inp)
return alg
def pointset_to_polydata_algorithm(inp):
"""Add a filter that casts PointSet to PolyData.
Parameters
----------
inp : pyvista.PointSet
Input point set to be cast to PolyData.
Returns
-------
:vtk:`vtkAlgorithm`
Filter that casts the input PointSet to PolyData.
"""
alg = PointSetToPolyDataAlgorithm()
set_algorithm_input(alg, inp)
return alg
@_deprecate_positional_args(allowed=['inp'])
def add_ids_algorithm(inp, point_ids: bool = True, cell_ids: bool = True): # noqa: FBT001, FBT002
"""Add a filter that adds point and/or cell IDs.
Parameters
----------
inp : pyvista.DataSet
The input data to which the IDs will be added.
point_ids : bool, default: True
If ``True``, point IDs will be added to the input data.
cell_ids : bool, default: True
If ``True``, cell IDs will be added to the input data.
Returns
-------
AddIDsAlgorithm
AddIDsAlgorithm filter.
"""
alg = AddIDsAlgorithm(point_ids=point_ids, cell_ids=cell_ids)
set_algorithm_input(alg, inp)
return alg
def crinkle_algorithm(clip, source):
"""Add a filter that crinkles a clip.
Parameters
----------
clip : pyvista.DataSet
The input data to be crinkled.
source : pyvista.DataSet
The source of the crinkle.
Returns
-------
CrinkleAlgorithm
CrinkleAlgorithm filter.
"""
alg = CrinkleAlgorithm()
set_algorithm_input(alg, clip, 0)
set_algorithm_input(alg, source, 1)
return alg
@_deprecate_positional_args(allowed=['inp'])
def cell_data_to_point_data_algorithm(inp, pass_cell_data: bool = False): # noqa: FBT001, FBT002
"""Add a filter that converts cell data to point data.
Parameters
----------
inp : pyvista.DataSet
The input data whose cell data will be converted to point data.
pass_cell_data : bool, default: False
If ``True``, the original cell data will be passed to the output.
Returns
-------
:vtk:`vtkCellDataToPointData`
The :vtk:`vtkCellDataToPointData` filter.
"""
alg = _vtk.vtkCellDataToPointData()
alg.SetPassCellData(pass_cell_data)
set_algorithm_input(alg, inp)
return alg
@_deprecate_positional_args(allowed=['inp'])
def point_data_to_cell_data_algorithm(inp, pass_point_data: bool = False): # noqa: FBT001, FBT002
"""Add a filter that converts point data to cell data.
Parameters
----------
inp : pyvista.DataSet
The input data whose point data will be converted to cell data.
pass_point_data : bool, default: False
If ``True``, the original point data will be passed to the output.
Returns
-------
:vtk:`vtkPointDataToCellData`
:vtk:`vtkPointDataToCellData` algorithm.
"""
alg = _vtk.vtkPointDataToCellData()
alg.SetPassPointData(pass_point_data)
set_algorithm_input(alg, inp)
return alg
def triangulate_algorithm(inp):
"""Triangulate the input data.
Parameters
----------
inp : :vtk:`vtkDataObject`
The input data to be triangulated.
Returns
-------
:vtk:`vtkTriangleFilter`
The triangle filter that has been applied to the input data.
"""
trifilter = _vtk.vtkTriangleFilter()
trifilter.PassVertsOff()
trifilter.PassLinesOff()
set_algorithm_input(trifilter, inp)
return trifilter
def decimation_algorithm(inp, target_reduction):
"""Decimate the input data to the target reduction.
Parameters
----------
inp : :vtk:`vtkDataObject`
The input data to be decimated.
target_reduction : float
The target reduction amount, as a fraction of the original data.
Returns
-------
:vtk:`vtkQuadricDecimation`
The decimation algorithm that has been applied to the input data.
"""
alg = _vtk.vtkQuadricDecimation()
alg.SetTargetReduction(target_reduction)
set_algorithm_input(alg, inp)
return alg
@@ -0,0 +1,131 @@
"""Cubemap utilities."""
from __future__ import annotations
from pathlib import Path
import pyvista
def cubemap(path='', prefix='', ext='.jpg'):
"""Construct a cubemap from 6 images from a directory.
Each of the 6 images must be in the following format:
- <prefix>negx<ext>
- <prefix>negy<ext>
- <prefix>negz<ext>
- <prefix>posx<ext>
- <prefix>posy<ext>
- <prefix>posz<ext>
Prefix may be empty, and extension will default to ``'.jpg'``
For example, if you have 6 images with the skybox2 prefix:
- ``'skybox2-negx.jpg'``
- ``'skybox2-negy.jpg'``
- ``'skybox2-negz.jpg'``
- ``'skybox2-posx.jpg'``
- ``'skybox2-posy.jpg'``
- ``'skybox2-posz.jpg'``
Parameters
----------
path : str, default: ""
Directory containing the cubemap images.
prefix : str, default: ""
Prefix to the filename.
ext : str, default: ".jpg"
The filename extension. For example ``'.jpg'``.
Returns
-------
pyvista.Texture
Texture with cubemap.
Notes
-----
Cubemap will appear flipped relative to the XY plane between VTK v9.1 and
VTK v9.2.
Examples
--------
Load a skybox given a directory, prefix, and file extension.
>>> import pyvista as pv
>>> skybox = pv.cubemap('my_directory', 'skybox', '.jpeg') # doctest:+SKIP
"""
sets = ['posx', 'negx', 'posy', 'negy', 'posz', 'negz']
image_paths = [str(Path(path) / f'{prefix}{suffix}{ext}') for suffix in sets]
return _cubemap_from_paths(image_paths)
def cubemap_from_filenames(image_paths):
"""Construct a cubemap from 6 images.
Images must be in the following order:
- Positive X
- Negative X
- Positive Y
- Negative Y
- Positive Z
- Negative Z
Parameters
----------
image_paths : sequence[str]
Paths of the individual cubemap images.
Returns
-------
pyvista.Texture
Texture with cubemap.
Examples
--------
Load a skybox given a list of image paths.
>>> image_paths = [
... '/home/user/_px.jpg',
... '/home/user/_nx.jpg',
... '/home/user/_py.jpg',
... '/home/user/_ny.jpg',
... '/home/user/_pz.jpg',
... '/home/user/_nz.jpg',
... ]
>>> skybox = pv.cubemap(image_paths=image_paths) # doctest:+SKIP
"""
if len(image_paths) != 6:
msg = 'image_paths must contain 6 paths'
raise ValueError(msg)
return _cubemap_from_paths(image_paths)
def _cubemap_from_paths(image_paths):
"""Construct a cubemap from image paths."""
for image_path in image_paths:
if not Path(image_path).is_file():
file_str = '\n'.join(image_paths)
msg = (
f'Unable to locate {image_path}\nExpected to find the following files:\n{file_str}'
)
raise FileNotFoundError(msg)
texture = pyvista.Texture() # type: ignore[abstract]
texture.SetMipmap(True)
texture.SetInterpolate(True)
texture.cube_map = True # Must be set prior to setting images
# add each image to the cubemap
for i, fn in enumerate(image_paths):
# Read and flip along y-axis
texture.SetInputDataObject(i, pyvista.read(fn)._flip_uniform(1))
return texture
@@ -0,0 +1,64 @@
"""Plotting GL checks."""
from __future__ import annotations
from pyvista.plotting import _vtk
def check_depth_peeling(number_of_peels=100, occlusion_ratio=0.0):
"""Check if depth peeling is available.
Attempts to use depth peeling to see if it is available for the
current environment. Returns ``True`` if depth peeling is
available and has been successfully leveraged, otherwise
``False``.
Parameters
----------
number_of_peels : int, default: 100
Maximum number of depth peels.
occlusion_ratio : float, default: 0.0
Occlusion ratio.
Returns
-------
bool
``True`` when system supports depth peeling with the specified
settings.
"""
# Try Depth Peeling with a basic scene
source = _vtk.vtkSphereSource()
mapper = _vtk.vtkPolyDataMapper()
mapper.SetInputConnection(source.GetOutputPort())
actor = _vtk.vtkActor()
actor.SetMapper(mapper)
# requires opacity < 1
actor.GetProperty().SetOpacity(0.5)
renderer = _vtk.vtkRenderer()
renderWindow = _vtk.vtkRenderWindow()
renderWindow.AddRenderer(renderer)
renderWindow.SetOffScreenRendering(True)
renderWindow.SetAlphaBitPlanes(True)
renderWindow.SetMultiSamples(0)
renderer.AddActor(actor)
renderer.SetUseDepthPeeling(True)
renderer.SetMaximumNumberOfPeels(number_of_peels)
renderer.SetOcclusionRatio(occlusion_ratio)
renderWindow.Render()
return renderer.GetLastRenderingUsedDepthPeeling() == 1
def uses_egl() -> bool:
"""Check if VTK has been compiled with EGL support via OSMesa.
Returns
-------
bool
``True`` if VTK has been compiled with EGL support via OSMesa,
otherwise ``False``.
"""
ren_win_str = str(type(_vtk.vtkRenderWindow()))
return 'EGL' in ren_win_str or 'OSOpenGL' in ren_win_str
@@ -0,0 +1,272 @@
"""Image regression module."""
from __future__ import annotations
from typing import TYPE_CHECKING
from typing import cast
import numpy as np
import pyvista
from pyvista._deprecate_positional_args import _deprecate_positional_args
from pyvista.core.utilities.arrays import point_array
from pyvista.core.utilities.helpers import wrap
from pyvista.plotting import _vtk
if TYPE_CHECKING:
from pyvista import ImageData
from pyvista.core._typing_core import NumpyArray
def remove_alpha(img):
"""Remove the alpha channel from a :vtk:`vtkImageData`.
Parameters
----------
img : :vtk:`vtkImageData`
The input image data with an alpha channel.
Returns
-------
ImageData
The output image data with the alpha channel removed.
"""
ec = _vtk.vtkImageExtractComponents()
ec.SetComponents(0, 1, 2)
ec.SetInputData(img)
ec.Update()
return pyvista.wrap(ec.GetOutput())
def wrap_image_array(arr):
"""Wrap a numpy array as a pyvista.ImageData.
Parameters
----------
arr : np.ndarray
A numpy array of shape (X, Y, (3 or 4)) and dtype ``np.uint8``. For
example, an array of shape ``(768, 1024, 3)``.
Raises
------
ValueError
If the input array does not have 3 dimensions, the third dimension of
the input array is not 3 or 4, or the input array is not of type
``np.uint8``.
Returns
-------
pyvista.ImageData
A PyVista ImageData object with the wrapped array data.
"""
if arr.ndim != 3:
msg = 'Expecting a X by Y by (3 or 4) array'
raise ValueError(msg)
if arr.shape[2] not in [3, 4]:
msg = 'Expecting a X by Y by (3 or 4) array'
raise ValueError(msg)
if arr.dtype != np.uint8:
msg = 'Expecting a np.uint8 array'
raise ValueError(msg)
img = _vtk.vtkImageData()
img.SetDimensions(arr.shape[1], arr.shape[0], 1)
wrap_img = pyvista.wrap(img)
wrap_img.point_data['PNGImage'] = arr[::-1].reshape(-1, arr.shape[2])
return wrap_img
def run_image_filter(imfilter: _vtk.vtkWindowToImageFilter) -> NumpyArray[float]:
"""Run a :vtk:`vtkWindowToImageFilter` and get output as array.
Parameters
----------
imfilter : :vtk:`vtkWindowToImageFilter`
The :vtk:`vtkWindowToImageFilter` instance to be processed.
Notes
-----
An empty array will be returned if an image cannot be extracted.
Returns
-------
numpy.ndarray
An array containing the filtered image data. The shape of the array
is given by (height, width, -1) where height and width are the
dimensions of the image.
"""
# Update filter and grab pixels
imfilter.Modified()
imfilter.Update()
image = cast('ImageData | None', wrap(imfilter.GetOutput()))
if image is None:
return np.empty((0, 0, 0))
img_size = image.dimensions
img_array = cast('NumpyArray[float]', point_array(image, 'ImageScalars'))
# Reshape and write
tgt_size = (img_size[1], img_size[0], -1)
return img_array.reshape(tgt_size)[::-1]
@_deprecate_positional_args(allowed=['render_window'])
def image_from_window( # noqa: PLR0917
render_window,
as_vtk: bool = False, # noqa: FBT001, FBT002
ignore_alpha: bool = False, # noqa: FBT001, FBT002
scale=1,
):
"""Extract the image from the render window as an array.
Parameters
----------
render_window : :vtk:`vtkRenderWindow`
The render window to extract the image from.
as_vtk : bool, default: False
If set to True, the image will be returned as a VTK object.
ignore_alpha : bool, default: False
If set to True, the image will be returned in RGB format,
otherwise, it will be returned in RGBA format.
scale : int, default: 1
The scaling factor of the extracted image. The default value is 1
which means that no scaling is applied.
Returns
-------
ndarray | :vtk:`vtkImageData`
The image as an array or as a VTK object depending on the ``as_vtk`` parameter.
"""
off = not render_window.GetInteractor().GetEnableRender()
if off:
render_window.GetInteractor().EnableRenderOn()
imfilter = _vtk.vtkWindowToImageFilter()
imfilter.SetInput(render_window)
imfilter.SetScale(scale)
imfilter.FixBoundaryOn()
imfilter.ReadFrontBufferOff()
imfilter.ShouldRerenderOff()
if ignore_alpha:
imfilter.SetInputBufferTypeToRGB()
else:
imfilter.SetInputBufferTypeToRGBA()
imfilter.ReadFrontBufferOn()
data = run_image_filter(imfilter)
if off:
# Critical for Trame and other offscreen tools
render_window.GetInteractor().EnableRenderOff()
if as_vtk:
return wrap_image_array(data)
return data
@_deprecate_positional_args(allowed=['im1', 'im2'])
def compare_images( # noqa: PLR0917
im1,
im2,
threshold=1,
use_vtk: bool = True, # noqa: FBT001, FBT002
):
"""Compare two different images of the same size.
Parameters
----------
im1 : str | numpy.ndarray | :vtk:`vtkRenderWindow` | :vtk:`vtkImageData`
Render window, numpy array representing the output of a render
window, or :vtk:`vtkImageData`.
im2 : str | numpy.ndarray | :vtk:`vtkRenderWindow` | :vtk:`vtkImageData`
Render window, numpy array representing the output of a render
window, or :vtk:`vtkImageData`.
threshold : int, default: 1
Threshold tolerance for pixel differences. This should be
greater than 0, otherwise it will always return an error, even
on identical images.
use_vtk : bool, default: True
When disabled, computes the mean pixel error over the entire
image using numpy. The difference between pixel is calculated
for each RGB channel, summed, and then divided by the number
of pixels. This is faster than using
:vtk:`vtkImageDifference` but potentially less accurate.
Returns
-------
float
Total error between the images if using ``use_vtk=True``, and
the mean pixel error when ``use_vtk=False``.
Examples
--------
Compare two active plotters.
>>> import pyvista as pv
>>> pl1 = pv.Plotter()
>>> _ = pl1.add_mesh(pv.Sphere(), smooth_shading=True)
>>> pl2 = pv.Plotter()
>>> _ = pl2.add_mesh(pv.Sphere(), smooth_shading=False)
>>> error = pv.compare_images(pl1, pl2)
Compare images from file.
>>> import pyvista as pv
>>> img1 = pv.read('img1.png') # doctest:+SKIP
>>> img2 = pv.read('img2.png') # doctest:+SKIP
>>> pv.compare_images(img1, img2) # doctest:+SKIP
"""
from pyvista import ImageData # noqa: PLC0415
from pyvista import Plotter # noqa: PLC0415
from pyvista import read # noqa: PLC0415
from pyvista import wrap # noqa: PLC0415
def to_img(img):
if isinstance(img, ImageData): # pragma: no cover
return img
elif isinstance(img, _vtk.vtkImageData):
return wrap(img)
elif isinstance(img, str):
return read(img)
elif isinstance(img, np.ndarray):
return wrap_image_array(img)
elif isinstance(img, Plotter):
if img._first_time: # must be rendered first else segfault
img._on_first_render_request()
img.render()
if img.render_window is None:
msg = 'Unable to extract image from Plotter as it has already been closed.'
raise RuntimeError(msg)
return image_from_window(img.render_window, as_vtk=True, ignore_alpha=True)
else:
msg = (
f'Unsupported data type {type(img)}. Should be '
'Either a np.ndarray, vtkRenderWindow, or vtkImageData'
)
raise TypeError(msg)
im1 = remove_alpha(to_img(im1))
im2 = remove_alpha(to_img(im2))
if im1.GetDimensions() != im2.GetDimensions():
msg = 'Input images are not the same size.'
raise RuntimeError(msg)
if use_vtk:
img_diff = _vtk.vtkImageDifference()
img_diff.SetThreshold(threshold)
img_diff.SetInputData(im1)
img_diff.SetImageData(im2)
img_diff.AllowShiftOff() # vastly increases compute time when enabled
# img_diff.AveragingOff() # increases compute time
img_diff.Update()
return img_diff.GetThresholdedError()
# otherwise, simply compute the mean pixel difference
diff = np.abs(im1.point_data[0] - im2.point_data[0])
return np.sum(diff) / im1.point_data[0].shape[0]
@@ -0,0 +1,229 @@
"""Utilities for using pyvista with sphinx-gallery."""
from __future__ import annotations
from pathlib import Path
import shutil
from typing import TYPE_CHECKING
import pyvista
from pyvista._deprecate_positional_args import _deprecate_positional_args
if TYPE_CHECKING:
from collections.abc import Iterator
BUILDING_GALLERY_ERROR_MSG = (
'pyvista.BUILDING_GALLERY must be set to True in your conf.py to capture '
'images within sphinx_gallery or when building documentation using the '
'pyvista-plot directive.'
)
def _get_sg_image_scraper():
"""Return the callable scraper to be used by Sphinx-Gallery.
It allows PyVista users to just use strings as they already can for
'matplotlib' and 'mayavi'. Details on this implementation can be found in
`sphinx-gallery/sphinx-gallery/494`_
This must be imported into the top level namespace of PyVista.
.. _sphinx-gallery/sphinx-gallery/494: https://github.com/sphinx-gallery/sphinx-gallery/pull/494
"""
return Scraper()
def html_rst(
figure_list,
sources_dir,
srcsetpaths=None,
): # pragma: no cover # numpydoc ignore=PR01,RT01
"""Generate reST for viewer with exported scene."""
from sphinx_gallery.scrapers import _get_srcset_st # noqa: PLC0415
from sphinx_gallery.scrapers import figure_rst # noqa: PLC0415
if srcsetpaths is None:
# this should never happen, but figure_rst is public, so
# this has to be a kwarg...
srcsetpaths = [{0: fl} for fl in figure_list]
images_rst = ''
for i, hinnames in enumerate(srcsetpaths):
srcset = _get_srcset_st(sources_dir, hinnames)
if srcset[-5:] == 'vtksz':
png_file = figure_list[i][:-5] + 'png'
indented_firgure_rst = '\n'.join(
' ' * 5 + line for line in figure_rst([png_file], sources_dir).split('\n')
)
images_rst += f"""
\n
\n
.. tab-set::\n
\n
.. tab-item:: Static Scene\n
\n
{indented_firgure_rst}
\n
.. tab-item:: Interactive Scene\n
\n
.. offlineviewer:: {figure_list[i]}\n\n"""
else:
images_rst += '\n' + figure_rst([figure_list[i]], sources_dir) + '\n\n'
return images_rst
def _process_events_before_scraping(plotter):
"""Process events such as changing the camera or an object before scraping."""
if plotter.iren is not None and plotter.iren.initialized:
# check for pyvistaqt app which can be specifically bound to pyvista plotter
# objects in order to interact with qt, then process the events from qt
if hasattr(plotter, 'app') and plotter.app is not None:
plotter.app.processEvents()
plotter.update()
@_deprecate_positional_args(allowed=['image_path_iterator'])
def generate_images(image_path_iterator: Iterator[str], dynamic: bool = False) -> list[str]: # noqa: FBT001, FBT002
"""Generate images from the current plotters.
The file names are taken from the ``image_path_iterator`` iterator.
A gif will be created if a plotter has a ``_gif_filename`` attribute.
Otherwise, depending on the value of ``dynamic``, either a ``.png`` static image
or a ``.vtksz`` file will be created.
Parameters
----------
image_path_iterator : Iterator[str]
An iterator that yields the path to the next image to be saved.
dynamic : bool, default: False
Whether to save a static ``.png`` image or a ``.vtksz`` (interactive)
file.
Returns
-------
list[str]
A list of the names of the images that were created.
"""
image_names = []
figures = pyvista.plotting.plotter._ALL_PLOTTERS
for plotter in figures.values():
_process_events_before_scraping(plotter)
fname = next(image_path_iterator)
# Make sure the extension is "png"
path = Path(fname)
fname_withoutextension = str(path.parent / path.stem)
fname = fname_withoutextension + '.png'
if (gif_filename := plotter._gif_filename) is not None:
# move gif to fname
fname = fname[:-3] + 'gif'
shutil.move(gif_filename, fname)
image_names.append(fname)
else:
plotter.screenshot(fname)
if not dynamic or plotter.last_vtksz is None:
image_names.append(fname)
else: # pragma: no cover
fname = fname[:-3] + 'vtksz'
with Path(fname).open('wb') as f:
f.write(plotter.last_vtksz) # type: ignore[arg-type]
image_names.append(fname)
pyvista.close_all() # close and clear all plotters
return image_names
class Scraper:
"""Save ``pyvista.Plotter`` objects.
Used by sphinx-gallery to generate the plots from the code in the examples.
Pass an instance of this class to ``sphinx_gallery_conf`` in your
``conf.py`` as the ``"image_scrapers"`` argument.
Be sure to set ``pyvista.BUILDING_GALLERY = True`` in your ``conf.py``.
"""
def __repr__(self) -> str:
"""Return a stable representation of the class instance."""
return f'<{type(self).__name__} object>'
def __call__(self, block, block_vars, gallery_conf): # noqa: ARG002
"""Save the figures generated after running example code.
Called by sphinx-gallery.
"""
from sphinx_gallery.scrapers import figure_rst # noqa: PLC0415
if not pyvista.BUILDING_GALLERY:
raise RuntimeError(BUILDING_GALLERY_ERROR_MSG)
image_path_iterator = block_vars['image_path_iterator']
image_names = generate_images(image_path_iterator, dynamic=False)
return figure_rst(image_names, gallery_conf['src_dir'])
class DynamicScraper: # pragma: no cover
"""Save ``pyvista.Plotter`` objects dynamically.
Used by sphinx-gallery to generate the plots from the code in the examples.
Pass an instance of this class to ``sphinx_gallery_conf`` in your
``conf.py`` as the ``"image_scrapers"`` argument.
Be sure to set ``pyvista.BUILDING_GALLERY = True`` in your ``conf.py``.
If the boolean variable ``PYVISTA_GALLERY_FORCE_STATIC_IN_DOCUMENT = True/False``
is set as a global variable in the document then its value will be used as default for the
force_static argument of the pyvista-plot command. see also the notes at :func:plot_directive
To alter the global value behavior just for some plots you may set the
boolean variable ``PYVISTA_GALLERY_FORCE_STATIC = True``/
``PYVISTA_GALLERY_FORCE_STATIC = False`` just before the appropriate ``plot`` command.
The default behavior of this scraper is to create interactive plots.
"""
def __repr__(self) -> str:
"""Return a stable representation of the class instance."""
return f'<{type(self).__name__} object>'
def __call__(self, block, block_vars, gallery_conf): # pragma: no cover
"""Save the figures generated after running example code.
Called by sphinx-gallery.
"""
if not pyvista.BUILDING_GALLERY:
raise RuntimeError(BUILDING_GALLERY_ERROR_MSG)
# read global option if it exists
force_static = block_vars['example_globals'].get(
'PYVISTA_GALLERY_FORCE_STATIC_IN_DOCUMENT',
False,
)
# override with block specific value if it exists
if 'PYVISTA_GALLERY_FORCE_STATIC = True' in block[1].split('\n'):
force_static = True
elif 'PYVISTA_GALLERY_FORCE_STATIC = False' in block[1].split('\n'):
force_static = False
if force_static is None:
# Just in case force_static is None at this point
force_static = False
dynamic = not force_static
image_path_iterator = block_vars['image_path_iterator']
image_names = generate_images(image_path_iterator, dynamic=dynamic)
return html_rst(image_names, gallery_conf['src_dir'])
@@ -0,0 +1,71 @@
"""Start xvfb from Python."""
from __future__ import annotations
import os
import time
import warnings
from pyvista.core.errors import PyVistaDeprecationWarning
XVFB_INSTALL_NOTES = """Please install Xvfb with:
Debian
$ sudo apt install libgl1-mesa-glx xvfb
CentOS / RHL
$ sudo yum install libgl1-mesa-glx xvfb
"""
def start_xvfb(wait=3, window_size=None):
"""Start the virtual framebuffer Xvfb.
Parameters
----------
wait : float, optional
Time to wait for the virtual framebuffer to start. Set to 0
to disable wait.
window_size : list, optional
Window size of the virtual frame buffer. Defaults to
:attr:`pyvista.global_theme.window_size
<pyvista.plotting.themes.Theme.window_size>`.
Notes
-----
Only available on Linux. Be sure to install ``xvfb``
and ``libgl1-mesa-glx`` in your package manager.
Examples
--------
>>> import pyvista as pv
>>> pv.start_xvfb() # doctest:+SKIP
"""
# Deprecated on 0.45.0, estimated removal on 0.48.0
warnings.warn(
'This function is deprecated and will be removed in future version of '
'PyVista. Use vtk-osmesa instead.',
PyVistaDeprecationWarning,
)
from pyvista import global_theme # noqa: PLC0415
if os.name != 'posix':
msg = '`start_xvfb` is only supported on Linux'
raise OSError(msg)
if os.system('which Xvfb > /dev/null'):
raise OSError(XVFB_INSTALL_NOTES)
# use current default window size
if window_size is None:
window_size = global_theme.window_size
window_size_parm = f'{window_size[0]:d}x{window_size[1]:d}x24'
display_num = ':99'
os.system(f'Xvfb {display_num} -screen 0 {window_size_parm} > /dev/null 2>&1 &')
os.environ['DISPLAY'] = display_num
if wait:
time.sleep(wait)