mediapy

mediapy: Read/write/show images and videos in an IPython/Jupyter notebook.

[GitHub source]   [API docs]   [PyPI package]   [Colab example]

See the example notebook, or better yet, open it in Colab.

Image examples

Display an image (2D or 3D numpy array):

checkerboard = np.kron([[0, 1] * 16, [1, 0] * 16] * 16, np.ones((4, 4)))
show_image(checkerboard)

Read and display an image (either local or from the Web):

IMAGE = 'https://github.com/hhoppe/data/raw/main/image.png'
show_image(read_image(IMAGE))

Read and display an image from a local file:

!wget -q -O /tmp/burano.png {IMAGE}
show_image(read_image('/tmp/burano.png'))

Show titled images side-by-side:

images = {
    'original': checkerboard,
    'darkened': checkerboard * 0.7,
    'random': np.random.rand(32, 32, 3),
}
show_images(images, vmin=0.0, vmax=1.0, border=True, height=64)

Compare two images using an interactive slider:

compare_images([checkerboard, np.random.rand(128, 128, 3)])

Video examples

Display a video (an iterable of images, e.g., a 3D or 4D array):

video = moving_circle((100, 100), num_images=10)
show_video(video, fps=10)

Show the video frames side-by-side:

show_images(video, columns=6, border=True, height=64)

Show the frames with their indices:

show_images({f'{i}': image for i, image in enumerate(video)}, width=32)

Read and display a video (either local or from the Web):

VIDEO = 'https://github.com/hhoppe/data/raw/main/video.mp4'
show_video(read_video(VIDEO))

Create and display a looping two-frame GIF video:

image1 = resize_image(np.random.rand(10, 10, 3), (50, 50))
show_video([image1, image1 * 0.8], fps=2, codec='gif')

Darken a video frame-by-frame:

output_path = '/tmp/out.mp4'
with VideoReader(VIDEO) as r:
  darken_image = lambda image: to_float01(image) * 0.5
  with VideoWriter(output_path, shape=r.shape, fps=r.fps, bps=r.bps) as w:
    for image in r:
      w.add_image(darken_image(image))
   1# Copyright 2026 The mediapy Authors.
   2#
   3# Licensed under the Apache License, Version 2.0 (the "License");
   4# you may not use this file except in compliance with the License.
   5# You may obtain a copy of the License at
   6#
   7#     http://www.apache.org/licenses/LICENSE-2.0
   8#
   9# Unless required by applicable law or agreed to in writing, software
  10# distributed under the License is distributed on an "AS IS" BASIS,
  11# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
  12# See the License for the specific language governing permissions and
  13# limitations under the License.
  14
  15"""`mediapy`: Read/write/show images and videos in an IPython/Jupyter notebook.
  16
  17[**[GitHub source]**](https://github.com/google/mediapy)  
  18[**[API docs]**](https://google.github.io/mediapy/)  
  19[**[PyPI package]**](https://pypi.org/project/mediapy/)  
  20[**[Colab
  21example]**](https://colab.research.google.com/github/google/mediapy/blob/main/mediapy_examples.ipynb)
  22
  23See the [example
  24notebook](https://github.com/google/mediapy/blob/main/mediapy_examples.ipynb),
  25or better yet, [**open it in
  26Colab**](https://colab.research.google.com/github/google/mediapy/blob/main/mediapy_examples.ipynb).
  27
  28## Image examples
  29
  30Display an image (2D or 3D `numpy` array):
  31```python
  32checkerboard = np.kron([[0, 1] * 16, [1, 0] * 16] * 16, np.ones((4, 4)))
  33show_image(checkerboard)
  34```
  35
  36Read and display an image (either local or from the Web):
  37```python
  38IMAGE = 'https://github.com/hhoppe/data/raw/main/image.png'
  39show_image(read_image(IMAGE))
  40```
  41
  42Read and display an image from a local file:
  43```python
  44!wget -q -O /tmp/burano.png {IMAGE}
  45show_image(read_image('/tmp/burano.png'))
  46```
  47
  48Show titled images side-by-side:
  49```python
  50images = {
  51    'original': checkerboard,
  52    'darkened': checkerboard * 0.7,
  53    'random': np.random.rand(32, 32, 3),
  54}
  55show_images(images, vmin=0.0, vmax=1.0, border=True, height=64)
  56```
  57
  58Compare two images using an interactive slider:
  59```python
  60compare_images([checkerboard, np.random.rand(128, 128, 3)])
  61```
  62
  63## Video examples
  64
  65Display a video (an iterable of images, e.g., a 3D or 4D array):
  66```python
  67video = moving_circle((100, 100), num_images=10)
  68show_video(video, fps=10)
  69```
  70
  71Show the video frames side-by-side:
  72```python
  73show_images(video, columns=6, border=True, height=64)
  74```
  75
  76Show the frames with their indices:
  77```python
  78show_images({f'{i}': image for i, image in enumerate(video)}, width=32)
  79```
  80
  81Read and display a video (either local or from the Web):
  82```python
  83VIDEO = 'https://github.com/hhoppe/data/raw/main/video.mp4'
  84show_video(read_video(VIDEO))
  85```
  86
  87Create and display a looping two-frame GIF video:
  88```python
  89image1 = resize_image(np.random.rand(10, 10, 3), (50, 50))
  90show_video([image1, image1 * 0.8], fps=2, codec='gif')
  91```
  92
  93Darken a video frame-by-frame:
  94```python
  95output_path = '/tmp/out.mp4'
  96with VideoReader(VIDEO) as r:
  97  darken_image = lambda image: to_float01(image) * 0.5
  98  with VideoWriter(output_path, shape=r.shape, fps=r.fps, bps=r.bps) as w:
  99    for image in r:
 100      w.add_image(darken_image(image))
 101```
 102"""
 103
 104from __future__ import annotations
 105
 106__docformat__ = 'google'
 107__version__ = '1.2.7'
 108__version_info__ = tuple(int(num) for num in __version__.split('.'))
 109
 110import base64
 111from collections.abc import Callable, Iterable, Iterator, Mapping, Sequence
 112import contextlib
 113import functools
 114import importlib
 115import io
 116import itertools
 117import math
 118import numbers
 119import os  # Package only needed for typing.TYPE_CHECKING.
 120import pathlib
 121import re
 122import shlex
 123import shutil
 124import subprocess
 125import sys
 126import tempfile
 127import typing
 128from typing import Any
 129import urllib.request
 130import warnings
 131
 132import IPython.display
 133import matplotlib.pyplot
 134import numpy as np
 135import numpy.typing as npt
 136import PIL.Image
 137import PIL.ImageOps
 138
 139
 140if not hasattr(PIL.Image, 'Resampling'):  # Allow Pillow<9.0.
 141  PIL.Image.Resampling = PIL.Image  # type: ignore
 142
 143# Selected and reordered here for pdoc documentation.
 144__all__ = [
 145    'show_image',
 146    'show_images',
 147    'compare_images',
 148    'show_video',
 149    'show_videos',
 150    'read_image',
 151    'write_image',
 152    'read_video',
 153    'write_video',
 154    'VideoReader',
 155    'VideoWriter',
 156    'VideoMetadata',
 157    'compress_image',
 158    'decompress_image',
 159    'compress_video',
 160    'decompress_video',
 161    'html_from_compressed_image',
 162    'html_from_compressed_video',
 163    'resize_image',
 164    'resize_video',
 165    'to_rgb',
 166    'to_type',
 167    'to_float01',
 168    'to_uint8',
 169    'set_output_height',
 170    'set_max_output_height',
 171    'color_ramp',
 172    'moving_circle',
 173    'set_show_save_dir',
 174    'set_ffmpeg',
 175    'video_is_available',
 176]
 177
 178if TYPE_CHECKING:
 179  _ArrayLike = npt.ArrayLike
 180  _DTypeLike = npt.DTypeLike
 181  _NDArray = npt.NDArray[Any]
 182  _DType = np.dtype[Any]
 183else:
 184  # Create named types for use in the `pdoc` documentation.
 185  _ArrayLike = TypeVar('_ArrayLike')
 186  _DTypeLike = TypeVar('_DTypeLike')
 187  _NDArray = TypeVar('_NDArray')
 188  _DType = TypeVar('_DType')  # pylint: disable=invalid-name
 189
 190_IPYTHON_HTML_SIZE_LIMIT = 10**10  # Unlimited seems to be OK now.
 191_T = TypeVar('_T')
 192_Path = Union[str, 'os.PathLike[str]']
 193
 194# Raw PCM sample formats (as in `ffmpeg -formats`) for the audio dtypes that
 195# `VideoWriter` is able to feed to `ffmpeg`.  Multi-byte formats lack the
 196# byte-order suffix ('le' or 'be'), which is appended for the native order.
 197_FFMPEG_AUDIO_FORMAT_FROM_DTYPE = {
 198    np.uint8: 'u8',
 199    np.int16: 's16',
 200    np.int32: 's32',
 201    np.float32: 'f32',
 202    np.float64: 'f64',
 203}
 204
 205_IMAGE_COMPARISON_HTML = """\
 206<script
 207  defer
 208  src="https://unpkg.com/img-comparison-slider@7/dist/index.js"
 209></script>
 210<link
 211  rel="stylesheet"
 212  href="https://unpkg.com/img-comparison-slider@7/dist/styles.css"
 213/>
 214
 215<img-comparison-slider>
 216  <img slot="first" src="data:image/png;base64,{b64_1}" />
 217  <img slot="second" src="data:image/png;base64,{b64_2}" />
 218</img-comparison-slider>
 219"""
 220
 221# ** Miscellaneous.
 222
 223
 224class _Config:
 225  ffmpeg_name_or_path: _Path = 'ffmpeg'
 226  show_save_dir: _Path | None = None
 227
 228
 229_config = _Config()
 230
 231
 232def _open(path: _Path, *args: Any, **kwargs: Any) -> Any:
 233  """Opens the file; this is a hook for the built-in `open()`."""
 234  return open(path, *args, **kwargs)
 235
 236
 237def _path_is_local(path: _Path) -> bool:
 238  """Returns True if the path is in the filesystem accessible by `ffmpeg`."""
 239  del path
 240  return True
 241
 242
 243def _search_for_ffmpeg_path() -> str | None:
 244  """Returns a path to the ffmpeg program, or None if not found."""
 245  if filename := shutil.which(_config.ffmpeg_name_or_path):
 246    return str(filename)
 247  return None
 248
 249
 250def _print_err(*args: str, **kwargs: Any) -> None:
 251  """Prints arguments to stderr immediately."""
 252  kwargs = {**dict(file=sys.stderr, flush=True), **kwargs}
 253  print(*args, **kwargs)
 254
 255
 256def _chunked(
 257    iterable: Iterable[_T], n: int | None = None
 258) -> Iterator[tuple[_T, ...]]:
 259  """Returns elements collected as tuples of length at most `n` if not None."""
 260
 261  def take(n: int | None, iterable: Iterable[_T]) -> tuple[_T, ...]:
 262    return tuple(itertools.islice(iterable, n))
 263
 264  return iter(functools.partial(take, n, iter(iterable)), ())
 265
 266
 267def _peek_first(iterator: Iterable[_T]) -> tuple[_T, Iterable[_T]]:
 268  """Given an iterator, returns first element and re-initialized iterator.
 269
 270  >>> first_image, images = _peek_first(moving_circle())
 271
 272  Args:
 273    iterator: An input iterator or iterable.
 274
 275  Returns:
 276    A tuple (first_element, iterator_reinitialized) containing:
 277      first_element: The first element of the input.
 278      iterator_reinitialized: A clone of the original iterator/iterable.
 279  """
 280  # Inspired from https://stackoverflow.com/a/12059829/1190077
 281  peeker, iterator_reinitialized = itertools.tee(iterator)
 282  first = next(peeker)
 283  return first, iterator_reinitialized
 284
 285
 286def _check_2d_shape(shape: tuple[int, int]) -> None:
 287  """Checks that `shape` is of the form (height, width) with two integers."""
 288  if len(shape) != 2:
 289    raise ValueError(f'Shape {shape} is not of the form (height, width).')
 290  if not all(isinstance(i, numbers.Integral) for i in shape):
 291    raise ValueError(f'Shape {shape} contains non-integers.')
 292
 293
 294def _run(args: str | Sequence[str]) -> None:
 295  """Executes command, printing output from stdout and stderr.
 296
 297  Args:
 298    args: Command to execute, which can be either a string or a sequence of word
 299      strings, as in `subprocess.run()`.  If `args` is a string, the shell is
 300      invoked to interpret it.
 301
 302  Raises:
 303    RuntimeError: If the command's exit code is nonzero.
 304  """
 305  proc = subprocess.run(
 306      args,
 307      shell=isinstance(args, str),
 308      stdout=subprocess.PIPE,
 309      stderr=subprocess.STDOUT,
 310      check=False,
 311      universal_newlines=True,
 312  )
 313  print(proc.stdout, end='', flush=True)
 314  if proc.returncode:
 315    raise RuntimeError(
 316        f"Command '{proc.args}' failed with code {proc.returncode}."
 317    )
 318
 319
 320def _display_html(text: str, /) -> None:
 321  """In a Jupyter notebook, display the HTML `text`."""
 322  IPython.display.display(IPython.display.HTML(text))  # type: ignore
 323
 324
 325def set_ffmpeg(name_or_path: _Path) -> None:
 326  """Specifies the name or path for the `ffmpeg` external program.
 327
 328  The `ffmpeg` program is required for compressing and decompressing video.
 329  (It is used in `read_video`, `write_video`, `show_video`, `show_videos`,
 330  etc.)
 331
 332  Args:
 333    name_or_path: Either a filename within a directory of `os.environ['PATH']`
 334      or a filepath.  The default setting is 'ffmpeg'.
 335  """
 336  _config.ffmpeg_name_or_path = name_or_path
 337
 338
 339def set_output_height(num_pixels: int) -> None:
 340  """Overrides the height of the current output cell, if using Colab."""
 341  try:
 342    # We want to fail gracefully for non-Colab IPython notebooks.
 343    output = importlib.import_module('google.colab.output')
 344    s = f'google.colab.output.setIframeHeight("{num_pixels}px")'
 345    output.eval_js(s)
 346  except (ModuleNotFoundError, AttributeError):
 347    pass
 348
 349
 350def set_max_output_height(num_pixels: int) -> None:
 351  """Sets the maximum height of the current output cell, if using Colab."""
 352  try:
 353    # We want to fail gracefully for non-Colab IPython notebooks.
 354    output = importlib.import_module('google.colab.output')
 355    s = (
 356        'google.colab.output.setIframeHeight('
 357        f'0, true, {{maxHeight: {num_pixels}}})'
 358    )
 359    output.eval_js(s)
 360  except (ModuleNotFoundError, AttributeError):
 361    pass
 362
 363
 364# ** Type conversions.
 365
 366
 367def _as_valid_media_type(dtype: _DTypeLike) -> _DType:
 368  """Returns validated media data type."""
 369  dtype = np.dtype(dtype)
 370  if not issubclass(dtype.type, (np.unsignedinteger, np.floating)):
 371    raise ValueError(
 372        f'Type {dtype} is not a valid media data type (uint or float).'
 373    )
 374  return dtype
 375
 376
 377def _as_valid_media_array(x: _ArrayLike) -> _NDArray:
 378  """Converts to ndarray (if not already), and checks validity of data type."""
 379  a = np.asarray(x)
 380  if a.dtype == bool:
 381    a = a.astype(np.uint8) * np.iinfo(np.uint8).max
 382  _as_valid_media_type(a.dtype)
 383  return a
 384
 385
 386def to_type(array: _ArrayLike, dtype: _DTypeLike) -> _NDArray:
 387  """Returns media array converted to specified type.
 388
 389  A "media array" is one in which the dtype is either a floating-point type
 390  (np.float32 or np.float64) or an unsigned integer type.  The array values are
 391  assumed to lie in the range [0.0, 1.0] for floating-point values, and in the
 392  full range for unsigned integers, e.g. [0, 255] for np.uint8.
 393
 394  Conversion between integers and floats maps uint(0) to 0.0 and uint(MAX) to
 395  1.0.  The input array may also be of type bool, whereby True maps to
 396  uint(MAX) or 1.0.  The values are scaled and clamped as appropriate during
 397  type conversions.
 398
 399  Args:
 400    array: Input array-like object (floating-point, unsigned int, or bool).
 401    dtype: Desired output type (floating-point or unsigned int).
 402
 403  Returns:
 404    Array `a` if it is already of the specified dtype, else a converted array.
 405  """
 406  a = np.asarray(array)
 407  dtype = np.dtype(dtype)
 408  del array
 409  if a.dtype != bool:
 410    _as_valid_media_type(a.dtype)  # Verify that 'a' has a valid dtype.
 411  if a.dtype == bool:
 412    result = a.astype(dtype)
 413    if np.issubdtype(dtype, np.unsignedinteger):
 414      result = result * dtype.type(np.iinfo(dtype).max)  # pyrefly: ignore[no-matching-overload]
 415  elif a.dtype == dtype:
 416    result = a
 417  elif np.issubdtype(dtype, np.unsignedinteger):
 418    if np.issubdtype(a.dtype, np.unsignedinteger):
 419      src_max: float = np.iinfo(a.dtype).max
 420    else:
 421      a = np.clip(a, 0.0, 1.0)
 422      src_max = 1.0
 423    dst_max = np.iinfo(dtype).max  # pyrefly: ignore[no-matching-overload]
 424    if dst_max <= np.iinfo(np.uint16).max:
 425      scale = np.array(dst_max / src_max, dtype=np.float32)
 426      result = (a * scale + 0.5).astype(dtype)
 427    elif dst_max <= np.iinfo(np.uint32).max:
 428      result = (a.astype(np.float64) * (dst_max / src_max) + 0.5).astype(dtype)
 429    else:
 430      # https://stackoverflow.com/a/66306123/
 431      a = a.astype(np.float64) * (dst_max / src_max) + 0.5
 432      dst = np.atleast_1d(a)
 433      values_too_large = dst >= np.float64(dst_max)
 434      with np.errstate(invalid='ignore'):
 435        dst = dst.astype(dtype)
 436      dst[values_too_large] = dst_max
 437      result = dst if a.ndim > 0 else dst[0]
 438  else:
 439    assert np.issubdtype(dtype, np.floating)
 440    result = a.astype(dtype)
 441    if np.issubdtype(a.dtype, np.unsignedinteger):
 442      result = result / dtype.type(np.iinfo(a.dtype).max)
 443  return result
 444
 445
 446def to_float01(a: _ArrayLike, dtype: _DTypeLike = np.float32) -> _NDArray:
 447  """If array has unsigned integers, rescales them to the range [0.0, 1.0].
 448
 449  Scaling is such that uint(0) maps to 0.0 and uint(MAX) maps to 1.0.  See
 450  `to_type`.
 451
 452  Args:
 453    a: Input array.
 454    dtype: Desired floating-point type if rescaling occurs.
 455
 456  Returns:
 457    A new array of dtype values in the range [0.0, 1.0] if the input array `a`
 458    contains unsigned integers; otherwise, array `a` is returned unchanged.
 459  """
 460  a = np.asarray(a)
 461  dtype = np.dtype(dtype)
 462  if not np.issubdtype(dtype, np.floating):
 463    raise ValueError(f'Type {dtype} is not floating-point.')
 464  if np.issubdtype(a.dtype, np.floating):
 465    return a
 466  return to_type(a, dtype)
 467
 468
 469def to_uint8(a: _ArrayLike) -> _NDArray:
 470  """Returns array converted to uint8 values; see `to_type`."""
 471  return to_type(a, np.uint8)
 472
 473
 474# ** Functions to generate example image and video data.
 475
 476
 477def color_ramp(
 478    shape: tuple[int, int] = (64, 64), *, dtype: _DTypeLike = np.float32
 479) -> _NDArray:
 480  """Returns an image of a red-green color gradient.
 481
 482  This is useful for quick experimentation and testing.  See also
 483  `moving_circle` to generate a sample video.
 484
 485  Args:
 486    shape: 2D spatial dimensions (height, width) of generated image.
 487    dtype: Type (uint or floating) of resulting pixel values.
 488  """
 489  _check_2d_shape(shape)
 490  dtype = _as_valid_media_type(dtype)
 491  yx = (np.moveaxis(np.indices(shape), 0, -1) + 0.5) / shape
 492  image = np.insert(yx, 2, 0.0, axis=-1)
 493  return to_type(image, dtype)
 494
 495
 496def moving_circle(
 497    shape: tuple[int, int] = (256, 256),
 498    num_images: int = 10,
 499    *,
 500    dtype: _DTypeLike = np.float32,
 501) -> _NDArray:
 502  """Returns a video of a circle moving in front of a color ramp.
 503
 504  This is useful for quick experimentation and testing.  See also `color_ramp`
 505  to generate a sample image.
 506
 507  >>> show_video(moving_circle((480, 640), 60), fps=60)
 508
 509  Args:
 510    shape: 2D spatial dimensions (height, width) of generated video.
 511    num_images: Number of video frames.
 512    dtype: Type (uint or floating) of resulting pixel values.
 513  """
 514  _check_2d_shape(shape)
 515  dtype = np.dtype(dtype)
 516
 517  def generate_image(image_index: int) -> _NDArray:
 518    """Returns a video frame image."""
 519    image = color_ramp(shape, dtype=dtype)
 520    yx = np.moveaxis(np.indices(shape), 0, -1)
 521    center = shape[0] * 0.6, shape[1] * (image_index + 0.5) / num_images
 522    radius_squared = (min(shape) * 0.1) ** 2
 523    inside = np.sum((yx - center) ** 2, axis=-1) < radius_squared
 524    white_circle_color = 1.0, 1.0, 1.0
 525    if np.issubdtype(dtype, np.unsignedinteger):
 526      white_circle_color = to_type([white_circle_color], dtype)[0]
 527    image[inside] = white_circle_color
 528    return image
 529
 530  return np.array([generate_image(i) for i in range(num_images)])
 531
 532
 533# ** Color-space conversions.
 534
 535# Same matrix values as in two sources:
 536# https://github.com/scikit-image/scikit-image/blob/master/skimage/color/colorconv.py#L377
 537# https://github.com/tensorflow/tensorflow/blob/r1.14/tensorflow/python/ops/image_ops_impl.py#L2754
 538_YUV_FROM_RGB_MATRIX = np.array(
 539    [
 540        [0.299, -0.14714119, 0.61497538],
 541        [0.587, -0.28886916, -0.51496512],
 542        [0.114, 0.43601035, -0.10001026],
 543    ],
 544    dtype=np.float32,
 545)
 546_RGB_FROM_YUV_MATRIX = np.linalg.inv(_YUV_FROM_RGB_MATRIX)
 547_YUV_CHROMA_OFFSET = np.array([0.0, 0.5, 0.5], dtype=np.float32)
 548
 549
 550def yuv_from_rgb(rgb: _ArrayLike) -> _NDArray:
 551  """Returns the RGB image/video mapped to YUV [0,1] color space.
 552
 553  Note that the "YUV" color space used by video compressors is actually YCbCr!
 554
 555  Args:
 556    rgb: Input image in sRGB space.
 557  """
 558  rgb = to_float01(rgb)
 559  if rgb.shape[-1] != 3:
 560    raise ValueError(f'The last dimension in {rgb.shape} is not 3.')
 561  return rgb @ _YUV_FROM_RGB_MATRIX + _YUV_CHROMA_OFFSET
 562
 563
 564def rgb_from_yuv(yuv: _ArrayLike) -> _NDArray:
 565  """Returns the YUV image/video mapped to RGB [0,1] color space."""
 566  yuv = to_float01(yuv)
 567  if yuv.shape[-1] != 3:
 568    raise ValueError(f'The last dimension in {yuv.shape} is not 3.')
 569  return (yuv - _YUV_CHROMA_OFFSET) @ _RGB_FROM_YUV_MATRIX
 570
 571
 572# Same matrix values as in
 573# https://github.com/scikit-image/scikit-image/blob/master/skimage/color/colorconv.py#L1654
 574# and https://en.wikipedia.org/wiki/YUV#Studio_swing_for_BT.601
 575_YCBCR_FROM_RGB_MATRIX = np.array(
 576    [
 577        [65.481, 128.553, 24.966],
 578        [-37.797, -74.203, 112.0],
 579        [112.0, -93.786, -18.214],
 580    ],
 581    dtype=np.float32,
 582).transpose()
 583_RGB_FROM_YCBCR_MATRIX = np.linalg.inv(_YCBCR_FROM_RGB_MATRIX)
 584_YCBCR_OFFSET = np.array([16.0, 128.0, 128.0], dtype=np.float32)
 585# Note that _YCBCR_FROM_RGB_MATRIX =~ _YUV_FROM_RGB_MATRIX * [219, 256, 182];
 586# https://en.wikipedia.org/wiki/YUV: "Y' values are conventionally shifted and
 587# scaled to the range [16, 235] (referred to as studio swing or 'TV levels')";
 588# "studio range of 16-240 for U and V".  (Where does value 182 come from?)
 589
 590
 591def ycbcr_from_rgb(rgb: _ArrayLike) -> _NDArray:
 592  """Returns the RGB image/video mapped to YCbCr [0,1] color space.
 593
 594  The YCbCr color space is the one called "YUV" by video compressors.
 595
 596  Args:
 597    rgb: Input image in sRGB space.
 598  """
 599  rgb = to_float01(rgb)
 600  if rgb.shape[-1] != 3:
 601    raise ValueError(f'The last dimension in {rgb.shape} is not 3.')
 602  return (rgb @ _YCBCR_FROM_RGB_MATRIX + _YCBCR_OFFSET) / 255.0
 603
 604
 605def rgb_from_ycbcr(ycbcr: _ArrayLike) -> _NDArray:
 606  """Returns the YCbCr image/video mapped to RGB [0,1] color space."""
 607  ycbcr = to_float01(ycbcr)
 608  if ycbcr.shape[-1] != 3:
 609    raise ValueError(f'The last dimension in {ycbcr.shape} is not 3.')
 610  return (ycbcr * 255.0 - _YCBCR_OFFSET) @ _RGB_FROM_YCBCR_MATRIX
 611
 612
 613# ** Image processing.
 614
 615
 616def _pil_image(image: _ArrayLike, mode: str | None = None) -> PIL.Image.Image:
 617  """Returns a PIL image given a numpy matrix (either uint8 or float [0,1])."""
 618  image = _as_valid_media_array(image)
 619  if image.ndim not in (2, 3):
 620    raise ValueError(f'Image shape {image.shape} is neither 2D nor 3D.')
 621  pil_image: PIL.Image.Image = PIL.Image.fromarray(image, mode=mode)
 622  return pil_image
 623
 624
 625def resize_image(image: _ArrayLike, shape: tuple[int, int]) -> _NDArray:
 626  """Resizes image to specified spatial dimensions using a Lanczos filter.
 627
 628  Args:
 629    image: Array-like 2D or 3D object, where dtype is uint or floating-point.
 630    shape: 2D spatial dimensions (height, width) of output image.
 631
 632  Returns:
 633    A resampled image whose spatial dimensions match `shape`.
 634  """
 635  image = _as_valid_media_array(image)
 636  if image.ndim not in (2, 3):
 637    raise ValueError(f'Image shape {image.shape} is neither 2D nor 3D.')
 638  _check_2d_shape(shape)
 639
 640  # A PIL image can be multichannel only if it has 3 or 4 uint8 channels,
 641  # and it can be resized only if it is uint8 or float32.
 642  supported_single_channel = (
 643      np.issubdtype(image.dtype, np.floating) or image.dtype == np.uint8
 644  ) and image.ndim == 2
 645  supported_multichannel = (
 646      image.dtype == np.uint8 and image.ndim == 3 and image.shape[2] in (3, 4)
 647  )
 648  if supported_single_channel or supported_multichannel:
 649    return np.array(
 650        _pil_image(image).resize(
 651            shape[::-1], resample=PIL.Image.Resampling.LANCZOS
 652        ),
 653        dtype=image.dtype,
 654    )
 655  if image.ndim == 2:
 656    # We convert to floating-point for resizing and convert back.
 657    return to_type(resize_image(to_float01(image), shape), image.dtype)
 658  # We resize each image channel individually.
 659  return np.dstack(
 660      [resize_image(channel, shape) for channel in np.moveaxis(image, -1, 0)]
 661  )
 662
 663
 664# ** Video processing.
 665
 666
 667def resize_video(video: Iterable[_NDArray], shape: tuple[int, int]) -> _NDArray:
 668  """Resizes `video` to specified spatial dimensions using a Lanczos filter.
 669
 670  Args:
 671    video: Iterable of images.
 672    shape: 2D spatial dimensions (height, width) of output video.
 673
 674  Returns:
 675    A resampled video whose spatial dimensions match `shape`.
 676  """
 677  _check_2d_shape(shape)
 678  return np.array([resize_image(image, shape) for image in video])
 679
 680
 681# ** General I/O.
 682
 683
 684def _is_url(path_or_url: _Path) -> bool:
 685  return isinstance(path_or_url, str) and path_or_url.startswith(
 686      ('http://', 'https://', 'file://')
 687  )
 688
 689
 690def read_contents(path_or_url: _Path) -> bytes:
 691  """Returns the contents of the file specified by either a path or URL."""
 692  data: bytes
 693  if _is_url(path_or_url):
 694    assert isinstance(path_or_url, str)
 695    headers = {'User-Agent': 'Chrome'}
 696    request = urllib.request.Request(path_or_url, headers=headers)
 697    with urllib.request.urlopen(request) as response:
 698      data = response.read()
 699  else:
 700    with _open(path_or_url, 'rb') as f:
 701      data = f.read()
 702  return data
 703
 704
 705@contextlib.contextmanager
 706def _read_via_local_file(path_or_url: _Path) -> Iterator[str]:
 707  """Context to copy a remote file locally to read from it.
 708
 709  Args:
 710    path_or_url: File, which may be remote.
 711
 712  Yields:
 713    The name of a local file which may be a copy of a remote file.
 714  """
 715  if _is_url(path_or_url) or not _path_is_local(path_or_url):
 716    suffix = pathlib.Path(path_or_url).suffix
 717    with tempfile.TemporaryDirectory() as directory_name:
 718      tmp_path = pathlib.Path(directory_name) / f'file{suffix}'
 719      tmp_path.write_bytes(read_contents(path_or_url))
 720      yield str(tmp_path)
 721  else:
 722    yield str(path_or_url)
 723
 724
 725@contextlib.contextmanager
 726def _write_via_local_file(path: _Path) -> Iterator[str]:
 727  """Context to write a temporary local file and subsequently copy it remotely.
 728
 729  Args:
 730    path: File, which may be remote.
 731
 732  Yields:
 733    The name of a local file which may be subsequently copied remotely.
 734  """
 735  if _path_is_local(path):
 736    yield str(path)
 737  else:
 738    suffix = pathlib.Path(path).suffix
 739    with tempfile.TemporaryDirectory() as directory_name:
 740      tmp_path = pathlib.Path(directory_name) / f'file{suffix}'
 741      yield str(tmp_path)
 742      with _open(path, mode='wb') as f:
 743        f.write(tmp_path.read_bytes())
 744
 745
 746@contextlib.contextmanager
 747def _audio_via_local_file(audio: _NDArray) -> Iterator[str]:
 748  """Context to write audio samples to a temporary local file.
 749
 750  Args:
 751    audio: Array of raw PCM audio samples, with shape (N,) for mono or (N, C)
 752      for C interleaved channels.
 753
 754  Yields:
 755    The name of a local file containing the raw samples.
 756  """
 757  with tempfile.TemporaryDirectory() as directory_name:
 758    tmp_path = pathlib.Path(directory_name) / 'audio.raw'
 759    tmp_path.write_bytes(audio.tobytes())
 760    yield str(tmp_path)
 761
 762
 763class set_show_save_dir:  # pylint: disable=invalid-name
 764  """Save all titled output from `show_*()` calls into files.
 765
 766  If the specified `directory` is not None, all titled images and videos
 767  displayed by `show_image`, `show_images`, `show_video`, and `show_videos` are
 768  also saved as files within the directory.
 769
 770  It can be used either to set the state or as a context manager:
 771
 772  >>> set_show_save_dir('/tmp')
 773  >>> show_image(color_ramp(), title='image1')  # Creates /tmp/image1.png.
 774  >>> show_video(moving_circle(), title='video2')  # Creates /tmp/video2.mp4.
 775  >>> set_show_save_dir(None)
 776
 777  >>> with set_show_save_dir('/tmp'):
 778  ...   show_image(color_ramp(), title='image1')  # Creates /tmp/image1.png.
 779  ...   show_video(moving_circle(), title='video2')  # Creates /tmp/video2.mp4.
 780  """
 781
 782  def __init__(self, directory: _Path | None):
 783    self._old_show_save_dir = _config.show_save_dir
 784    _config.show_save_dir = directory
 785
 786  def __enter__(self) -> None:
 787    pass
 788
 789  def __exit__(self, *_: Any) -> None:
 790    _config.show_save_dir = self._old_show_save_dir
 791
 792
 793# ** Image I/O.
 794
 795
 796def read_image(
 797    path_or_url: _Path,
 798    *,
 799    apply_exif_transpose: bool = True,
 800    dtype: _DTypeLike = None,  # pyrefly: ignore[bad-function-definition]
 801) -> _NDArray:
 802  """Returns an image read from a file path or URL.
 803
 804  Decoding is performed using `PIL`, which supports `uint8` images with 1, 3,
 805  or 4 channels and `uint16` images with a single channel.
 806
 807  Args:
 808    path_or_url: Path of input file.
 809    apply_exif_transpose: If True, rotate image according to EXIF orientation.
 810    dtype: Data type of the returned array.  If None, `np.uint8` or `np.uint16`
 811      is inferred automatically.
 812  """
 813  data = read_contents(path_or_url)
 814  return decompress_image(data, dtype, apply_exif_transpose)
 815
 816
 817def write_image(
 818    path: _Path, image: _ArrayLike, fmt: str = 'png', **kwargs: Any
 819) -> None:
 820  """Writes an image to a file.
 821
 822  Encoding is performed using `PIL`, which supports `uint8` images with 1, 3,
 823  or 4 channels and `uint16` images with a single channel.
 824
 825  File format is explicitly provided by `fmt` and not inferred by `path`.
 826
 827  Args:
 828    path: Path of output file.
 829    image: Array-like object.  If its type is float, it is converted to np.uint8
 830      using `to_uint8` (thus clamping to the input to the range [0.0, 1.0]).
 831      Otherwise it must be np.uint8 or np.uint16.
 832    fmt: Desired compression encoding, e.g. 'png'.
 833    **kwargs: Additional parameters for `PIL.Image.save()`.
 834  """
 835  image = _as_valid_media_array(image)
 836  if np.issubdtype(image.dtype, np.floating):
 837    image = to_uint8(image)
 838  with _open(path, 'wb') as f:
 839    _pil_image(image).save(f, format=fmt, **kwargs)
 840
 841
 842def to_rgb(
 843    array: _ArrayLike,
 844    *,
 845    vmin: float | None = None,
 846    vmax: float | None = None,
 847    cmap: str | Callable[[_ArrayLike], _NDArray] = 'gray',
 848) -> _NDArray:
 849  """Maps scalar values to RGB using value bounds and a color map.
 850
 851  Args:
 852    array: Scalar values, with arbitrary shape.
 853    vmin: Explicit min value for remapping; if None, it is obtained as the
 854      minimum finite value of `array`.
 855    vmax: Explicit max value for remapping; if None, it is obtained as the
 856      maximum finite value of `array`.
 857    cmap: A `pyplot` color map or callable, to map from 1D value to 3D or 4D
 858      color.
 859
 860  Returns:
 861    A new array in which each element is affinely mapped from [vmin, vmax]
 862    to [0.0, 1.0] and then color-mapped.
 863  """
 864  a = _as_valid_media_array(array)
 865  del array
 866  # For future numpy version 1.7.0:
 867  # vmin = np.amin(a, where=np.isfinite(a)) if vmin is None else vmin
 868  # vmax = np.amax(a, where=np.isfinite(a)) if vmax is None else vmax
 869  vmin = np.amin(np.where(np.isfinite(a), a, np.inf)) if vmin is None else vmin
 870  vmax = np.amax(np.where(np.isfinite(a), a, -np.inf)) if vmax is None else vmax
 871  a = (a.astype('float') - vmin) / (vmax - vmin + np.finfo(float).eps)
 872  if isinstance(cmap, str):
 873    if hasattr(matplotlib, 'colormaps'):
 874      rgb_from_scalar: Any = matplotlib.colormaps[cmap]  # Newer version.
 875    else:
 876      rgb_from_scalar = matplotlib.pyplot.cm.get_cmap(cmap)  # pylint: disable=no-member
 877  else:
 878    rgb_from_scalar = cmap
 879  a = cast(_NDArray, rgb_from_scalar(a))
 880  # If there is a fully opaque alpha channel, remove it.
 881  if a.shape[-1] == 4 and np.all(to_float01(a[..., 3])) == 1.0:
 882    a = a[..., :3]
 883  return a
 884
 885
 886def compress_image(
 887    image: _ArrayLike, *, fmt: str = 'png', **kwargs: Any
 888) -> bytes:
 889  """Returns a buffer containing a compressed image.
 890
 891  Args:
 892    image: Array in a format supported by `PIL`, e.g. np.uint8 or np.uint16.
 893    fmt: Desired compression encoding, e.g. 'png'.
 894    **kwargs: Options for `PIL.save()`, e.g. `optimize=True` for greater
 895      compression.
 896  """
 897  image = _as_valid_media_array(image)
 898  with io.BytesIO() as output:
 899    _pil_image(image).save(output, format=fmt, **kwargs)
 900    return output.getvalue()
 901
 902
 903def decompress_image(
 904    data: bytes, dtype: _DTypeLike = None, apply_exif_transpose: bool = True  # pyrefly: ignore[bad-function-definition]
 905) -> _NDArray:
 906  """Returns an image from a compressed data buffer.
 907
 908  Decoding is performed using `PIL`, which supports `uint8` images with 1, 3,
 909  or 4 channels and `uint16` images with a single channel.
 910
 911  Args:
 912    data: Buffer containing compressed image.
 913    dtype: Data type of the returned array.  If None, `np.uint8` or `np.uint16`
 914      is inferred automatically.
 915    apply_exif_transpose: If True, rotate image according to EXIF orientation.
 916  """
 917  pil_image: PIL.Image.Image = PIL.Image.open(io.BytesIO(data))
 918  if apply_exif_transpose:
 919    tmp_image = PIL.ImageOps.exif_transpose(pil_image)  # Future: in_place=True.
 920    assert tmp_image
 921    pil_image = tmp_image
 922  if dtype is None:
 923    dtype = np.uint16 if pil_image.mode.startswith('I') else np.uint8
 924  return np.array(pil_image, dtype=dtype)
 925
 926
 927def html_from_compressed_image(
 928    data: bytes,
 929    width: int,
 930    height: int,
 931    *,
 932    title: str | None = None,
 933    border: bool | str = False,
 934    pixelated: bool = True,
 935    fmt: str = 'png',
 936) -> str:
 937  """Returns an HTML string with an image tag containing encoded data.
 938
 939  Args:
 940    data: Compressed image bytes.
 941    width: Width of HTML image in pixels.
 942    height: Height of HTML image in pixels.
 943    title: Optional text shown centered above image.
 944    border: If `bool`, whether to place a black boundary around the image, or if
 945      `str`, the boundary CSS style.
 946    pixelated: If True, sets the CSS style to 'image-rendering: pixelated;'.
 947    fmt: Compression encoding.
 948  """
 949  b64 = base64.b64encode(data).decode('utf-8')
 950  if isinstance(border, str):
 951    border = f'{border}; '
 952  elif border:
 953    border = 'border:1px solid black; '
 954  else:
 955    border = ''
 956  s_pixelated = 'pixelated' if pixelated else 'auto'
 957  s = (
 958      f'<img width="{width}" height="{height}"'
 959      f' style="{border}image-rendering:{s_pixelated}; object-fit:cover;"'
 960      f' src="data:image/{fmt};base64,{b64}"/>'
 961  )
 962  if title is not None:
 963    s = f"""<div style="display:flex; align-items:left;">
 964      <div style="display:flex; flex-direction:column; align-items:center;">
 965      <div>{title}</div><div>{s}</div></div></div>"""
 966  return s
 967
 968
 969def _get_width_height(
 970    width: int | None, height: int | None, shape: tuple[int, int]
 971) -> tuple[int, int]:
 972  """Returns (width, height) given optional parameters and image shape."""
 973  assert len(shape) == 2, shape
 974  if width and height:
 975    return width, height
 976  if width and not height:
 977    return width, int(width * (shape[0] / shape[1]) + 0.5)
 978  if height and not width:
 979    return int(height * (shape[1] / shape[0]) + 0.5), height
 980  return shape[::-1]  # pyrefly: ignore[bad-return]
 981
 982
 983def _ensure_mapped_to_rgb(
 984    image: _ArrayLike,
 985    *,
 986    vmin: float | None = None,
 987    vmax: float | None = None,
 988    cmap: str | Callable[[_ArrayLike], _NDArray] = 'gray',
 989) -> _NDArray:
 990  """Ensure image is mapped to RGB."""
 991  image = _as_valid_media_array(image)
 992  if not (image.ndim == 2 or (image.ndim == 3 and image.shape[2] in (1, 3, 4))):
 993    raise ValueError(
 994        f'Image with shape {image.shape} is neither a 2D array'
 995        ' nor a 3D array with 1, 3, or 4 channels.'
 996    )
 997  if image.ndim == 3 and image.shape[2] == 1:
 998    image = image[:, :, 0]
 999  if image.ndim == 2:
1000    image = to_rgb(image, vmin=vmin, vmax=vmax, cmap=cmap)
1001  return image
1002
1003
1004def show_image(
1005    image: _ArrayLike, *, title: str | None = None, **kwargs: Any
1006) -> str | None:
1007  """Displays an image in the notebook and optionally saves it to a file.
1008
1009  See `show_images`.
1010
1011  >>> show_image(np.random.rand(100, 100))
1012  >>> show_image(np.random.randint(0, 256, size=(80, 80, 3), dtype='uint8'))
1013  >>> show_image(np.random.rand(10, 10) - 0.5, cmap='bwr', height=100)
1014  >>> show_image(read_image('/tmp/image.png'))
1015  >>> url = 'https://github.com/hhoppe/data/raw/main/image.png'
1016  >>> show_image(read_image(url))
1017
1018  Args:
1019    image: 2D array-like, or 3D array-like with 1, 3, or 4 channels.
1020    title: Optional text shown centered above the image.
1021    **kwargs: See `show_images`.
1022
1023  Returns:
1024    html string if `return_html` is `True`.
1025  """
1026  return show_images([np.asarray(image)], [title], **kwargs)
1027
1028
1029def show_images(
1030    images: Iterable[_ArrayLike] | Mapping[str, _ArrayLike],
1031    titles: Iterable[str | None] | None = None,
1032    *,
1033    width: int | None = None,
1034    height: int | None = None,
1035    downsample: bool = True,
1036    columns: int | None = None,
1037    vmin: float | None = None,
1038    vmax: float | None = None,
1039    cmap: str | Callable[[_ArrayLike], _NDArray] = 'gray',
1040    border: bool | str = False,
1041    ylabel: str = '',
1042    html_class: str = 'show_images',
1043    pixelated: bool | None = None,
1044    return_html: bool = False,
1045) -> str | None:
1046  """Displays a row of images in the IPython/Jupyter notebook.
1047
1048  If a directory has been specified using `set_show_save_dir`, also saves each
1049  titled image to a file in that directory based on its title.
1050
1051  >>> image1, image2 = np.random.rand(64, 64, 3), color_ramp((64, 64))
1052  >>> show_images([image1, image2])
1053  >>> show_images({'random image': image1, 'color ramp': image2}, height=128)
1054  >>> show_images([image1, image2] * 5, columns=4, border=True)
1055
1056  Args:
1057    images: Iterable of images, or dictionary of `{title: image}`.  Each image
1058      must be either a 2D array or a 3D array with 1, 3, or 4 channels.
1059    titles: Optional strings shown above the corresponding images.
1060    width: Optional, overrides displayed width (in pixels).
1061    height: Optional, overrides displayed height (in pixels).
1062    downsample: If True, each image whose width or height is greater than the
1063      specified `width` or `height` is resampled to the display resolution. This
1064      improves antialiasing and reduces the size of the notebook.
1065    columns: Optional, maximum number of images per row.
1066    vmin: For single-channel image, explicit min value for display.
1067    vmax: For single-channel image, explicit max value for display.
1068    cmap: For single-channel image, `pyplot` color map or callable to map 1D to
1069      3D color.
1070    border: If `bool`, whether to place a black boundary around the image, or if
1071      `str`, the boundary CSS style.
1072    ylabel: Text (rotated by 90 degrees) shown on the left of each row.
1073    html_class: CSS class name used in definition of HTML element.
1074    pixelated: If True, sets the CSS style to 'image-rendering: pixelated;'; if
1075      False, sets 'image-rendering: auto'; if None, uses pixelated rendering
1076      only on images for which `width` or `height` introduces magnification.
1077    return_html: If `True` return the raw HTML `str` instead of displaying.
1078
1079  Returns:
1080    html string if `return_html` is `True`.
1081  """
1082  if isinstance(images, Mapping):
1083    if titles is not None:
1084      raise ValueError('Cannot have images dictionary and titles parameter.')
1085    list_titles, list_images = list(images.keys()), list(images.values())
1086  else:
1087    list_images = list(images)
1088    list_titles = [None] * len(list_images) if titles is None else list(titles)
1089    if len(list_images) != len(list_titles):
1090      raise ValueError(
1091          'Number of images does not match number of titles'
1092          f' ({len(list_images)} vs {len(list_titles)}).'
1093      )
1094
1095  list_images = [
1096      _ensure_mapped_to_rgb(image, vmin=vmin, vmax=vmax, cmap=cmap)
1097      for image in list_images
1098  ]
1099
1100  def maybe_downsample(image: _NDArray) -> _NDArray:
1101    shape = image.shape[0], image.shape[1]
1102    w, h = _get_width_height(width, height, shape)
1103    if w < shape[1] or h < shape[0]:
1104      image = resize_image(image, (h, w))
1105    return image
1106
1107  if downsample:
1108    list_images = [maybe_downsample(image) for image in list_images]
1109  png_datas = [compress_image(to_uint8(image)) for image in list_images]
1110
1111  for title, png_data in zip(list_titles, png_datas):
1112    if title is not None and _config.show_save_dir:
1113      path = pathlib.Path(_config.show_save_dir) / f'{title}.png'
1114      with _open(path, mode='wb') as f:
1115        f.write(png_data)
1116
1117  def html_from_compressed_images() -> str:
1118    html_strings = []
1119    for image, title, png_data in zip(list_images, list_titles, png_datas):
1120      w, h = _get_width_height(width, height, image.shape[:2])  # pyrefly: ignore[missing-attribute]
1121      magnified = h > image.shape[0] or w > image.shape[1]  # pyrefly: ignore[missing-attribute]
1122      pixelated2 = pixelated if pixelated is not None else magnified
1123      html_strings.append(
1124          html_from_compressed_image(
1125              png_data, w, h, title=title, border=border, pixelated=pixelated2  # pyrefly: ignore[bad-argument-type]
1126          )
1127      )
1128    # Create single-row tables each with no more than 'columns' elements.
1129    table_strings = []
1130    for row_html_strings in _chunked(html_strings, columns):
1131      td = '<td style="padding:1px;">'
1132      s = ''.join(f'{td}{e}</td>' for e in row_html_strings)
1133      if ylabel:
1134        style = 'writing-mode:vertical-lr; transform:rotate(180deg);'
1135        s = f'{td}<span style="{style}">{ylabel}</span></td>' + s
1136      table_strings.append(
1137          f'<table class="{html_class}"'
1138          f' style="border-spacing:0px;"><tr>{s}</tr></table>'
1139      )
1140    return ''.join(table_strings)
1141
1142  s = html_from_compressed_images()
1143  while len(s) > _IPYTHON_HTML_SIZE_LIMIT * 0.5:
1144    warnings.warn('mediapy: subsampling images to reduce HTML size')
1145    list_images = [image[::2, ::2] for image in list_images]
1146    png_datas = [compress_image(to_uint8(image)) for image in list_images]
1147    s = html_from_compressed_images()
1148  if return_html:
1149    return s
1150  _display_html(s)
1151  return None
1152
1153
1154def compare_images(
1155    images: Iterable[_ArrayLike],
1156    *,
1157    vmin: float | None = None,
1158    vmax: float | None = None,
1159    cmap: str | Callable[[_ArrayLike], _NDArray] = 'gray',
1160) -> None:
1161  """Compare two images using an interactive slider.
1162
1163  Displays an HTML slider component to interactively swipe between two images.
1164  The slider functionality requires that the web browser have Internet access.
1165  See additional info in `https://github.com/sneas/img-comparison-slider`.
1166
1167  >>> image1, image2 = np.random.rand(64, 64, 3), color_ramp((64, 64))
1168  >>> compare_images([image1, image2])
1169
1170  Args:
1171    images: Iterable of images.  Each image must be either a 2D array or a 3D
1172      array with 1, 3, or 4 channels.  There must be exactly two images.
1173    vmin: For single-channel image, explicit min value for display.
1174    vmax: For single-channel image, explicit max value for display.
1175    cmap: For single-channel image, `pyplot` color map or callable to map 1D to
1176      3D color.
1177  """
1178  list_images = [
1179      _ensure_mapped_to_rgb(image, vmin=vmin, vmax=vmax, cmap=cmap)
1180      for image in images
1181  ]
1182  if len(list_images) != 2:
1183    raise ValueError('The number of images must be 2.')
1184  png_datas = [compress_image(to_uint8(image)) for image in list_images]
1185  b64_1, b64_2 = [
1186      base64.b64encode(png_data).decode('utf-8') for png_data in png_datas
1187  ]
1188  s = _IMAGE_COMPARISON_HTML.replace('{b64_1}', b64_1).replace('{b64_2}', b64_2)
1189  _display_html(s)
1190
1191
1192# ** Video I/O.
1193
1194
1195def _filename_suffix_from_codec(codec: str) -> str:
1196  if codec == 'gif':
1197    return '.gif'
1198  if codec == 'vp9':
1199    return '.webm'
1200
1201  return '.mp4'
1202
1203
1204def _get_ffmpeg_path() -> str:
1205  path = _search_for_ffmpeg_path()
1206  if not path:
1207    raise RuntimeError(
1208        f"Program '{_config.ffmpeg_name_or_path}' is not found;"
1209        " perhaps install ffmpeg using 'apt install ffmpeg'."
1210    )
1211  return path
1212
1213
1214@typing.overload
1215def _run_ffmpeg(
1216    ffmpeg_args: Sequence[str],
1217    stdin: int | None = None,
1218    stdout: int | None = None,
1219    stderr: int | None = None,
1220    encoding: None = None,  # No encoding -> bytes
1221    allowed_input_files: Sequence[str] | None = None,
1222    allowed_output_files: Sequence[str] | None = None,
1223    sandbox_max_run_time_secs: int | None = None,
1224) -> subprocess.Popen[bytes]:
1225  ...
1226
1227
1228@typing.overload
1229def _run_ffmpeg(
1230    ffmpeg_args: Sequence[str],
1231    stdin: int | None = None,
1232    stdout: int | None = None,
1233    stderr: int | None = None,
1234    encoding: str = ...,  # Encoding -> str
1235    allowed_input_files: Sequence[str] | None = None,
1236    allowed_output_files: Sequence[str] | None = None,
1237    sandbox_max_run_time_secs: int | None = None,
1238) -> subprocess.Popen[str]:
1239  ...
1240
1241
1242def _run_ffmpeg(
1243    ffmpeg_args: Sequence[str],
1244    stdin: int | None = None,
1245    stdout: int | None = None,
1246    stderr: int | None = None,
1247    encoding: str | None = None,
1248    allowed_input_files: Sequence[str] | None = None,
1249    allowed_output_files: Sequence[str] | None = None,
1250    sandbox_max_run_time_secs: int | None = None,
1251) -> subprocess.Popen[bytes] | subprocess.Popen[str]:
1252  """Runs ffmpeg with the given args.
1253
1254  Args:
1255    ffmpeg_args: The args to pass to ffmpeg.
1256    stdin: Same as in `subprocess.Popen`.
1257    stdout: Same as in `subprocess.Popen`.
1258    stderr: Same as in `subprocess.Popen`.
1259    encoding: Same as in `subprocess.Popen`.
1260    allowed_input_files: The input files to allow for ffmpeg.
1261    allowed_output_files: The output files to allow for ffmpeg.
1262    sandbox_max_run_time_secs: The maximum time in seconds to run the sandbox.
1263      If None, the default limit is 30 minutes.
1264
1265  Returns:
1266    The subprocess.Popen object with running ffmpeg process.
1267  """
1268  argv = []
1269  # In open source, keep env=None to preserve default behavior.
1270  # Context: https://github.com/google/mediapy/pull/62
1271  env: Any = None  # pylint: disable=unused-variable
1272  ffmpeg_path = _get_ffmpeg_path()
1273
1274  # Sandbox max runtime, allowed input and ouput files are not supported in
1275  # open source.
1276  del allowed_input_files
1277  del allowed_output_files
1278  del sandbox_max_run_time_secs
1279
1280  argv.append(ffmpeg_path)
1281  argv.extend(ffmpeg_args)
1282
1283  return subprocess.Popen(
1284      argv,
1285      stdin=stdin,
1286      stdout=stdout,
1287      stderr=stderr,
1288      encoding=encoding,
1289      env=env,
1290  )
1291
1292
1293def video_is_available() -> bool:
1294  """Returns True if the program `ffmpeg` is found.
1295
1296  See also `set_ffmpeg`.
1297  """
1298  return _search_for_ffmpeg_path() is not None
1299
1300
1301class VideoMetadata(NamedTuple):
1302  """Represents the data stored in a video container header.
1303
1304  Attributes:
1305    num_images: Number of frames that is expected from the video stream.  This
1306      is estimated from the framerate and the duration stored in the video
1307      header, so it might be inexact.  We set the value to -1 if number of
1308      frames is not found in the header.
1309    shape: The dimensions (height, width) of each video frame.
1310    fps: The framerate in frames per second.
1311    bps: The estimated bitrate of the video stream in bits per second, retrieved
1312      from the video header.
1313  """
1314
1315  num_images: int
1316  shape: tuple[int, int]
1317  fps: float
1318  bps: int | None
1319
1320
1321def _get_video_metadata(path: _Path) -> VideoMetadata:
1322  """Returns attributes of video stored in the specified local file."""
1323  if not pathlib.Path(path).is_file():
1324    raise RuntimeError(f"Video file '{path}' is not found.")
1325
1326  command = [
1327      '-nostdin',
1328      '-i',
1329      str(path),
1330      '-acodec',
1331      'copy',
1332      # Necessary to get "frame= *(\d+)" using newer ffmpeg versions.
1333      # Previously, was `'-vcodec', 'copy'`
1334      '-vf',
1335      'select=1',
1336      '-vsync',
1337      '0',
1338      '-f',
1339      'null',
1340      '-',
1341  ]
1342  with _run_ffmpeg(
1343      command,
1344      allowed_input_files=[str(path)],
1345      stderr=subprocess.PIPE,
1346      encoding='utf-8',
1347  ) as proc:
1348    _, err = proc.communicate()
1349  bps = fps = num_images = width = height = rotation = None
1350  before_output_info = True
1351  for line in err.split('\n'):
1352    if line.startswith('Output '):
1353      before_output_info = False
1354    if match := re.search(r', bitrate: *([\d.]+) kb/s', line):
1355      bps = int(match.group(1)) * 1000
1356    if matches := re.findall(r'frame= *(\d+) ', line):
1357      num_images = int(matches[-1])
1358    if 'Stream #0:' in line and ': Video:' in line and before_output_info:
1359      if not (match := re.search(r', (\d+)x(\d+)', line)):
1360        raise RuntimeError(f'Unable to parse video dimensions in line {line}')
1361      width, height = int(match.group(1)), int(match.group(2))
1362      # Some videos have a 'k' suffix for kiloframes per second. Thus an extra
1363      # `(k?)` is used for this case. This usually happens when we don't know
1364      # the exact framerate. However, in order to not raise an error here, we
1365      # will try to parse the framerate as x1000.
1366      if match := re.search(r', ([\d.]+)(k?) fps', line):
1367        number = float(match.group(1))
1368        if match.group(2) == 'k':
1369          fps = number * 1000
1370        else:
1371          fps = number
1372      elif str(path).endswith('.gif'):
1373        # Some GIF files lack a framerate attribute; use a reasonable default.
1374        fps = 10
1375      else:
1376        raise RuntimeError(f'Unable to parse video framerate in line {line}')
1377    if match := re.fullmatch(r'\s*rotate\s*:\s*(\d+)', line):
1378      rotation = int(match.group(1))
1379    if match := re.fullmatch(r'.*rotation of (-?\d+).*\sdegrees', line):
1380      rotation = int(match.group(1))
1381  if not num_images:
1382    num_images = -1
1383  if not width:
1384    raise RuntimeError(f'Unable to parse video header: {err}')
1385  # By default, ffmpeg enables "-autorotate"; we just fix the dimensions.
1386  if rotation in (90, 270, -90, -270):
1387    width, height = height, width
1388  assert height is not None and width is not None
1389  shape = height, width
1390  assert fps is not None
1391  return VideoMetadata(num_images, shape, fps, bps)
1392
1393
1394class _VideoIO:
1395  """Base class for `VideoReader` and `VideoWriter`."""
1396
1397  def _get_pix_fmt(self, dtype: _DType, image_format: str) -> str:
1398    """Returns ffmpeg pix_fmt given data type and image format."""
1399    native_endian_suffix = {'little': 'le', 'big': 'be'}[sys.byteorder]
1400    return {
1401        np.uint8: {
1402            'rgb': 'rgb24',
1403            'yuv': 'yuv444p',
1404            'gray': 'gray',
1405        },
1406        np.uint16: {
1407            'rgb': 'rgb48' + native_endian_suffix,
1408            'yuv': 'yuv444p16' + native_endian_suffix,
1409            'gray': 'gray16' + native_endian_suffix,
1410        },
1411    }[dtype.type][image_format]  # pyrefly: ignore[bad-index]
1412
1413
1414class VideoReader(_VideoIO):
1415  """Context to read a compressed video as an iterable over its images.
1416
1417  >>> with VideoReader('/tmp/river.mp4') as reader:
1418  ...   print(f'Video has {reader.num_images} images with shape={reader.shape},'
1419  ...         f' at {reader.fps} frames/sec and {reader.bps} bits/sec.')
1420  ...   for image in reader:
1421  ...     print(image.shape)
1422
1423  >>> with VideoReader('/tmp/river.mp4') as reader:
1424  ...   video = np.array(tuple(reader))
1425
1426  >>> url = 'https://github.com/hhoppe/data/raw/main/video.mp4'
1427  >>> with VideoReader(url) as reader:
1428  ...   show_video(reader)
1429
1430  Attributes:
1431    path_or_url: Location of input video.
1432    output_format: Format of output images (default 'rgb').  If 'rgb', each
1433      image has shape=(height, width, 3) with R, G, B values.  If 'yuv', each
1434      image has shape=(height, width, 3) with Y, U, V values.  If 'gray', each
1435      image has shape=(height, width).
1436    dtype: Data type for output images.  The default is `np.uint8`.  Use of
1437      `np.uint16` allows reading 10-bit or 12-bit data without precision loss.
1438    metadata: Object storing the information retrieved from the video header.
1439      Its attributes are copied as attributes in this class.
1440    num_images: Number of frames that is expected from the video stream.  This
1441      is estimated from the framerate and the duration stored in the video
1442      header, so it might be inexact.
1443    shape: The dimensions (height, width) of each video frame.
1444    fps: The framerate in frames per second.
1445    bps: The estimated bitrate of the video stream in bits per second, retrieved
1446      from the video header.
1447    stream_index: The stream index to read from. The default is 0.
1448    sandbox_max_run_time_secs: The maximum time in seconds to run the sandbox.
1449      If None, the default limit is 30 minutes. Unused in open source.
1450  """
1451
1452  path_or_url: _Path
1453  output_format: str
1454  dtype: _DType
1455  metadata: VideoMetadata
1456  num_images: int
1457  shape: tuple[int, int]
1458  fps: float
1459  bps: int | None
1460  stream_index: int
1461  _num_bytes_per_image: int
1462
1463  def __init__(
1464      self,
1465      path_or_url: _Path,
1466      *,
1467      stream_index: int = 0,
1468      output_format: str = 'rgb',
1469      dtype: _DTypeLike = np.uint8,
1470      sandbox_max_run_time_secs: int | None = None,
1471  ):
1472    if output_format not in {'rgb', 'yuv', 'gray'}:
1473      raise ValueError(
1474          f'Output format {output_format} is not rgb, yuv, or gray.'
1475      )
1476    self.path_or_url = path_or_url
1477    self.output_format = output_format
1478    self.stream_index = stream_index
1479    self.dtype = np.dtype(dtype)
1480    if self.dtype.type not in (np.uint8, np.uint16):
1481      raise ValueError(f'Type {dtype} is not np.uint8 or np.uint16.')
1482    self.sandbox_max_run_time_secs = sandbox_max_run_time_secs
1483    self._read_via_local_file: Any = None
1484    self._popen: subprocess.Popen[bytes] | None = None
1485    self._proc: subprocess.Popen[bytes] | None = None
1486
1487  def __enter__(self) -> 'VideoReader':
1488    try:
1489      self._read_via_local_file = _read_via_local_file(self.path_or_url)
1490      # pylint: disable-next=no-member
1491      tmp_name = self._read_via_local_file.__enter__()
1492
1493      self.metadata = _get_video_metadata(tmp_name)
1494      self.num_images, self.shape, self.fps, self.bps = self.metadata
1495      pix_fmt = self._get_pix_fmt(self.dtype, self.output_format)
1496      num_channels = {'rgb': 3, 'yuv': 3, 'gray': 1}[self.output_format]
1497      bytes_per_channel = self.dtype.itemsize
1498      self._num_bytes_per_image = (
1499          math.prod(self.shape) * num_channels * bytes_per_channel
1500      )
1501
1502      command = [
1503          '-v',
1504          'panic',
1505          '-nostdin',
1506          '-i',
1507          tmp_name,
1508          '-vcodec',
1509          'rawvideo',
1510          '-f',
1511          'image2pipe',
1512          '-map',
1513          f'0:v:{self.stream_index}',
1514          '-pix_fmt',
1515          pix_fmt,
1516          '-vsync',
1517          'vfr',
1518          '-',
1519      ]
1520      self._popen = _run_ffmpeg(
1521          command,
1522          stdout=subprocess.PIPE,
1523          stderr=subprocess.PIPE,
1524          allowed_input_files=[tmp_name],
1525          sandbox_max_run_time_secs=self.sandbox_max_run_time_secs,
1526      )
1527      self._proc = self._popen.__enter__()
1528    except Exception:
1529      self.__exit__(None, None, None)
1530      raise
1531    return self
1532
1533  def __exit__(self, *_: Any) -> None:
1534    self.close()
1535
1536  def read(self) -> _NDArray | None:
1537    """Reads a video image frame (or None if at end of file).
1538
1539    Returns:
1540      A numpy array in the format specified by `output_format`, i.e., a 3D
1541      array with 3 color channels, except for format 'gray' which is 2D.
1542
1543    Raises:
1544      RuntimeError: If there is an error reading from the output file.
1545    """
1546    assert self._proc, 'Error: reading from an already closed context.'
1547    stdout = self._proc.stdout
1548    assert stdout is not None
1549    data = stdout.read(self._num_bytes_per_image)
1550    if not data:  # Due to either end-of-file or subprocess error.
1551      self.close()  # Raises exception if subprocess had error.
1552      return None  # To indicate end-of-file.
1553    if len(data) != self._num_bytes_per_image:
1554      self._proc.wait()
1555      stderr = self._proc.stderr
1556      stderr_output = ''
1557      if stderr is not None:
1558        stderr_output = stderr.read().decode('utf-8', errors='replace').strip()
1559      raise RuntimeError(
1560          f'ffmpeg exited with code {self._proc.returncode}.\nIncomplete'
1561          f' frame read: expected {self._num_bytes_per_image} bytes, but got'
1562          f' {len(data)}.\nffmpeg stderr:\n{stderr_output}'
1563      )
1564    image = np.frombuffer(data, dtype=self.dtype)
1565    if self.output_format == 'rgb':
1566      image = image.reshape(*self.shape, 3)
1567    elif self.output_format == 'yuv':  # Convert from planar YUV to pixel YUV.
1568      image = np.moveaxis(image.reshape(3, *self.shape), 0, 2)
1569    elif self.output_format == 'gray':  # Generate 2D rather than 3D ndimage.
1570      image = image.reshape(*self.shape)
1571    else:
1572      raise AssertionError
1573    return image
1574
1575  def __iter__(self) -> Iterator[_NDArray]:
1576    while True:
1577      image = self.read()
1578      if image is None:
1579        return
1580      yield image
1581
1582  def close(self) -> None:
1583    """Terminates video reader.  (Called automatically at end of context.)"""
1584    if self._popen:
1585      self._popen.__exit__(None, None, None)
1586      self._popen = None
1587      self._proc = None
1588    if self._read_via_local_file:
1589      # pylint: disable-next=no-member
1590      self._read_via_local_file.__exit__(None, None, None)
1591      self._read_via_local_file = None
1592
1593
1594class VideoWriter(_VideoIO):
1595  """Context to write a compressed video.
1596
1597  >>> shape = 480, 640
1598  >>> with VideoWriter('/tmp/v.mp4', shape, fps=60) as writer:
1599  ...   for image in moving_circle(shape, num_images=60):
1600  ...     writer.add_image(image)
1601  >>> show_video(read_video('/tmp/v.mp4'))
1602
1603
1604  Bitrate control may be specified using at most one of: `bps`, `qp`, or `crf`.
1605  If none are specified, `qp` is set to a default value.
1606  See https://slhck.info/video/2017/03/01/rate-control.html
1607
1608  If codec is 'gif', the args `bps`, `qp`, `crf`, and `encoded_format` are
1609  ignored.
1610
1611  Attributes:
1612    path: Output video.  Its suffix (e.g. '.mp4') determines the video container
1613      format.  The suffix must be '.gif' if the codec is 'gif'.
1614    shape: 2D spatial dimensions (height, width) of video image frames.  The
1615      dimensions must be even if 'encoded_format' has subsampled chroma (e.g.,
1616      'yuv420p' or 'yuv420p10le').
1617    codec: Compression algorithm as defined by "ffmpeg -codecs" (e.g., 'h264',
1618      'hevc', 'vp9', or 'gif').
1619    metadata: Optional VideoMetadata object whose `fps` and `bps` attributes are
1620      used if not specified as explicit parameters.
1621    fps: Frames-per-second framerate (default is 60.0 except 25.0 for 'gif').
1622    bps: Requested average bits-per-second bitrate (default None).
1623    qp: Quantization parameter for video compression quality (default None).
1624    crf: Constant rate factor for video compression quality (default None).
1625    ffmpeg_args: Additional arguments for `ffmpeg` command, e.g. '-g 30' to
1626      introduce I-frames, or '-bf 0' to omit B-frames.
1627    input_format: Format of input images (default 'rgb').  If 'rgb', each image
1628      has shape=(height, width, 3) or (height, width).  If 'yuv', each image has
1629      shape=(height, width, 3) with Y, U, V values.  If 'gray', each image has
1630      shape=(height, width).
1631    dtype: Expected data type for input images (any float input images are
1632      converted to `dtype`).  The default is `np.uint8`.  Use of `np.uint16` is
1633      necessary when encoding >8 bits/channel.
1634    encoded_format: Pixel format as defined by `ffmpeg -pix_fmts`, e.g.,
1635      'yuv420p' (2x2-subsampled chroma), 'yuv444p' (full-res chroma),
1636      'yuv420p10le' (10-bit per channel), etc.  The default (None) selects
1637      'yuv420p' if all shape dimensions are even, else 'yuv444p'.
1638    sandbox_max_run_time_secs: The maximum time in seconds to run the sandbox.
1639      If None, the default limit is 30 minutes. Unused in open source.
1640    audio: Optional audio data as a NumPy array.  It should have shape (N,) for
1641      mono or (N, C) for multi-channel audio, where N is the number of samples
1642      and C is the number of channels.  The dtype must be one of `np.uint8`,
1643      `np.int16`, `np.int32`, `np.float32`, or `np.float64`, in native byte
1644      order.  Float samples are nominally in [-1.0, 1.0], and values outside
1645      this range typically clip; integer samples span the full range of their
1646      type, where `np.uint8` is unsigned with silence at 128.  The two streams
1647      are not truncated to a common length: if the audio is shorter than the
1648      video, the remainder is silent; if it is longer, the audio continues past
1649      the end of the video stream (most players keep showing the last frame).
1650    audio_sample_rate: Sample rate of the audio in Hz. Required if `audio` is
1651      provided.
1652    audio_codec: Audio compression algorithm as defined by "ffmpeg -codecs"
1653      (default 'aac').  It must be supported by the video container, e.g., 'aac'
1654      for MP4 ('h264' or 'hevc') or 'libopus' for WebM ('vp9').  Ignored if
1655      `audio` is None.
1656  """
1657
1658  def __init__(
1659      self,
1660      path: _Path,
1661      shape: tuple[int, int],
1662      *,
1663      codec: str = 'h264',
1664      metadata: VideoMetadata | None = None,
1665      fps: float | None = None,
1666      bps: int | None = None,
1667      qp: int | None = None,
1668      crf: float | None = None,
1669      ffmpeg_args: str | Sequence[str] = '',
1670      input_format: str = 'rgb',
1671      dtype: _DTypeLike = np.uint8,
1672      encoded_format: str | None = None,
1673      sandbox_max_run_time_secs: int | None = None,
1674      audio: _NDArray | None = None,
1675      audio_sample_rate: int | None = None,
1676      audio_codec: str = 'aac',
1677  ) -> None:
1678    _check_2d_shape(shape)
1679    if fps is None and metadata:
1680      fps = metadata.fps
1681    if fps is None:
1682      fps = 25.0 if codec == 'gif' else 60.0
1683    if fps <= 0.0:
1684      raise ValueError(f'Frame-per-second value {fps} is invalid.')
1685    if bps is None and metadata:
1686      bps = metadata.bps
1687    bps = int(bps) if bps is not None else None
1688    if bps is not None and bps <= 0:
1689      raise ValueError(f'Bitrate value {bps} is invalid.')
1690    if qp is not None and (not isinstance(qp, int) or qp < 0):
1691      raise ValueError(
1692          f'Quantization parameter {qp} cannot be negative. It must be a'
1693          ' non-negative integer.'
1694      )
1695    num_rate_specifications = sum(x is not None for x in (bps, qp, crf))
1696    if num_rate_specifications > 1:
1697      raise ValueError(
1698          f'Must specify at most one of bps, qp, or crf ({bps}, {qp}, {crf}).'
1699      )
1700    ffmpeg_args = (
1701        shlex.split(ffmpeg_args)
1702        if isinstance(ffmpeg_args, str)
1703        else list(ffmpeg_args)
1704    )
1705    if input_format not in {'rgb', 'yuv', 'gray'}:
1706      raise ValueError(f'Input format {input_format} is not rgb, yuv, or gray.')
1707    dtype = np.dtype(dtype)
1708    if dtype.type not in (np.uint8, np.uint16):
1709      raise ValueError(f'Type {dtype} is not np.uint8 or np.uint16.')
1710    if audio is not None:
1711      if audio_sample_rate is None:
1712        raise ValueError('The audio_sample_rate must be set if audio is set.')
1713      if audio.ndim not in (1, 2):
1714        raise ValueError(f'Audio shape {audio.shape} is not (N,) or (N, C).')
1715      if audio.dtype.type not in _FFMPEG_AUDIO_FORMAT_FROM_DTYPE:
1716        raise ValueError(f'Audio type {audio.dtype} is unsupported.')
1717      if not audio.dtype.isnative:
1718        raise ValueError(
1719            f'Audio type {audio.dtype} is not in native byte order.'
1720        )
1721      if codec == 'gif':
1722        raise ValueError('Audio is not supported with the gif codec.')
1723    self.path = pathlib.Path(path)
1724    self.shape = shape
1725    all_dimensions_are_even = all(dim % 2 == 0 for dim in shape)
1726    if encoded_format is None:
1727      encoded_format = 'yuv420p' if all_dimensions_are_even else 'yuv444p'
1728    if not all_dimensions_are_even and encoded_format.startswith(
1729        ('yuv42', 'yuvj42')
1730    ):
1731      raise ValueError(
1732          f'With encoded_format {encoded_format}, video dimensions must be'
1733          f' even, but shape is {shape}.'
1734      )
1735    self.fps = fps
1736    self.codec = codec
1737    self.bps = bps
1738    self.qp = qp
1739    self.crf = crf
1740    self.ffmpeg_args = ffmpeg_args
1741    self.input_format = input_format
1742    self.dtype = dtype
1743    self.encoded_format = encoded_format
1744    self.sandbox_max_run_time_secs = sandbox_max_run_time_secs
1745    self.audio = audio
1746    self.audio_sample_rate = audio_sample_rate
1747    self.audio_codec = audio_codec
1748    if num_rate_specifications == 0 and not ffmpeg_args:
1749      qp = 20 if math.prod(self.shape) <= 640 * 480 else 28
1750    self._bitrate_args = (
1751        (['-vb', f'{bps}'] if bps is not None else [])
1752        + (['-qp', f'{qp}'] if qp is not None else [])
1753        + (['-vb', '0', '-crf', f'{crf}'] if crf is not None else [])
1754    )
1755    if self.codec == 'gif':
1756      if self.path.suffix != '.gif':
1757        raise ValueError(f"File '{self.path}' does not have a .gif suffix.")
1758      self.encoded_format = 'pal8'
1759      self._bitrate_args = []
1760      video_filter = 'split[s0][s1];[s0]palettegen[p];[s1][p]paletteuse'
1761      # Less common (and likely less useful) is a per-frame color palette:
1762      # video_filter = ('split[s0][s1];[s0]palettegen=stats_mode=single[p];'
1763      #                 '[s1][p]paletteuse=new=1')
1764      self.ffmpeg_args = ['-vf', video_filter, '-f', 'gif'] + self.ffmpeg_args
1765    self._exit_stack: contextlib.ExitStack | None = None
1766    self._proc: subprocess.Popen[bytes] | None = None
1767
1768  def __enter__(self) -> 'VideoWriter':
1769    input_pix_fmt = self._get_pix_fmt(self.dtype, self.input_format)
1770    try:
1771      self._exit_stack = contextlib.ExitStack()
1772      tmp_name = self._exit_stack.enter_context(
1773          _write_via_local_file(self.path)
1774      )
1775
1776      # Writing to stdout using ('-f', 'mp4', '-') would require
1777      # ('-movflags', 'frag_keyframe+empty_moov') and the result is nonportable.
1778      height, width = self.shape
1779
1780      audio_input_args = []
1781      audio_output_args = ['-an']
1782      allowed_input_files = []
1783
1784      if self.audio is not None:
1785        audio_path = self._exit_stack.enter_context(
1786            _audio_via_local_file(self.audio)
1787        )
1788        channels = 1 if self.audio.ndim == 1 else self.audio.shape[1]
1789        audio_format = _FFMPEG_AUDIO_FORMAT_FROM_DTYPE[self.audio.dtype.type]
1790        if self.audio.dtype.itemsize > 1:
1791          audio_format += {'little': 'le', 'big': 'be'}[sys.byteorder]
1792        audio_input_args = [
1793            '-f',
1794            audio_format,
1795            '-ar',
1796            str(self.audio_sample_rate),
1797            '-ac',
1798            str(channels),
1799            '-i',
1800            audio_path,
1801        ]
1802        audio_output_args = ['-c:a', self.audio_codec]
1803        allowed_input_files.append(audio_path)
1804
1805      command = (
1806          [
1807              '-v',
1808              'error',
1809              '-f',
1810              'rawvideo',
1811              '-vcodec',
1812              'rawvideo',
1813              '-pix_fmt',
1814              input_pix_fmt,
1815              '-s',
1816              f'{width}x{height}',
1817              '-r',
1818              f'{self.fps}',
1819              '-i',
1820              '-',
1821          ]
1822          + audio_input_args
1823          + audio_output_args
1824          + [
1825              '-vcodec',
1826              self.codec,
1827              '-pix_fmt',
1828              self.encoded_format,
1829          ]
1830          + self._bitrate_args
1831          + self.ffmpeg_args
1832          + ['-y', tmp_name]
1833      )
1834      self._proc = self._exit_stack.enter_context(
1835          _run_ffmpeg(
1836              command,
1837              stdin=subprocess.PIPE,
1838              stderr=subprocess.PIPE,
1839              # `_run_ffmpeg` omits the sandbox flag only for None, so an empty
1840              # list would pass an empty '--sandbox_read_access_files'.
1841              allowed_input_files=allowed_input_files or None,
1842              allowed_output_files=[tmp_name],
1843              sandbox_max_run_time_secs=self.sandbox_max_run_time_secs,
1844          )
1845      )
1846    except Exception:
1847      self.__exit__(None, None, None)
1848      raise
1849    return self
1850
1851  def __exit__(self, *_: Any) -> None:
1852    self.close()
1853
1854  def add_image(self, image: _NDArray) -> None:
1855    """Writes a video frame.
1856
1857    Args:
1858      image: Array whose dtype and first two dimensions must match the `dtype`
1859        and `shape` specified in `VideoWriter` initialization.  If
1860        `input_format` is 'gray', the image must be 2D.  For the 'rgb'
1861        input_format, the image may be either 2D (interpreted as grayscale) or
1862        3D with three (R, G, B) channels.  For the 'yuv' input_format, the image
1863        must be 3D with three (Y, U, V) channels.
1864
1865    Raises:
1866      RuntimeError: If there is an error writing to the output file.
1867    """
1868    assert self._proc, 'Error: writing to an already closed context.'
1869    if issubclass(image.dtype.type, (np.floating, np.bool_)):
1870      image = to_type(image, self.dtype)
1871    if image.dtype != self.dtype:
1872      raise ValueError(f'Image type {image.dtype} != {self.dtype}.')
1873    if self.input_format == 'gray':
1874      if image.ndim != 2:
1875        raise ValueError(f'Image dimensions {image.shape} are not 2D.')
1876    else:
1877      if image.ndim == 2 and self.input_format == 'rgb':
1878        image = np.dstack((image, image, image))
1879      if not (image.ndim == 3 and image.shape[2] == 3):
1880        raise ValueError(f'Image dimensions {image.shape} are invalid.')
1881    if image.shape[:2] != self.shape:
1882      raise ValueError(
1883          f'Image dimensions {image.shape[:2]} do not match'
1884          f' those of the initialized video {self.shape}.'
1885      )
1886    if self.input_format == 'yuv':  # Convert from per-pixel YUV to planar YUV.
1887      image = np.moveaxis(image, 2, 0)
1888    data = image.tobytes()
1889    stdin = self._proc.stdin
1890    assert stdin is not None
1891    if stdin.write(data) != len(data):
1892      self._proc.wait()
1893      stderr = self._proc.stderr
1894      assert stderr is not None
1895      s = stderr.read().decode('utf-8')
1896      raise RuntimeError(f"Error writing '{self.path}': {s}")
1897
1898  def close(self) -> None:
1899    """Finishes writing the video.  (Called automatically at end of context.)"""
1900    if self._exit_stack is None:
1901      return
1902    # Unwinding the stack terminates the `ffmpeg` process, removes the
1903    # temporary audio file, and copies the encoded video to a remote `path`.
1904    # Raising within the `with` propagates the error into those contexts, so
1905    # an incomplete video is discarded rather than copied.
1906    with self._exit_stack:
1907      self._exit_stack = None
1908      if self._proc is not None:
1909        proc, self._proc = self._proc, None
1910        stdin = proc.stdin
1911        assert stdin is not None
1912        stdin.close()
1913        if proc.wait():
1914          stderr = proc.stderr
1915          assert stderr is not None
1916          s = stderr.read().decode('utf-8')
1917          raise RuntimeError(f"Error writing '{self.path}': {s}")
1918
1919
1920class _VideoArray(np.ndarray):
1921  """Wrapper to add a VideoMetadata `metadata` attribute to a numpy array."""
1922
1923  metadata: VideoMetadata | None
1924
1925  def __new__(
1926      cls: Type['_VideoArray'],
1927      input_array: _NDArray,
1928      metadata: VideoMetadata | None = None,
1929  ) -> '_VideoArray':
1930    obj: _VideoArray = np.asarray(input_array).view(cls)
1931    obj.metadata = metadata
1932    return obj
1933
1934  def __array_finalize__(self, obj: Any) -> None:
1935    if obj is None:
1936      return
1937    self.metadata = getattr(obj, 'metadata', None)
1938
1939
1940def read_video(path_or_url: _Path, **kwargs: Any) -> _VideoArray:
1941  """Returns an array containing all images read from a compressed video file.
1942
1943  >>> video = read_video('/tmp/river.mp4')
1944  >>> print(f'The framerate is {video.metadata.fps} frames/s.')
1945  >>> show_video(video)
1946
1947  >>> url = 'https://github.com/hhoppe/data/raw/main/video.mp4'
1948  >>> show_video(read_video(url))
1949
1950  Args:
1951    path_or_url: Input video file.
1952    **kwargs: Additional parameters for `VideoReader`.
1953
1954  Returns:
1955    A 4D `numpy` array with dimensions (frame, height, width, channel), or a 3D
1956    array if `output_format` is specified as 'gray'.  The returned array has an
1957    attribute `metadata` containing `VideoMetadata` information.  This enables
1958    `show_video` to retrieve the framerate in `metadata.fps`.  Note that the
1959    metadata attribute is lost in most subsequent `numpy` operations.
1960  """
1961  with VideoReader(path_or_url, **kwargs) as reader:
1962    return _VideoArray(np.array(tuple(reader)), metadata=reader.metadata)
1963
1964
1965def write_video(path: _Path, images: Iterable[_NDArray], **kwargs: Any) -> None:
1966  """Writes images to a compressed video file.
1967
1968  >>> video = moving_circle((480, 640), num_images=60)
1969  >>> write_video('/tmp/v.mp4', video, fps=60, qp=18)
1970  >>> show_video(read_video('/tmp/v.mp4'))
1971
1972  Args:
1973    path: Output video file.
1974    images: Iterable over video frames, e.g. a 4D array or a list of 2D or 3D
1975      arrays.
1976    **kwargs: Additional parameters for `VideoWriter`.
1977  """
1978  first_image, images = _peek_first(images)
1979  shape = first_image.shape[0], first_image.shape[1]
1980  dtype = first_image.dtype
1981  if dtype == bool:
1982    dtype = np.dtype(np.uint8)
1983  elif np.issubdtype(dtype, np.floating):
1984    dtype = np.dtype(np.uint16)
1985  kwargs = {'metadata': getattr(images, 'metadata', None), **kwargs}
1986  with VideoWriter(path, shape=shape, dtype=dtype, **kwargs) as writer:
1987    for image in images:
1988      writer.add_image(image)
1989
1990
1991def compress_video(
1992    images: Iterable[_NDArray], *, codec: str = 'h264', **kwargs: Any
1993) -> bytes:
1994  """Returns a buffer containing a compressed video.
1995
1996  The video container is 'gif' for 'gif' codec, 'webm' for 'vp9' codec,
1997  and mp4 otherwise.
1998
1999  >>> video = read_video('/tmp/river.mp4')
2000  >>> data = compress_video(video, bps=10_000_000)
2001  >>> print(len(data))
2002
2003  >>> data = compress_video(moving_circle((100, 100), num_images=10), fps=10)
2004
2005  Args:
2006    images: Iterable over video frames.
2007    codec: Compression algorithm as defined by `ffmpeg -codecs` (e.g., 'h264',
2008      'hevc', 'vp9', or 'gif').
2009    **kwargs: Additional parameters for `VideoWriter`.
2010
2011  Returns:
2012    A bytes buffer containing the compressed video.
2013  """
2014  suffix = _filename_suffix_from_codec(codec)
2015  with tempfile.TemporaryDirectory() as directory_name:
2016    tmp_path = pathlib.Path(directory_name) / f'file{suffix}'
2017    write_video(tmp_path, images, codec=codec, **kwargs)
2018    return tmp_path.read_bytes()
2019
2020
2021def decompress_video(data: bytes, **kwargs: Any) -> _NDArray:
2022  """Returns video images from an MP4-compressed data buffer."""
2023  with tempfile.TemporaryDirectory() as directory_name:
2024    tmp_path = pathlib.Path(directory_name) / 'file.mp4'
2025    tmp_path.write_bytes(data)
2026    return read_video(tmp_path, **kwargs)
2027
2028
2029def html_from_compressed_video(
2030    data: bytes,
2031    width: int,
2032    height: int,
2033    *,
2034    title: str | None = None,
2035    border: bool | str = False,
2036    loop: bool = True,
2037    autoplay: bool = True,
2038) -> str:
2039  """Returns an HTML string with a video tag containing H264-encoded data.
2040
2041  Args:
2042    data: MP4-compressed video bytes.
2043    width: Width of HTML video in pixels.
2044    height: Height of HTML video in pixels.
2045    title: Optional text shown centered above the video.
2046    border: If `bool`, whether to place a black boundary around the image, or if
2047      `str`, the boundary CSS style.
2048    loop: If True, the playback repeats forever.
2049    autoplay: If True, video playback starts without having to click.
2050  """
2051  b64 = base64.b64encode(data).decode('utf-8')
2052  if isinstance(border, str):
2053    border = f'{border}; '
2054  elif border:
2055    border = 'border:1px solid black; '
2056  else:
2057    border = ''
2058  options = (
2059      f'controls width="{width}" height="{height}"'
2060      f' style="{border}object-fit:cover;"'
2061      f'{" loop" if loop else ""}'
2062      f'{" autoplay muted" if autoplay else ""}'
2063  )
2064  s = f"""<video {options}>
2065      <source src="data:video/mp4;base64,{b64}" type="video/mp4"/>
2066      This browser does not support the video tag.
2067      </video>"""
2068  if title is not None:
2069    s = f"""<div style="display:flex; align-items:left;">
2070      <div style="display:flex; flex-direction:column; align-items:center;">
2071      <div>{title}</div><div>{s}</div></div></div>"""
2072  return s
2073
2074
2075def show_video(
2076    images: Iterable[_NDArray], *, title: str | None = None, **kwargs: Any
2077) -> str | None:
2078  """Displays a video in the IPython notebook and optionally saves it to a file.
2079
2080  See `show_videos`.
2081
2082  >>> video = read_video('https://github.com/hhoppe/data/raw/main/video.mp4')
2083  >>> show_video(video, title='River video')
2084
2085  >>> show_video(moving_circle((80, 80), num_images=10), fps=5, border=True)
2086
2087  >>> show_video(read_video('/tmp/river.mp4'))
2088
2089  Args:
2090    images: Iterable of video frames (e.g., a 4D array or a list of 2D or 3D
2091      arrays).
2092    title: Optional text shown centered above the video.
2093    **kwargs: See `show_videos`.
2094
2095  Returns:
2096    html string if `return_html` is `True`.
2097  """
2098  return show_videos([images], [title], **kwargs)
2099
2100
2101def show_videos(
2102    videos: Iterable[Iterable[_NDArray]] | Mapping[str, Iterable[_NDArray]],
2103    titles: Iterable[str | None] | None = None,
2104    *,
2105    width: int | None = None,
2106    height: int | None = None,
2107    downsample: bool = True,
2108    columns: int | None = None,
2109    fps: float | None = None,
2110    bps: int | None = None,
2111    qp: int | None = None,
2112    codec: str = 'h264',
2113    ylabel: str = '',
2114    html_class: str = 'show_videos',
2115    return_html: bool = False,
2116    audios: Iterable[_NDArray] | Mapping[str, _NDArray] | None = None,
2117    audio_sample_rate: int | None = None,
2118    audio_codec: str = 'aac',
2119    **kwargs: Any,
2120) -> str | None:
2121  """Displays a row of videos in the IPython notebook.
2122
2123  Creates HTML with `<video>` tags containing embedded H264-encoded bytestrings.
2124  If `codec` is set to 'gif', we instead use `<img>` tags containing embedded
2125  GIF-encoded bytestrings.  Note that the resulting GIF animations skip frames
2126  when the `fps` period is not a multiple of 10 ms units (GIF frame delay
2127  units).  Encoding at `fps` = 20.0, 25.0, or 50.0 works fine.
2128
2129  If a directory has been specified using `set_show_save_dir`, also saves each
2130  titled video to a file in that directory based on its title.
2131
2132  Args:
2133    videos: Iterable of videos, or dictionary of `{title: video}`.  Each video
2134      must be an iterable of images.  If a video object has a `metadata`
2135      (`VideoMetadata`) attribute, its `fps` field provides a default framerate.
2136    titles: Optional strings shown above the corresponding videos.
2137    width: Optional, overrides displayed width (in pixels).
2138    height: Optional, overrides displayed height (in pixels).
2139    downsample: If True, each video whose width or height is greater than the
2140      specified `width` or `height` is resampled to the display resolution. This
2141      improves antialiasing and reduces the size of the notebook.
2142    columns: Optional, maximum number of videos per row.
2143    fps: Frames-per-second framerate (default is 60.0 except 25.0 for GIF).
2144    bps: Bits-per-second bitrate (default None).
2145    qp: Quantization parameter for video compression quality (default None).
2146    codec: Compression algorithm; must be either 'h264' or 'gif'.
2147    ylabel: Text (rotated by 90 degrees) shown on the left of each row.
2148    html_class: CSS class name used in definition of HTML element.
2149    return_html: If `True` return the raw HTML `str` instead of displaying.
2150    audios: Optional iterable of audio tracks, or dictionary of `{title:
2151      audio}`; see `VideoWriter`.  Each track is attached to the corresponding
2152      video, so an iterable must have one entry (possibly None) per video.  With
2153      a dictionary, its keys must match the video titles exactly (a value may be
2154      None for no audio).  Audio is incompatible with `codec` 'gif'.  With the
2155      default `autoplay=True`, the video starts muted (browsers block unmuted
2156      autoplay); unmute it in the player controls, or pass `autoplay=False`.
2157    audio_sample_rate: Sample rate of the audio in Hz, shared by all tracks.
2158      Required if `audios` is provided.
2159    audio_codec: Audio compression algorithm (default 'aac'); see `VideoWriter`.
2160    **kwargs: Additional parameters (`border`, `loop`, `autoplay`) for
2161      `html_from_compressed_video`.
2162
2163  Returns:
2164    html string if `return_html` is `True`.
2165  """
2166  if isinstance(videos, Mapping):
2167    if titles is not None:
2168      raise ValueError(
2169          'Cannot have both a video dictionary and a titles parameter.'
2170      )
2171    list_titles = list(videos.keys())
2172    list_videos = list(videos.values())
2173  else:
2174    list_videos = list(cast('Iterable[_NDArray]', videos))
2175    list_titles = [None] * len(list_videos) if titles is None else list(titles)
2176    if len(list_videos) != len(list_titles):
2177      raise ValueError(
2178          'Number of videos does not match number of titles'
2179          f' ({len(list_videos)} vs {len(list_titles)}).'
2180      )
2181
2182  if audios is None:
2183    list_audios = [None] * len(list_videos)
2184  elif isinstance(audios, Mapping):
2185    missing = set(list_titles).difference(audios)
2186    extra = set(audios).difference(list_titles)
2187    if missing or extra:
2188      raise ValueError(
2189          'The audios dictionary keys must match the video titles (use None as'
2190          f' the value for no audio); missing: {sorted(missing, key=str)},'
2191          f' extra: {sorted(extra, key=str)}.'
2192      )
2193    list_audios = [audios.get(title) for title in list_titles]  # pyrefly: ignore[bad-argument-type]
2194  else:
2195    list_audios = list(audios)
2196
2197  if len(list_videos) != len(list_audios):
2198    raise ValueError(
2199        'Number of videos does not match number of audio'
2200        f' ({len(list_videos)} vs {len(list_audios)}).'
2201    )
2202
2203  if codec not in {'h264', 'gif'}:
2204    raise ValueError(f'Codec {codec} is neither h264 or gif.')
2205
2206  html_strings = []
2207  for video, title, video_audio in zip(list_videos, list_titles, list_audios):
2208    metadata: VideoMetadata | None = getattr(video, 'metadata', None)
2209    first_image, video = _peek_first(video)
2210    w, h = _get_width_height(width, height, first_image.shape[:2])
2211    if downsample and (w < first_image.shape[1] or h < first_image.shape[0]):
2212      # Not resize_video() because each image may have different depth and type.
2213      video = [resize_image(image, (h, w)) for image in video]
2214      first_image = video[0]
2215    data = compress_video(
2216        video,
2217        metadata=metadata,
2218        fps=fps,
2219        bps=bps,
2220        qp=qp,
2221        codec=codec,
2222        audio=video_audio,
2223        audio_sample_rate=audio_sample_rate,
2224        audio_codec=audio_codec,
2225    )
2226    if title is not None and _config.show_save_dir:
2227      suffix = _filename_suffix_from_codec(codec)
2228      path = pathlib.Path(_config.show_save_dir) / f'{title}{suffix}'
2229      with _open(path, mode='wb') as f:
2230        f.write(data)
2231    if codec == 'gif':
2232      pixelated = h > first_image.shape[0] or w > first_image.shape[1]
2233      html_string = html_from_compressed_image(
2234          data, w, h, title=title, fmt='gif', pixelated=pixelated, **kwargs  # pyrefly: ignore[bad-argument-type]
2235      )
2236    else:
2237      html_string = html_from_compressed_video(
2238          data, w, h, title=title, **kwargs  # pyrefly: ignore[bad-argument-type]
2239      )
2240    html_strings.append(html_string)
2241
2242  # Create single-row tables each with no more than 'columns' elements.
2243  table_strings = []
2244  for row_html_strings in _chunked(html_strings, columns):
2245    td = '<td style="padding:1px;">'
2246    s = ''.join(f'{td}{e}</td>' for e in row_html_strings)
2247    if ylabel:
2248      style = 'writing-mode:vertical-lr; transform:rotate(180deg);'
2249      s = f'{td}<span style="{style}">{ylabel}</span></td>' + s
2250    table_strings.append(
2251        f'<table class="{html_class}"'
2252        f' style="border-spacing:0px;"><tr>{s}</tr></table>'
2253    )
2254  s = ''.join(table_strings)
2255  if return_html:
2256    return s
2257  _display_html(s)
2258  return None
2259
2260
2261# Local Variables:
2262# fill-column: 80
2263# End:
def show_image( image: ArrayLike, *, title: str | None = None, **kwargs: Any) -> str | None:
1005def show_image(
1006    image: _ArrayLike, *, title: str | None = None, **kwargs: Any
1007) -> str | None:
1008  """Displays an image in the notebook and optionally saves it to a file.
1009
1010  See `show_images`.
1011
1012  >>> show_image(np.random.rand(100, 100))
1013  >>> show_image(np.random.randint(0, 256, size=(80, 80, 3), dtype='uint8'))
1014  >>> show_image(np.random.rand(10, 10) - 0.5, cmap='bwr', height=100)
1015  >>> show_image(read_image('/tmp/image.png'))
1016  >>> url = 'https://github.com/hhoppe/data/raw/main/image.png'
1017  >>> show_image(read_image(url))
1018
1019  Args:
1020    image: 2D array-like, or 3D array-like with 1, 3, or 4 channels.
1021    title: Optional text shown centered above the image.
1022    **kwargs: See `show_images`.
1023
1024  Returns:
1025    html string if `return_html` is `True`.
1026  """
1027  return show_images([np.asarray(image)], [title], **kwargs)

Displays an image in the notebook and optionally saves it to a file.

See show_images.

>>> show_image(np.random.rand(100, 100))
>>> show_image(np.random.randint(0, 256, size=(80, 80, 3), dtype='uint8'))
>>> show_image(np.random.rand(10, 10) - 0.5, cmap='bwr', height=100)
>>> show_image(read_image('/tmp/image.png'))
>>> url = 'https://github.com/hhoppe/data/raw/main/image.png'
>>> show_image(read_image(url))
Arguments:
  • image: 2D array-like, or 3D array-like with 1, 3, or 4 channels.
  • title: Optional text shown centered above the image.
  • **kwargs: See show_images.
Returns:

html string if return_html is True.

def show_images( images: Iterable[ArrayLike] | Mapping[str, ArrayLike], titles: Iterable[str | None] | None = None, *, width: int | None = None, height: int | None = None, downsample: bool = True, columns: int | None = None, vmin: float | None = None, vmax: float | None = None, cmap: str | Callable[[ArrayLike], np.ndarray] = 'gray', border: bool | str = False, ylabel: str = '', html_class: str = 'show_images', pixelated: bool | None = None, return_html: bool = False) -> str | None:
1030def show_images(
1031    images: Iterable[_ArrayLike] | Mapping[str, _ArrayLike],
1032    titles: Iterable[str | None] | None = None,
1033    *,
1034    width: int | None = None,
1035    height: int | None = None,
1036    downsample: bool = True,
1037    columns: int | None = None,
1038    vmin: float | None = None,
1039    vmax: float | None = None,
1040    cmap: str | Callable[[_ArrayLike], _NDArray] = 'gray',
1041    border: bool | str = False,
1042    ylabel: str = '',
1043    html_class: str = 'show_images',
1044    pixelated: bool | None = None,
1045    return_html: bool = False,
1046) -> str | None:
1047  """Displays a row of images in the IPython/Jupyter notebook.
1048
1049  If a directory has been specified using `set_show_save_dir`, also saves each
1050  titled image to a file in that directory based on its title.
1051
1052  >>> image1, image2 = np.random.rand(64, 64, 3), color_ramp((64, 64))
1053  >>> show_images([image1, image2])
1054  >>> show_images({'random image': image1, 'color ramp': image2}, height=128)
1055  >>> show_images([image1, image2] * 5, columns=4, border=True)
1056
1057  Args:
1058    images: Iterable of images, or dictionary of `{title: image}`.  Each image
1059      must be either a 2D array or a 3D array with 1, 3, or 4 channels.
1060    titles: Optional strings shown above the corresponding images.
1061    width: Optional, overrides displayed width (in pixels).
1062    height: Optional, overrides displayed height (in pixels).
1063    downsample: If True, each image whose width or height is greater than the
1064      specified `width` or `height` is resampled to the display resolution. This
1065      improves antialiasing and reduces the size of the notebook.
1066    columns: Optional, maximum number of images per row.
1067    vmin: For single-channel image, explicit min value for display.
1068    vmax: For single-channel image, explicit max value for display.
1069    cmap: For single-channel image, `pyplot` color map or callable to map 1D to
1070      3D color.
1071    border: If `bool`, whether to place a black boundary around the image, or if
1072      `str`, the boundary CSS style.
1073    ylabel: Text (rotated by 90 degrees) shown on the left of each row.
1074    html_class: CSS class name used in definition of HTML element.
1075    pixelated: If True, sets the CSS style to 'image-rendering: pixelated;'; if
1076      False, sets 'image-rendering: auto'; if None, uses pixelated rendering
1077      only on images for which `width` or `height` introduces magnification.
1078    return_html: If `True` return the raw HTML `str` instead of displaying.
1079
1080  Returns:
1081    html string if `return_html` is `True`.
1082  """
1083  if isinstance(images, Mapping):
1084    if titles is not None:
1085      raise ValueError('Cannot have images dictionary and titles parameter.')
1086    list_titles, list_images = list(images.keys()), list(images.values())
1087  else:
1088    list_images = list(images)
1089    list_titles = [None] * len(list_images) if titles is None else list(titles)
1090    if len(list_images) != len(list_titles):
1091      raise ValueError(
1092          'Number of images does not match number of titles'
1093          f' ({len(list_images)} vs {len(list_titles)}).'
1094      )
1095
1096  list_images = [
1097      _ensure_mapped_to_rgb(image, vmin=vmin, vmax=vmax, cmap=cmap)
1098      for image in list_images
1099  ]
1100
1101  def maybe_downsample(image: _NDArray) -> _NDArray:
1102    shape = image.shape[0], image.shape[1]
1103    w, h = _get_width_height(width, height, shape)
1104    if w < shape[1] or h < shape[0]:
1105      image = resize_image(image, (h, w))
1106    return image
1107
1108  if downsample:
1109    list_images = [maybe_downsample(image) for image in list_images]
1110  png_datas = [compress_image(to_uint8(image)) for image in list_images]
1111
1112  for title, png_data in zip(list_titles, png_datas):
1113    if title is not None and _config.show_save_dir:
1114      path = pathlib.Path(_config.show_save_dir) / f'{title}.png'
1115      with _open(path, mode='wb') as f:
1116        f.write(png_data)
1117
1118  def html_from_compressed_images() -> str:
1119    html_strings = []
1120    for image, title, png_data in zip(list_images, list_titles, png_datas):
1121      w, h = _get_width_height(width, height, image.shape[:2])  # pyrefly: ignore[missing-attribute]
1122      magnified = h > image.shape[0] or w > image.shape[1]  # pyrefly: ignore[missing-attribute]
1123      pixelated2 = pixelated if pixelated is not None else magnified
1124      html_strings.append(
1125          html_from_compressed_image(
1126              png_data, w, h, title=title, border=border, pixelated=pixelated2  # pyrefly: ignore[bad-argument-type]
1127          )
1128      )
1129    # Create single-row tables each with no more than 'columns' elements.
1130    table_strings = []
1131    for row_html_strings in _chunked(html_strings, columns):
1132      td = '<td style="padding:1px;">'
1133      s = ''.join(f'{td}{e}</td>' for e in row_html_strings)
1134      if ylabel:
1135        style = 'writing-mode:vertical-lr; transform:rotate(180deg);'
1136        s = f'{td}<span style="{style}">{ylabel}</span></td>' + s
1137      table_strings.append(
1138          f'<table class="{html_class}"'
1139          f' style="border-spacing:0px;"><tr>{s}</tr></table>'
1140      )
1141    return ''.join(table_strings)
1142
1143  s = html_from_compressed_images()
1144  while len(s) > _IPYTHON_HTML_SIZE_LIMIT * 0.5:
1145    warnings.warn('mediapy: subsampling images to reduce HTML size')
1146    list_images = [image[::2, ::2] for image in list_images]
1147    png_datas = [compress_image(to_uint8(image)) for image in list_images]
1148    s = html_from_compressed_images()
1149  if return_html:
1150    return s
1151  _display_html(s)
1152  return None

Displays a row of images in the IPython/Jupyter notebook.

If a directory has been specified using set_show_save_dir, also saves each titled image to a file in that directory based on its title.

>>> image1, image2 = np.random.rand(64, 64, 3), color_ramp((64, 64))
>>> show_images([image1, image2])
>>> show_images({'random image': image1, 'color ramp': image2}, height=128)
>>> show_images([image1, image2] * 5, columns=4, border=True)
Arguments:
  • images: Iterable of images, or dictionary of {title: image}. Each image must be either a 2D array or a 3D array with 1, 3, or 4 channels.
  • titles: Optional strings shown above the corresponding images.
  • width: Optional, overrides displayed width (in pixels).
  • height: Optional, overrides displayed height (in pixels).
  • downsample: If True, each image whose width or height is greater than the specified width or height is resampled to the display resolution. This improves antialiasing and reduces the size of the notebook.
  • columns: Optional, maximum number of images per row.
  • vmin: For single-channel image, explicit min value for display.
  • vmax: For single-channel image, explicit max value for display.
  • cmap: For single-channel image, pyplot color map or callable to map 1D to 3D color.
  • border: If bool, whether to place a black boundary around the image, or if str, the boundary CSS style.
  • ylabel: Text (rotated by 90 degrees) shown on the left of each row.
  • html_class: CSS class name used in definition of HTML element.
  • pixelated: If True, sets the CSS style to 'image-rendering: pixelated;'; if False, sets 'image-rendering: auto'; if None, uses pixelated rendering only on images for which width or height introduces magnification.
  • return_html: If True return the raw HTML str instead of displaying.
Returns:

html string if return_html is True.

def compare_images( images: Iterable[ArrayLike], *, vmin: float | None = None, vmax: float | None = None, cmap: str | Callable[[ArrayLike], np.ndarray] = 'gray') -> None:
1155def compare_images(
1156    images: Iterable[_ArrayLike],
1157    *,
1158    vmin: float | None = None,
1159    vmax: float | None = None,
1160    cmap: str | Callable[[_ArrayLike], _NDArray] = 'gray',
1161) -> None:
1162  """Compare two images using an interactive slider.
1163
1164  Displays an HTML slider component to interactively swipe between two images.
1165  The slider functionality requires that the web browser have Internet access.
1166  See additional info in `https://github.com/sneas/img-comparison-slider`.
1167
1168  >>> image1, image2 = np.random.rand(64, 64, 3), color_ramp((64, 64))
1169  >>> compare_images([image1, image2])
1170
1171  Args:
1172    images: Iterable of images.  Each image must be either a 2D array or a 3D
1173      array with 1, 3, or 4 channels.  There must be exactly two images.
1174    vmin: For single-channel image, explicit min value for display.
1175    vmax: For single-channel image, explicit max value for display.
1176    cmap: For single-channel image, `pyplot` color map or callable to map 1D to
1177      3D color.
1178  """
1179  list_images = [
1180      _ensure_mapped_to_rgb(image, vmin=vmin, vmax=vmax, cmap=cmap)
1181      for image in images
1182  ]
1183  if len(list_images) != 2:
1184    raise ValueError('The number of images must be 2.')
1185  png_datas = [compress_image(to_uint8(image)) for image in list_images]
1186  b64_1, b64_2 = [
1187      base64.b64encode(png_data).decode('utf-8') for png_data in png_datas
1188  ]
1189  s = _IMAGE_COMPARISON_HTML.replace('{b64_1}', b64_1).replace('{b64_2}', b64_2)
1190  _display_html(s)

Compare two images using an interactive slider.

Displays an HTML slider component to interactively swipe between two images. The slider functionality requires that the web browser have Internet access. See additional info in https://github.com/sneas/img-comparison-slider.

>>> image1, image2 = np.random.rand(64, 64, 3), color_ramp((64, 64))
>>> compare_images([image1, image2])
Arguments:
  • images: Iterable of images. Each image must be either a 2D array or a 3D array with 1, 3, or 4 channels. There must be exactly two images.
  • vmin: For single-channel image, explicit min value for display.
  • vmax: For single-channel image, explicit max value for display.
  • cmap: For single-channel image, pyplot color map or callable to map 1D to 3D color.
def show_video( images: Iterable[np.ndarray], *, title: str | None = None, **kwargs: Any) -> str | None:
2076def show_video(
2077    images: Iterable[_NDArray], *, title: str | None = None, **kwargs: Any
2078) -> str | None:
2079  """Displays a video in the IPython notebook and optionally saves it to a file.
2080
2081  See `show_videos`.
2082
2083  >>> video = read_video('https://github.com/hhoppe/data/raw/main/video.mp4')
2084  >>> show_video(video, title='River video')
2085
2086  >>> show_video(moving_circle((80, 80), num_images=10), fps=5, border=True)
2087
2088  >>> show_video(read_video('/tmp/river.mp4'))
2089
2090  Args:
2091    images: Iterable of video frames (e.g., a 4D array or a list of 2D or 3D
2092      arrays).
2093    title: Optional text shown centered above the video.
2094    **kwargs: See `show_videos`.
2095
2096  Returns:
2097    html string if `return_html` is `True`.
2098  """
2099  return show_videos([images], [title], **kwargs)

Displays a video in the IPython notebook and optionally saves it to a file.

See show_videos.

>>> video = read_video('https://github.com/hhoppe/data/raw/main/video.mp4')
>>> show_video(video, title='River video')
>>> show_video(moving_circle((80, 80), num_images=10), fps=5, border=True)
>>> show_video(read_video('/tmp/river.mp4'))
Arguments:
  • images: Iterable of video frames (e.g., a 4D array or a list of 2D or 3D arrays).
  • title: Optional text shown centered above the video.
  • **kwargs: See show_videos.
Returns:

html string if return_html is True.

def show_videos( videos: Iterable[Iterable[np.ndarray]] | Mapping[str, Iterable[np.ndarray]], titles: Iterable[str | None] | None = None, *, width: int | None = None, height: int | None = None, downsample: bool = True, columns: int | None = None, fps: float | None = None, bps: int | None = None, qp: int | None = None, codec: str = 'h264', ylabel: str = '', html_class: str = 'show_videos', return_html: bool = False, audios: Iterable[np.ndarray] | Mapping[str, np.ndarray] | None = None, audio_sample_rate: int | None = None, audio_codec: str = 'aac', **kwargs: Any) -> str | None:
2102def show_videos(
2103    videos: Iterable[Iterable[_NDArray]] | Mapping[str, Iterable[_NDArray]],
2104    titles: Iterable[str | None] | None = None,
2105    *,
2106    width: int | None = None,
2107    height: int | None = None,
2108    downsample: bool = True,
2109    columns: int | None = None,
2110    fps: float | None = None,
2111    bps: int | None = None,
2112    qp: int | None = None,
2113    codec: str = 'h264',
2114    ylabel: str = '',
2115    html_class: str = 'show_videos',
2116    return_html: bool = False,
2117    audios: Iterable[_NDArray] | Mapping[str, _NDArray] | None = None,
2118    audio_sample_rate: int | None = None,
2119    audio_codec: str = 'aac',
2120    **kwargs: Any,
2121) -> str | None:
2122  """Displays a row of videos in the IPython notebook.
2123
2124  Creates HTML with `<video>` tags containing embedded H264-encoded bytestrings.
2125  If `codec` is set to 'gif', we instead use `<img>` tags containing embedded
2126  GIF-encoded bytestrings.  Note that the resulting GIF animations skip frames
2127  when the `fps` period is not a multiple of 10 ms units (GIF frame delay
2128  units).  Encoding at `fps` = 20.0, 25.0, or 50.0 works fine.
2129
2130  If a directory has been specified using `set_show_save_dir`, also saves each
2131  titled video to a file in that directory based on its title.
2132
2133  Args:
2134    videos: Iterable of videos, or dictionary of `{title: video}`.  Each video
2135      must be an iterable of images.  If a video object has a `metadata`
2136      (`VideoMetadata`) attribute, its `fps` field provides a default framerate.
2137    titles: Optional strings shown above the corresponding videos.
2138    width: Optional, overrides displayed width (in pixels).
2139    height: Optional, overrides displayed height (in pixels).
2140    downsample: If True, each video whose width or height is greater than the
2141      specified `width` or `height` is resampled to the display resolution. This
2142      improves antialiasing and reduces the size of the notebook.
2143    columns: Optional, maximum number of videos per row.
2144    fps: Frames-per-second framerate (default is 60.0 except 25.0 for GIF).
2145    bps: Bits-per-second bitrate (default None).
2146    qp: Quantization parameter for video compression quality (default None).
2147    codec: Compression algorithm; must be either 'h264' or 'gif'.
2148    ylabel: Text (rotated by 90 degrees) shown on the left of each row.
2149    html_class: CSS class name used in definition of HTML element.
2150    return_html: If `True` return the raw HTML `str` instead of displaying.
2151    audios: Optional iterable of audio tracks, or dictionary of `{title:
2152      audio}`; see `VideoWriter`.  Each track is attached to the corresponding
2153      video, so an iterable must have one entry (possibly None) per video.  With
2154      a dictionary, its keys must match the video titles exactly (a value may be
2155      None for no audio).  Audio is incompatible with `codec` 'gif'.  With the
2156      default `autoplay=True`, the video starts muted (browsers block unmuted
2157      autoplay); unmute it in the player controls, or pass `autoplay=False`.
2158    audio_sample_rate: Sample rate of the audio in Hz, shared by all tracks.
2159      Required if `audios` is provided.
2160    audio_codec: Audio compression algorithm (default 'aac'); see `VideoWriter`.
2161    **kwargs: Additional parameters (`border`, `loop`, `autoplay`) for
2162      `html_from_compressed_video`.
2163
2164  Returns:
2165    html string if `return_html` is `True`.
2166  """
2167  if isinstance(videos, Mapping):
2168    if titles is not None:
2169      raise ValueError(
2170          'Cannot have both a video dictionary and a titles parameter.'
2171      )
2172    list_titles = list(videos.keys())
2173    list_videos = list(videos.values())
2174  else:
2175    list_videos = list(cast('Iterable[_NDArray]', videos))
2176    list_titles = [None] * len(list_videos) if titles is None else list(titles)
2177    if len(list_videos) != len(list_titles):
2178      raise ValueError(
2179          'Number of videos does not match number of titles'
2180          f' ({len(list_videos)} vs {len(list_titles)}).'
2181      )
2182
2183  if audios is None:
2184    list_audios = [None] * len(list_videos)
2185  elif isinstance(audios, Mapping):
2186    missing = set(list_titles).difference(audios)
2187    extra = set(audios).difference(list_titles)
2188    if missing or extra:
2189      raise ValueError(
2190          'The audios dictionary keys must match the video titles (use None as'
2191          f' the value for no audio); missing: {sorted(missing, key=str)},'
2192          f' extra: {sorted(extra, key=str)}.'
2193      )
2194    list_audios = [audios.get(title) for title in list_titles]  # pyrefly: ignore[bad-argument-type]
2195  else:
2196    list_audios = list(audios)
2197
2198  if len(list_videos) != len(list_audios):
2199    raise ValueError(
2200        'Number of videos does not match number of audio'
2201        f' ({len(list_videos)} vs {len(list_audios)}).'
2202    )
2203
2204  if codec not in {'h264', 'gif'}:
2205    raise ValueError(f'Codec {codec} is neither h264 or gif.')
2206
2207  html_strings = []
2208  for video, title, video_audio in zip(list_videos, list_titles, list_audios):
2209    metadata: VideoMetadata | None = getattr(video, 'metadata', None)
2210    first_image, video = _peek_first(video)
2211    w, h = _get_width_height(width, height, first_image.shape[:2])
2212    if downsample and (w < first_image.shape[1] or h < first_image.shape[0]):
2213      # Not resize_video() because each image may have different depth and type.
2214      video = [resize_image(image, (h, w)) for image in video]
2215      first_image = video[0]
2216    data = compress_video(
2217        video,
2218        metadata=metadata,
2219        fps=fps,
2220        bps=bps,
2221        qp=qp,
2222        codec=codec,
2223        audio=video_audio,
2224        audio_sample_rate=audio_sample_rate,
2225        audio_codec=audio_codec,
2226    )
2227    if title is not None and _config.show_save_dir:
2228      suffix = _filename_suffix_from_codec(codec)
2229      path = pathlib.Path(_config.show_save_dir) / f'{title}{suffix}'
2230      with _open(path, mode='wb') as f:
2231        f.write(data)
2232    if codec == 'gif':
2233      pixelated = h > first_image.shape[0] or w > first_image.shape[1]
2234      html_string = html_from_compressed_image(
2235          data, w, h, title=title, fmt='gif', pixelated=pixelated, **kwargs  # pyrefly: ignore[bad-argument-type]
2236      )
2237    else:
2238      html_string = html_from_compressed_video(
2239          data, w, h, title=title, **kwargs  # pyrefly: ignore[bad-argument-type]
2240      )
2241    html_strings.append(html_string)
2242
2243  # Create single-row tables each with no more than 'columns' elements.
2244  table_strings = []
2245  for row_html_strings in _chunked(html_strings, columns):
2246    td = '<td style="padding:1px;">'
2247    s = ''.join(f'{td}{e}</td>' for e in row_html_strings)
2248    if ylabel:
2249      style = 'writing-mode:vertical-lr; transform:rotate(180deg);'
2250      s = f'{td}<span style="{style}">{ylabel}</span></td>' + s
2251    table_strings.append(
2252        f'<table class="{html_class}"'
2253        f' style="border-spacing:0px;"><tr>{s}</tr></table>'
2254    )
2255  s = ''.join(table_strings)
2256  if return_html:
2257    return s
2258  _display_html(s)
2259  return None

Displays a row of videos in the IPython notebook.

Creates HTML with <video> tags containing embedded H264-encoded bytestrings. If codec is set to 'gif', we instead use <img> tags containing embedded GIF-encoded bytestrings. Note that the resulting GIF animations skip frames when the fps period is not a multiple of 10 ms units (GIF frame delay units). Encoding at fps = 20.0, 25.0, or 50.0 works fine.

If a directory has been specified using set_show_save_dir, also saves each titled video to a file in that directory based on its title.

Arguments:
  • videos: Iterable of videos, or dictionary of {title: video}. Each video must be an iterable of images. If a video object has a metadata (VideoMetadata) attribute, its fps field provides a default framerate.
  • titles: Optional strings shown above the corresponding videos.
  • width: Optional, overrides displayed width (in pixels).
  • height: Optional, overrides displayed height (in pixels).
  • downsample: If True, each video whose width or height is greater than the specified width or height is resampled to the display resolution. This improves antialiasing and reduces the size of the notebook.
  • columns: Optional, maximum number of videos per row.
  • fps: Frames-per-second framerate (default is 60.0 except 25.0 for GIF).
  • bps: Bits-per-second bitrate (default None).
  • qp: Quantization parameter for video compression quality (default None).
  • codec: Compression algorithm; must be either 'h264' or 'gif'.
  • ylabel: Text (rotated by 90 degrees) shown on the left of each row.
  • html_class: CSS class name used in definition of HTML element.
  • return_html: If True return the raw HTML str instead of displaying.
  • audios: Optional iterable of audio tracks, or dictionary of {title: audio}; see VideoWriter. Each track is attached to the corresponding video, so an iterable must have one entry (possibly None) per video. With a dictionary, its keys must match the video titles exactly (a value may be None for no audio). Audio is incompatible with codec 'gif'. With the default autoplay=True, the video starts muted (browsers block unmuted autoplay); unmute it in the player controls, or pass autoplay=False.
  • audio_sample_rate: Sample rate of the audio in Hz, shared by all tracks. Required if audios is provided.
  • audio_codec: Audio compression algorithm (default 'aac'); see VideoWriter.
  • **kwargs: Additional parameters (border, loop, autoplay) for html_from_compressed_video.
Returns:

html string if return_html is True.

def read_image( path_or_url: str | os.PathLike[str], *, apply_exif_transpose: bool = True, dtype: DTypeLike = None) -> np.ndarray:
797def read_image(
798    path_or_url: _Path,
799    *,
800    apply_exif_transpose: bool = True,
801    dtype: _DTypeLike = None,  # pyrefly: ignore[bad-function-definition]
802) -> _NDArray:
803  """Returns an image read from a file path or URL.
804
805  Decoding is performed using `PIL`, which supports `uint8` images with 1, 3,
806  or 4 channels and `uint16` images with a single channel.
807
808  Args:
809    path_or_url: Path of input file.
810    apply_exif_transpose: If True, rotate image according to EXIF orientation.
811    dtype: Data type of the returned array.  If None, `np.uint8` or `np.uint16`
812      is inferred automatically.
813  """
814  data = read_contents(path_or_url)
815  return decompress_image(data, dtype, apply_exif_transpose)

Returns an image read from a file path or URL.

Decoding is performed using PIL, which supports uint8 images with 1, 3, or 4 channels and uint16 images with a single channel.

Arguments:
  • path_or_url: Path of input file.
  • apply_exif_transpose: If True, rotate image according to EXIF orientation.
  • dtype: Data type of the returned array. If None, np.uint8 or np.uint16 is inferred automatically.
def write_image( path: str | os.PathLike[str], image: ArrayLike, fmt: str = 'png', **kwargs: Any) -> None:
818def write_image(
819    path: _Path, image: _ArrayLike, fmt: str = 'png', **kwargs: Any
820) -> None:
821  """Writes an image to a file.
822
823  Encoding is performed using `PIL`, which supports `uint8` images with 1, 3,
824  or 4 channels and `uint16` images with a single channel.
825
826  File format is explicitly provided by `fmt` and not inferred by `path`.
827
828  Args:
829    path: Path of output file.
830    image: Array-like object.  If its type is float, it is converted to np.uint8
831      using `to_uint8` (thus clamping to the input to the range [0.0, 1.0]).
832      Otherwise it must be np.uint8 or np.uint16.
833    fmt: Desired compression encoding, e.g. 'png'.
834    **kwargs: Additional parameters for `PIL.Image.save()`.
835  """
836  image = _as_valid_media_array(image)
837  if np.issubdtype(image.dtype, np.floating):
838    image = to_uint8(image)
839  with _open(path, 'wb') as f:
840    _pil_image(image).save(f, format=fmt, **kwargs)

Writes an image to a file.

Encoding is performed using PIL, which supports uint8 images with 1, 3, or 4 channels and uint16 images with a single channel.

File format is explicitly provided by fmt and not inferred by path.

Arguments:
  • path: Path of output file.
  • image: Array-like object. If its type is float, it is converted to np.uint8 using to_uint8 (thus clamping to the input to the range [0.0, 1.0]). Otherwise it must be np.uint8 or np.uint16.
  • fmt: Desired compression encoding, e.g. 'png'.
  • **kwargs: Additional parameters for PIL.Image.save().
def read_video( path_or_url: str | os.PathLike[str], **kwargs: Any) -> mediapy._VideoArray:
1941def read_video(path_or_url: _Path, **kwargs: Any) -> _VideoArray:
1942  """Returns an array containing all images read from a compressed video file.
1943
1944  >>> video = read_video('/tmp/river.mp4')
1945  >>> print(f'The framerate is {video.metadata.fps} frames/s.')
1946  >>> show_video(video)
1947
1948  >>> url = 'https://github.com/hhoppe/data/raw/main/video.mp4'
1949  >>> show_video(read_video(url))
1950
1951  Args:
1952    path_or_url: Input video file.
1953    **kwargs: Additional parameters for `VideoReader`.
1954
1955  Returns:
1956    A 4D `numpy` array with dimensions (frame, height, width, channel), or a 3D
1957    array if `output_format` is specified as 'gray'.  The returned array has an
1958    attribute `metadata` containing `VideoMetadata` information.  This enables
1959    `show_video` to retrieve the framerate in `metadata.fps`.  Note that the
1960    metadata attribute is lost in most subsequent `numpy` operations.
1961  """
1962  with VideoReader(path_or_url, **kwargs) as reader:
1963    return _VideoArray(np.array(tuple(reader)), metadata=reader.metadata)

Returns an array containing all images read from a compressed video file.

>>> video = read_video('/tmp/river.mp4')
>>> print(f'The framerate is {video.metadata.fps} frames/s.')
>>> show_video(video)
>>> url = 'https://github.com/hhoppe/data/raw/main/video.mp4'
>>> show_video(read_video(url))
Arguments:
  • path_or_url: Input video file.
  • **kwargs: Additional parameters for VideoReader.
Returns:

A 4D numpy array with dimensions (frame, height, width, channel), or a 3D array if output_format is specified as 'gray'. The returned array has an attribute metadata containing VideoMetadata information. This enables show_video to retrieve the framerate in metadata.fps. Note that the metadata attribute is lost in most subsequent numpy operations.

def write_video( path: str | os.PathLike[str], images: Iterable[np.ndarray], **kwargs: Any) -> None:
1966def write_video(path: _Path, images: Iterable[_NDArray], **kwargs: Any) -> None:
1967  """Writes images to a compressed video file.
1968
1969  >>> video = moving_circle((480, 640), num_images=60)
1970  >>> write_video('/tmp/v.mp4', video, fps=60, qp=18)
1971  >>> show_video(read_video('/tmp/v.mp4'))
1972
1973  Args:
1974    path: Output video file.
1975    images: Iterable over video frames, e.g. a 4D array or a list of 2D or 3D
1976      arrays.
1977    **kwargs: Additional parameters for `VideoWriter`.
1978  """
1979  first_image, images = _peek_first(images)
1980  shape = first_image.shape[0], first_image.shape[1]
1981  dtype = first_image.dtype
1982  if dtype == bool:
1983    dtype = np.dtype(np.uint8)
1984  elif np.issubdtype(dtype, np.floating):
1985    dtype = np.dtype(np.uint16)
1986  kwargs = {'metadata': getattr(images, 'metadata', None), **kwargs}
1987  with VideoWriter(path, shape=shape, dtype=dtype, **kwargs) as writer:
1988    for image in images:
1989      writer.add_image(image)

Writes images to a compressed video file.

>>> video = moving_circle((480, 640), num_images=60)
>>> write_video('/tmp/v.mp4', video, fps=60, qp=18)
>>> show_video(read_video('/tmp/v.mp4'))
Arguments:
  • path: Output video file.
  • images: Iterable over video frames, e.g. a 4D array or a list of 2D or 3D arrays.
  • **kwargs: Additional parameters for VideoWriter.
class VideoReader(_VideoIO):
1415class VideoReader(_VideoIO):
1416  """Context to read a compressed video as an iterable over its images.
1417
1418  >>> with VideoReader('/tmp/river.mp4') as reader:
1419  ...   print(f'Video has {reader.num_images} images with shape={reader.shape},'
1420  ...         f' at {reader.fps} frames/sec and {reader.bps} bits/sec.')
1421  ...   for image in reader:
1422  ...     print(image.shape)
1423
1424  >>> with VideoReader('/tmp/river.mp4') as reader:
1425  ...   video = np.array(tuple(reader))
1426
1427  >>> url = 'https://github.com/hhoppe/data/raw/main/video.mp4'
1428  >>> with VideoReader(url) as reader:
1429  ...   show_video(reader)
1430
1431  Attributes:
1432    path_or_url: Location of input video.
1433    output_format: Format of output images (default 'rgb').  If 'rgb', each
1434      image has shape=(height, width, 3) with R, G, B values.  If 'yuv', each
1435      image has shape=(height, width, 3) with Y, U, V values.  If 'gray', each
1436      image has shape=(height, width).
1437    dtype: Data type for output images.  The default is `np.uint8`.  Use of
1438      `np.uint16` allows reading 10-bit or 12-bit data without precision loss.
1439    metadata: Object storing the information retrieved from the video header.
1440      Its attributes are copied as attributes in this class.
1441    num_images: Number of frames that is expected from the video stream.  This
1442      is estimated from the framerate and the duration stored in the video
1443      header, so it might be inexact.
1444    shape: The dimensions (height, width) of each video frame.
1445    fps: The framerate in frames per second.
1446    bps: The estimated bitrate of the video stream in bits per second, retrieved
1447      from the video header.
1448    stream_index: The stream index to read from. The default is 0.
1449    sandbox_max_run_time_secs: The maximum time in seconds to run the sandbox.
1450      If None, the default limit is 30 minutes. Unused in open source.
1451  """
1452
1453  path_or_url: _Path
1454  output_format: str
1455  dtype: _DType
1456  metadata: VideoMetadata
1457  num_images: int
1458  shape: tuple[int, int]
1459  fps: float
1460  bps: int | None
1461  stream_index: int
1462  _num_bytes_per_image: int
1463
1464  def __init__(
1465      self,
1466      path_or_url: _Path,
1467      *,
1468      stream_index: int = 0,
1469      output_format: str = 'rgb',
1470      dtype: _DTypeLike = np.uint8,
1471      sandbox_max_run_time_secs: int | None = None,
1472  ):
1473    if output_format not in {'rgb', 'yuv', 'gray'}:
1474      raise ValueError(
1475          f'Output format {output_format} is not rgb, yuv, or gray.'
1476      )
1477    self.path_or_url = path_or_url
1478    self.output_format = output_format
1479    self.stream_index = stream_index
1480    self.dtype = np.dtype(dtype)
1481    if self.dtype.type not in (np.uint8, np.uint16):
1482      raise ValueError(f'Type {dtype} is not np.uint8 or np.uint16.')
1483    self.sandbox_max_run_time_secs = sandbox_max_run_time_secs
1484    self._read_via_local_file: Any = None
1485    self._popen: subprocess.Popen[bytes] | None = None
1486    self._proc: subprocess.Popen[bytes] | None = None
1487
1488  def __enter__(self) -> 'VideoReader':
1489    try:
1490      self._read_via_local_file = _read_via_local_file(self.path_or_url)
1491      # pylint: disable-next=no-member
1492      tmp_name = self._read_via_local_file.__enter__()
1493
1494      self.metadata = _get_video_metadata(tmp_name)
1495      self.num_images, self.shape, self.fps, self.bps = self.metadata
1496      pix_fmt = self._get_pix_fmt(self.dtype, self.output_format)
1497      num_channels = {'rgb': 3, 'yuv': 3, 'gray': 1}[self.output_format]
1498      bytes_per_channel = self.dtype.itemsize
1499      self._num_bytes_per_image = (
1500          math.prod(self.shape) * num_channels * bytes_per_channel
1501      )
1502
1503      command = [
1504          '-v',
1505          'panic',
1506          '-nostdin',
1507          '-i',
1508          tmp_name,
1509          '-vcodec',
1510          'rawvideo',
1511          '-f',
1512          'image2pipe',
1513          '-map',
1514          f'0:v:{self.stream_index}',
1515          '-pix_fmt',
1516          pix_fmt,
1517          '-vsync',
1518          'vfr',
1519          '-',
1520      ]
1521      self._popen = _run_ffmpeg(
1522          command,
1523          stdout=subprocess.PIPE,
1524          stderr=subprocess.PIPE,
1525          allowed_input_files=[tmp_name],
1526          sandbox_max_run_time_secs=self.sandbox_max_run_time_secs,
1527      )
1528      self._proc = self._popen.__enter__()
1529    except Exception:
1530      self.__exit__(None, None, None)
1531      raise
1532    return self
1533
1534  def __exit__(self, *_: Any) -> None:
1535    self.close()
1536
1537  def read(self) -> _NDArray | None:
1538    """Reads a video image frame (or None if at end of file).
1539
1540    Returns:
1541      A numpy array in the format specified by `output_format`, i.e., a 3D
1542      array with 3 color channels, except for format 'gray' which is 2D.
1543
1544    Raises:
1545      RuntimeError: If there is an error reading from the output file.
1546    """
1547    assert self._proc, 'Error: reading from an already closed context.'
1548    stdout = self._proc.stdout
1549    assert stdout is not None
1550    data = stdout.read(self._num_bytes_per_image)
1551    if not data:  # Due to either end-of-file or subprocess error.
1552      self.close()  # Raises exception if subprocess had error.
1553      return None  # To indicate end-of-file.
1554    if len(data) != self._num_bytes_per_image:
1555      self._proc.wait()
1556      stderr = self._proc.stderr
1557      stderr_output = ''
1558      if stderr is not None:
1559        stderr_output = stderr.read().decode('utf-8', errors='replace').strip()
1560      raise RuntimeError(
1561          f'ffmpeg exited with code {self._proc.returncode}.\nIncomplete'
1562          f' frame read: expected {self._num_bytes_per_image} bytes, but got'
1563          f' {len(data)}.\nffmpeg stderr:\n{stderr_output}'
1564      )
1565    image = np.frombuffer(data, dtype=self.dtype)
1566    if self.output_format == 'rgb':
1567      image = image.reshape(*self.shape, 3)
1568    elif self.output_format == 'yuv':  # Convert from planar YUV to pixel YUV.
1569      image = np.moveaxis(image.reshape(3, *self.shape), 0, 2)
1570    elif self.output_format == 'gray':  # Generate 2D rather than 3D ndimage.
1571      image = image.reshape(*self.shape)
1572    else:
1573      raise AssertionError
1574    return image
1575
1576  def __iter__(self) -> Iterator[_NDArray]:
1577    while True:
1578      image = self.read()
1579      if image is None:
1580        return
1581      yield image
1582
1583  def close(self) -> None:
1584    """Terminates video reader.  (Called automatically at end of context.)"""
1585    if self._popen:
1586      self._popen.__exit__(None, None, None)
1587      self._popen = None
1588      self._proc = None
1589    if self._read_via_local_file:
1590      # pylint: disable-next=no-member
1591      self._read_via_local_file.__exit__(None, None, None)
1592      self._read_via_local_file = None

Context to read a compressed video as an iterable over its images.

>>> with VideoReader('/tmp/river.mp4') as reader:
...   print(f'Video has {reader.num_images} images with shape={reader.shape},'
...         f' at {reader.fps} frames/sec and {reader.bps} bits/sec.')
...   for image in reader:
...     print(image.shape)
>>> with VideoReader('/tmp/river.mp4') as reader:
...   video = np.array(tuple(reader))
>>> url = 'https://github.com/hhoppe/data/raw/main/video.mp4'
>>> with VideoReader(url) as reader:
...   show_video(reader)
Attributes:
  • path_or_url: Location of input video.
  • output_format: Format of output images (default 'rgb'). If 'rgb', each image has shape=(height, width, 3) with R, G, B values. If 'yuv', each image has shape=(height, width, 3) with Y, U, V values. If 'gray', each image has shape=(height, width).
  • dtype: Data type for output images. The default is np.uint8. Use of np.uint16 allows reading 10-bit or 12-bit data without precision loss.
  • metadata: Object storing the information retrieved from the video header. Its attributes are copied as attributes in this class.
  • num_images: Number of frames that is expected from the video stream. This is estimated from the framerate and the duration stored in the video header, so it might be inexact.
  • shape: The dimensions (height, width) of each video frame.
  • fps: The framerate in frames per second.
  • bps: The estimated bitrate of the video stream in bits per second, retrieved from the video header.
  • stream_index: The stream index to read from. The default is 0.
  • sandbox_max_run_time_secs: The maximum time in seconds to run the sandbox. If None, the default limit is 30 minutes. Unused in open source.
VideoReader( path_or_url: str | os.PathLike[str], *, stream_index: int = 0, output_format: str = 'rgb', dtype: DTypeLike = <class 'numpy.uint8'>, sandbox_max_run_time_secs: int | None = None)
1464  def __init__(
1465      self,
1466      path_or_url: _Path,
1467      *,
1468      stream_index: int = 0,
1469      output_format: str = 'rgb',
1470      dtype: _DTypeLike = np.uint8,
1471      sandbox_max_run_time_secs: int | None = None,
1472  ):
1473    if output_format not in {'rgb', 'yuv', 'gray'}:
1474      raise ValueError(
1475          f'Output format {output_format} is not rgb, yuv, or gray.'
1476      )
1477    self.path_or_url = path_or_url
1478    self.output_format = output_format
1479    self.stream_index = stream_index
1480    self.dtype = np.dtype(dtype)
1481    if self.dtype.type not in (np.uint8, np.uint16):
1482      raise ValueError(f'Type {dtype} is not np.uint8 or np.uint16.')
1483    self.sandbox_max_run_time_secs = sandbox_max_run_time_secs
1484    self._read_via_local_file: Any = None
1485    self._popen: subprocess.Popen[bytes] | None = None
1486    self._proc: subprocess.Popen[bytes] | None = None
path_or_url: str | os.PathLike[str]
output_format: str
dtype: ~_DType
metadata: VideoMetadata
num_images: int
shape: tuple[int, int]
fps: float
bps: int | None
stream_index: int
sandbox_max_run_time_secs
def read(self) -> np.ndarray | None:
1537  def read(self) -> _NDArray | None:
1538    """Reads a video image frame (or None if at end of file).
1539
1540    Returns:
1541      A numpy array in the format specified by `output_format`, i.e., a 3D
1542      array with 3 color channels, except for format 'gray' which is 2D.
1543
1544    Raises:
1545      RuntimeError: If there is an error reading from the output file.
1546    """
1547    assert self._proc, 'Error: reading from an already closed context.'
1548    stdout = self._proc.stdout
1549    assert stdout is not None
1550    data = stdout.read(self._num_bytes_per_image)
1551    if not data:  # Due to either end-of-file or subprocess error.
1552      self.close()  # Raises exception if subprocess had error.
1553      return None  # To indicate end-of-file.
1554    if len(data) != self._num_bytes_per_image:
1555      self._proc.wait()
1556      stderr = self._proc.stderr
1557      stderr_output = ''
1558      if stderr is not None:
1559        stderr_output = stderr.read().decode('utf-8', errors='replace').strip()
1560      raise RuntimeError(
1561          f'ffmpeg exited with code {self._proc.returncode}.\nIncomplete'
1562          f' frame read: expected {self._num_bytes_per_image} bytes, but got'
1563          f' {len(data)}.\nffmpeg stderr:\n{stderr_output}'
1564      )
1565    image = np.frombuffer(data, dtype=self.dtype)
1566    if self.output_format == 'rgb':
1567      image = image.reshape(*self.shape, 3)
1568    elif self.output_format == 'yuv':  # Convert from planar YUV to pixel YUV.
1569      image = np.moveaxis(image.reshape(3, *self.shape), 0, 2)
1570    elif self.output_format == 'gray':  # Generate 2D rather than 3D ndimage.
1571      image = image.reshape(*self.shape)
1572    else:
1573      raise AssertionError
1574    return image

Reads a video image frame (or None if at end of file).

Returns:

A numpy array in the format specified by output_format, i.e., a 3D array with 3 color channels, except for format 'gray' which is 2D.

Raises:
  • RuntimeError: If there is an error reading from the output file.
def close(self) -> None:
1583  def close(self) -> None:
1584    """Terminates video reader.  (Called automatically at end of context.)"""
1585    if self._popen:
1586      self._popen.__exit__(None, None, None)
1587      self._popen = None
1588      self._proc = None
1589    if self._read_via_local_file:
1590      # pylint: disable-next=no-member
1591      self._read_via_local_file.__exit__(None, None, None)
1592      self._read_via_local_file = None

Terminates video reader. (Called automatically at end of context.)

class VideoWriter(_VideoIO):
1595class VideoWriter(_VideoIO):
1596  """Context to write a compressed video.
1597
1598  >>> shape = 480, 640
1599  >>> with VideoWriter('/tmp/v.mp4', shape, fps=60) as writer:
1600  ...   for image in moving_circle(shape, num_images=60):
1601  ...     writer.add_image(image)
1602  >>> show_video(read_video('/tmp/v.mp4'))
1603
1604
1605  Bitrate control may be specified using at most one of: `bps`, `qp`, or `crf`.
1606  If none are specified, `qp` is set to a default value.
1607  See https://slhck.info/video/2017/03/01/rate-control.html
1608
1609  If codec is 'gif', the args `bps`, `qp`, `crf`, and `encoded_format` are
1610  ignored.
1611
1612  Attributes:
1613    path: Output video.  Its suffix (e.g. '.mp4') determines the video container
1614      format.  The suffix must be '.gif' if the codec is 'gif'.
1615    shape: 2D spatial dimensions (height, width) of video image frames.  The
1616      dimensions must be even if 'encoded_format' has subsampled chroma (e.g.,
1617      'yuv420p' or 'yuv420p10le').
1618    codec: Compression algorithm as defined by "ffmpeg -codecs" (e.g., 'h264',
1619      'hevc', 'vp9', or 'gif').
1620    metadata: Optional VideoMetadata object whose `fps` and `bps` attributes are
1621      used if not specified as explicit parameters.
1622    fps: Frames-per-second framerate (default is 60.0 except 25.0 for 'gif').
1623    bps: Requested average bits-per-second bitrate (default None).
1624    qp: Quantization parameter for video compression quality (default None).
1625    crf: Constant rate factor for video compression quality (default None).
1626    ffmpeg_args: Additional arguments for `ffmpeg` command, e.g. '-g 30' to
1627      introduce I-frames, or '-bf 0' to omit B-frames.
1628    input_format: Format of input images (default 'rgb').  If 'rgb', each image
1629      has shape=(height, width, 3) or (height, width).  If 'yuv', each image has
1630      shape=(height, width, 3) with Y, U, V values.  If 'gray', each image has
1631      shape=(height, width).
1632    dtype: Expected data type for input images (any float input images are
1633      converted to `dtype`).  The default is `np.uint8`.  Use of `np.uint16` is
1634      necessary when encoding >8 bits/channel.
1635    encoded_format: Pixel format as defined by `ffmpeg -pix_fmts`, e.g.,
1636      'yuv420p' (2x2-subsampled chroma), 'yuv444p' (full-res chroma),
1637      'yuv420p10le' (10-bit per channel), etc.  The default (None) selects
1638      'yuv420p' if all shape dimensions are even, else 'yuv444p'.
1639    sandbox_max_run_time_secs: The maximum time in seconds to run the sandbox.
1640      If None, the default limit is 30 minutes. Unused in open source.
1641    audio: Optional audio data as a NumPy array.  It should have shape (N,) for
1642      mono or (N, C) for multi-channel audio, where N is the number of samples
1643      and C is the number of channels.  The dtype must be one of `np.uint8`,
1644      `np.int16`, `np.int32`, `np.float32`, or `np.float64`, in native byte
1645      order.  Float samples are nominally in [-1.0, 1.0], and values outside
1646      this range typically clip; integer samples span the full range of their
1647      type, where `np.uint8` is unsigned with silence at 128.  The two streams
1648      are not truncated to a common length: if the audio is shorter than the
1649      video, the remainder is silent; if it is longer, the audio continues past
1650      the end of the video stream (most players keep showing the last frame).
1651    audio_sample_rate: Sample rate of the audio in Hz. Required if `audio` is
1652      provided.
1653    audio_codec: Audio compression algorithm as defined by "ffmpeg -codecs"
1654      (default 'aac').  It must be supported by the video container, e.g., 'aac'
1655      for MP4 ('h264' or 'hevc') or 'libopus' for WebM ('vp9').  Ignored if
1656      `audio` is None.
1657  """
1658
1659  def __init__(
1660      self,
1661      path: _Path,
1662      shape: tuple[int, int],
1663      *,
1664      codec: str = 'h264',
1665      metadata: VideoMetadata | None = None,
1666      fps: float | None = None,
1667      bps: int | None = None,
1668      qp: int | None = None,
1669      crf: float | None = None,
1670      ffmpeg_args: str | Sequence[str] = '',
1671      input_format: str = 'rgb',
1672      dtype: _DTypeLike = np.uint8,
1673      encoded_format: str | None = None,
1674      sandbox_max_run_time_secs: int | None = None,
1675      audio: _NDArray | None = None,
1676      audio_sample_rate: int | None = None,
1677      audio_codec: str = 'aac',
1678  ) -> None:
1679    _check_2d_shape(shape)
1680    if fps is None and metadata:
1681      fps = metadata.fps
1682    if fps is None:
1683      fps = 25.0 if codec == 'gif' else 60.0
1684    if fps <= 0.0:
1685      raise ValueError(f'Frame-per-second value {fps} is invalid.')
1686    if bps is None and metadata:
1687      bps = metadata.bps
1688    bps = int(bps) if bps is not None else None
1689    if bps is not None and bps <= 0:
1690      raise ValueError(f'Bitrate value {bps} is invalid.')
1691    if qp is not None and (not isinstance(qp, int) or qp < 0):
1692      raise ValueError(
1693          f'Quantization parameter {qp} cannot be negative. It must be a'
1694          ' non-negative integer.'
1695      )
1696    num_rate_specifications = sum(x is not None for x in (bps, qp, crf))
1697    if num_rate_specifications > 1:
1698      raise ValueError(
1699          f'Must specify at most one of bps, qp, or crf ({bps}, {qp}, {crf}).'
1700      )
1701    ffmpeg_args = (
1702        shlex.split(ffmpeg_args)
1703        if isinstance(ffmpeg_args, str)
1704        else list(ffmpeg_args)
1705    )
1706    if input_format not in {'rgb', 'yuv', 'gray'}:
1707      raise ValueError(f'Input format {input_format} is not rgb, yuv, or gray.')
1708    dtype = np.dtype(dtype)
1709    if dtype.type not in (np.uint8, np.uint16):
1710      raise ValueError(f'Type {dtype} is not np.uint8 or np.uint16.')
1711    if audio is not None:
1712      if audio_sample_rate is None:
1713        raise ValueError('The audio_sample_rate must be set if audio is set.')
1714      if audio.ndim not in (1, 2):
1715        raise ValueError(f'Audio shape {audio.shape} is not (N,) or (N, C).')
1716      if audio.dtype.type not in _FFMPEG_AUDIO_FORMAT_FROM_DTYPE:
1717        raise ValueError(f'Audio type {audio.dtype} is unsupported.')
1718      if not audio.dtype.isnative:
1719        raise ValueError(
1720            f'Audio type {audio.dtype} is not in native byte order.'
1721        )
1722      if codec == 'gif':
1723        raise ValueError('Audio is not supported with the gif codec.')
1724    self.path = pathlib.Path(path)
1725    self.shape = shape
1726    all_dimensions_are_even = all(dim % 2 == 0 for dim in shape)
1727    if encoded_format is None:
1728      encoded_format = 'yuv420p' if all_dimensions_are_even else 'yuv444p'
1729    if not all_dimensions_are_even and encoded_format.startswith(
1730        ('yuv42', 'yuvj42')
1731    ):
1732      raise ValueError(
1733          f'With encoded_format {encoded_format}, video dimensions must be'
1734          f' even, but shape is {shape}.'
1735      )
1736    self.fps = fps
1737    self.codec = codec
1738    self.bps = bps
1739    self.qp = qp
1740    self.crf = crf
1741    self.ffmpeg_args = ffmpeg_args
1742    self.input_format = input_format
1743    self.dtype = dtype
1744    self.encoded_format = encoded_format
1745    self.sandbox_max_run_time_secs = sandbox_max_run_time_secs
1746    self.audio = audio
1747    self.audio_sample_rate = audio_sample_rate
1748    self.audio_codec = audio_codec
1749    if num_rate_specifications == 0 and not ffmpeg_args:
1750      qp = 20 if math.prod(self.shape) <= 640 * 480 else 28
1751    self._bitrate_args = (
1752        (['-vb', f'{bps}'] if bps is not None else [])
1753        + (['-qp', f'{qp}'] if qp is not None else [])
1754        + (['-vb', '0', '-crf', f'{crf}'] if crf is not None else [])
1755    )
1756    if self.codec == 'gif':
1757      if self.path.suffix != '.gif':
1758        raise ValueError(f"File '{self.path}' does not have a .gif suffix.")
1759      self.encoded_format = 'pal8'
1760      self._bitrate_args = []
1761      video_filter = 'split[s0][s1];[s0]palettegen[p];[s1][p]paletteuse'
1762      # Less common (and likely less useful) is a per-frame color palette:
1763      # video_filter = ('split[s0][s1];[s0]palettegen=stats_mode=single[p];'
1764      #                 '[s1][p]paletteuse=new=1')
1765      self.ffmpeg_args = ['-vf', video_filter, '-f', 'gif'] + self.ffmpeg_args
1766    self._exit_stack: contextlib.ExitStack | None = None
1767    self._proc: subprocess.Popen[bytes] | None = None
1768
1769  def __enter__(self) -> 'VideoWriter':
1770    input_pix_fmt = self._get_pix_fmt(self.dtype, self.input_format)
1771    try:
1772      self._exit_stack = contextlib.ExitStack()
1773      tmp_name = self._exit_stack.enter_context(
1774          _write_via_local_file(self.path)
1775      )
1776
1777      # Writing to stdout using ('-f', 'mp4', '-') would require
1778      # ('-movflags', 'frag_keyframe+empty_moov') and the result is nonportable.
1779      height, width = self.shape
1780
1781      audio_input_args = []
1782      audio_output_args = ['-an']
1783      allowed_input_files = []
1784
1785      if self.audio is not None:
1786        audio_path = self._exit_stack.enter_context(
1787            _audio_via_local_file(self.audio)
1788        )
1789        channels = 1 if self.audio.ndim == 1 else self.audio.shape[1]
1790        audio_format = _FFMPEG_AUDIO_FORMAT_FROM_DTYPE[self.audio.dtype.type]
1791        if self.audio.dtype.itemsize > 1:
1792          audio_format += {'little': 'le', 'big': 'be'}[sys.byteorder]
1793        audio_input_args = [
1794            '-f',
1795            audio_format,
1796            '-ar',
1797            str(self.audio_sample_rate),
1798            '-ac',
1799            str(channels),
1800            '-i',
1801            audio_path,
1802        ]
1803        audio_output_args = ['-c:a', self.audio_codec]
1804        allowed_input_files.append(audio_path)
1805
1806      command = (
1807          [
1808              '-v',
1809              'error',
1810              '-f',
1811              'rawvideo',
1812              '-vcodec',
1813              'rawvideo',
1814              '-pix_fmt',
1815              input_pix_fmt,
1816              '-s',
1817              f'{width}x{height}',
1818              '-r',
1819              f'{self.fps}',
1820              '-i',
1821              '-',
1822          ]
1823          + audio_input_args
1824          + audio_output_args
1825          + [
1826              '-vcodec',
1827              self.codec,
1828              '-pix_fmt',
1829              self.encoded_format,
1830          ]
1831          + self._bitrate_args
1832          + self.ffmpeg_args
1833          + ['-y', tmp_name]
1834      )
1835      self._proc = self._exit_stack.enter_context(
1836          _run_ffmpeg(
1837              command,
1838              stdin=subprocess.PIPE,
1839              stderr=subprocess.PIPE,
1840              # `_run_ffmpeg` omits the sandbox flag only for None, so an empty
1841              # list would pass an empty '--sandbox_read_access_files'.
1842              allowed_input_files=allowed_input_files or None,
1843              allowed_output_files=[tmp_name],
1844              sandbox_max_run_time_secs=self.sandbox_max_run_time_secs,
1845          )
1846      )
1847    except Exception:
1848      self.__exit__(None, None, None)
1849      raise
1850    return self
1851
1852  def __exit__(self, *_: Any) -> None:
1853    self.close()
1854
1855  def add_image(self, image: _NDArray) -> None:
1856    """Writes a video frame.
1857
1858    Args:
1859      image: Array whose dtype and first two dimensions must match the `dtype`
1860        and `shape` specified in `VideoWriter` initialization.  If
1861        `input_format` is 'gray', the image must be 2D.  For the 'rgb'
1862        input_format, the image may be either 2D (interpreted as grayscale) or
1863        3D with three (R, G, B) channels.  For the 'yuv' input_format, the image
1864        must be 3D with three (Y, U, V) channels.
1865
1866    Raises:
1867      RuntimeError: If there is an error writing to the output file.
1868    """
1869    assert self._proc, 'Error: writing to an already closed context.'
1870    if issubclass(image.dtype.type, (np.floating, np.bool_)):
1871      image = to_type(image, self.dtype)
1872    if image.dtype != self.dtype:
1873      raise ValueError(f'Image type {image.dtype} != {self.dtype}.')
1874    if self.input_format == 'gray':
1875      if image.ndim != 2:
1876        raise ValueError(f'Image dimensions {image.shape} are not 2D.')
1877    else:
1878      if image.ndim == 2 and self.input_format == 'rgb':
1879        image = np.dstack((image, image, image))
1880      if not (image.ndim == 3 and image.shape[2] == 3):
1881        raise ValueError(f'Image dimensions {image.shape} are invalid.')
1882    if image.shape[:2] != self.shape:
1883      raise ValueError(
1884          f'Image dimensions {image.shape[:2]} do not match'
1885          f' those of the initialized video {self.shape}.'
1886      )
1887    if self.input_format == 'yuv':  # Convert from per-pixel YUV to planar YUV.
1888      image = np.moveaxis(image, 2, 0)
1889    data = image.tobytes()
1890    stdin = self._proc.stdin
1891    assert stdin is not None
1892    if stdin.write(data) != len(data):
1893      self._proc.wait()
1894      stderr = self._proc.stderr
1895      assert stderr is not None
1896      s = stderr.read().decode('utf-8')
1897      raise RuntimeError(f"Error writing '{self.path}': {s}")
1898
1899  def close(self) -> None:
1900    """Finishes writing the video.  (Called automatically at end of context.)"""
1901    if self._exit_stack is None:
1902      return
1903    # Unwinding the stack terminates the `ffmpeg` process, removes the
1904    # temporary audio file, and copies the encoded video to a remote `path`.
1905    # Raising within the `with` propagates the error into those contexts, so
1906    # an incomplete video is discarded rather than copied.
1907    with self._exit_stack:
1908      self._exit_stack = None
1909      if self._proc is not None:
1910        proc, self._proc = self._proc, None
1911        stdin = proc.stdin
1912        assert stdin is not None
1913        stdin.close()
1914        if proc.wait():
1915          stderr = proc.stderr
1916          assert stderr is not None
1917          s = stderr.read().decode('utf-8')
1918          raise RuntimeError(f"Error writing '{self.path}': {s}")

Context to write a compressed video.

>>> shape = 480, 640
>>> with VideoWriter('/tmp/v.mp4', shape, fps=60) as writer:
...   for image in moving_circle(shape, num_images=60):
...     writer.add_image(image)
>>> show_video(read_video('/tmp/v.mp4'))

Bitrate control may be specified using at most one of: bps, qp, or crf. If none are specified, qp is set to a default value. See https://slhck.info/video/2017/03/01/rate-control.html

If codec is 'gif', the args bps, qp, crf, and encoded_format are ignored.

Attributes:
  • path: Output video. Its suffix (e.g. '.mp4') determines the video container format. The suffix must be '.gif' if the codec is 'gif'.
  • shape: 2D spatial dimensions (height, width) of video image frames. The dimensions must be even if 'encoded_format' has subsampled chroma (e.g., 'yuv420p' or 'yuv420p10le').
  • codec: Compression algorithm as defined by "ffmpeg -codecs" (e.g., 'h264', 'hevc', 'vp9', or 'gif').
  • metadata: Optional VideoMetadata object whose fps and bps attributes are used if not specified as explicit parameters.
  • fps: Frames-per-second framerate (default is 60.0 except 25.0 for 'gif').
  • bps: Requested average bits-per-second bitrate (default None).
  • qp: Quantization parameter for video compression quality (default None).
  • crf: Constant rate factor for video compression quality (default None).
  • ffmpeg_args: Additional arguments for ffmpeg command, e.g. '-g 30' to introduce I-frames, or '-bf 0' to omit B-frames.
  • input_format: Format of input images (default 'rgb'). If 'rgb', each image has shape=(height, width, 3) or (height, width). If 'yuv', each image has shape=(height, width, 3) with Y, U, V values. If 'gray', each image has shape=(height, width).
  • dtype: Expected data type for input images (any float input images are converted to dtype). The default is np.uint8. Use of np.uint16 is necessary when encoding >8 bits/channel.
  • encoded_format: Pixel format as defined by ffmpeg -pix_fmts, e.g., 'yuv420p' (2x2-subsampled chroma), 'yuv444p' (full-res chroma), 'yuv420p10le' (10-bit per channel), etc. The default (None) selects 'yuv420p' if all shape dimensions are even, else 'yuv444p'.
  • sandbox_max_run_time_secs: The maximum time in seconds to run the sandbox. If None, the default limit is 30 minutes. Unused in open source.
  • audio: Optional audio data as a NumPy array. It should have shape (N,) for mono or (N, C) for multi-channel audio, where N is the number of samples and C is the number of channels. The dtype must be one of np.uint8, np.int16, np.int32, np.float32, or np.float64, in native byte order. Float samples are nominally in [-1.0, 1.0], and values outside this range typically clip; integer samples span the full range of their type, where np.uint8 is unsigned with silence at 128. The two streams are not truncated to a common length: if the audio is shorter than the video, the remainder is silent; if it is longer, the audio continues past the end of the video stream (most players keep showing the last frame).
  • audio_sample_rate: Sample rate of the audio in Hz. Required if audio is provided.
  • audio_codec: Audio compression algorithm as defined by "ffmpeg -codecs" (default 'aac'). It must be supported by the video container, e.g., 'aac' for MP4 ('h264' or 'hevc') or 'libopus' for WebM ('vp9'). Ignored if audio is None.
VideoWriter( path: str | os.PathLike[str], shape: tuple[int, int], *, codec: str = 'h264', metadata: VideoMetadata | None = None, fps: float | None = None, bps: int | None = None, qp: int | None = None, crf: float | None = None, ffmpeg_args: str | Sequence[str] = '', input_format: str = 'rgb', dtype: DTypeLike = <class 'numpy.uint8'>, encoded_format: str | None = None, sandbox_max_run_time_secs: int | None = None, audio: np.ndarray | None = None, audio_sample_rate: int | None = None, audio_codec: str = 'aac')
1659  def __init__(
1660      self,
1661      path: _Path,
1662      shape: tuple[int, int],
1663      *,
1664      codec: str = 'h264',
1665      metadata: VideoMetadata | None = None,
1666      fps: float | None = None,
1667      bps: int | None = None,
1668      qp: int | None = None,
1669      crf: float | None = None,
1670      ffmpeg_args: str | Sequence[str] = '',
1671      input_format: str = 'rgb',
1672      dtype: _DTypeLike = np.uint8,
1673      encoded_format: str | None = None,
1674      sandbox_max_run_time_secs: int | None = None,
1675      audio: _NDArray | None = None,
1676      audio_sample_rate: int | None = None,
1677      audio_codec: str = 'aac',
1678  ) -> None:
1679    _check_2d_shape(shape)
1680    if fps is None and metadata:
1681      fps = metadata.fps
1682    if fps is None:
1683      fps = 25.0 if codec == 'gif' else 60.0
1684    if fps <= 0.0:
1685      raise ValueError(f'Frame-per-second value {fps} is invalid.')
1686    if bps is None and metadata:
1687      bps = metadata.bps
1688    bps = int(bps) if bps is not None else None
1689    if bps is not None and bps <= 0:
1690      raise ValueError(f'Bitrate value {bps} is invalid.')
1691    if qp is not None and (not isinstance(qp, int) or qp < 0):
1692      raise ValueError(
1693          f'Quantization parameter {qp} cannot be negative. It must be a'
1694          ' non-negative integer.'
1695      )
1696    num_rate_specifications = sum(x is not None for x in (bps, qp, crf))
1697    if num_rate_specifications > 1:
1698      raise ValueError(
1699          f'Must specify at most one of bps, qp, or crf ({bps}, {qp}, {crf}).'
1700      )
1701    ffmpeg_args = (
1702        shlex.split(ffmpeg_args)
1703        if isinstance(ffmpeg_args, str)
1704        else list(ffmpeg_args)
1705    )
1706    if input_format not in {'rgb', 'yuv', 'gray'}:
1707      raise ValueError(f'Input format {input_format} is not rgb, yuv, or gray.')
1708    dtype = np.dtype(dtype)
1709    if dtype.type not in (np.uint8, np.uint16):
1710      raise ValueError(f'Type {dtype} is not np.uint8 or np.uint16.')
1711    if audio is not None:
1712      if audio_sample_rate is None:
1713        raise ValueError('The audio_sample_rate must be set if audio is set.')
1714      if audio.ndim not in (1, 2):
1715        raise ValueError(f'Audio shape {audio.shape} is not (N,) or (N, C).')
1716      if audio.dtype.type not in _FFMPEG_AUDIO_FORMAT_FROM_DTYPE:
1717        raise ValueError(f'Audio type {audio.dtype} is unsupported.')
1718      if not audio.dtype.isnative:
1719        raise ValueError(
1720            f'Audio type {audio.dtype} is not in native byte order.'
1721        )
1722      if codec == 'gif':
1723        raise ValueError('Audio is not supported with the gif codec.')
1724    self.path = pathlib.Path(path)
1725    self.shape = shape
1726    all_dimensions_are_even = all(dim % 2 == 0 for dim in shape)
1727    if encoded_format is None:
1728      encoded_format = 'yuv420p' if all_dimensions_are_even else 'yuv444p'
1729    if not all_dimensions_are_even and encoded_format.startswith(
1730        ('yuv42', 'yuvj42')
1731    ):
1732      raise ValueError(
1733          f'With encoded_format {encoded_format}, video dimensions must be'
1734          f' even, but shape is {shape}.'
1735      )
1736    self.fps = fps
1737    self.codec = codec
1738    self.bps = bps
1739    self.qp = qp
1740    self.crf = crf
1741    self.ffmpeg_args = ffmpeg_args
1742    self.input_format = input_format
1743    self.dtype = dtype
1744    self.encoded_format = encoded_format
1745    self.sandbox_max_run_time_secs = sandbox_max_run_time_secs
1746    self.audio = audio
1747    self.audio_sample_rate = audio_sample_rate
1748    self.audio_codec = audio_codec
1749    if num_rate_specifications == 0 and not ffmpeg_args:
1750      qp = 20 if math.prod(self.shape) <= 640 * 480 else 28
1751    self._bitrate_args = (
1752        (['-vb', f'{bps}'] if bps is not None else [])
1753        + (['-qp', f'{qp}'] if qp is not None else [])
1754        + (['-vb', '0', '-crf', f'{crf}'] if crf is not None else [])
1755    )
1756    if self.codec == 'gif':
1757      if self.path.suffix != '.gif':
1758        raise ValueError(f"File '{self.path}' does not have a .gif suffix.")
1759      self.encoded_format = 'pal8'
1760      self._bitrate_args = []
1761      video_filter = 'split[s0][s1];[s0]palettegen[p];[s1][p]paletteuse'
1762      # Less common (and likely less useful) is a per-frame color palette:
1763      # video_filter = ('split[s0][s1];[s0]palettegen=stats_mode=single[p];'
1764      #                 '[s1][p]paletteuse=new=1')
1765      self.ffmpeg_args = ['-vf', video_filter, '-f', 'gif'] + self.ffmpeg_args
1766    self._exit_stack: contextlib.ExitStack | None = None
1767    self._proc: subprocess.Popen[bytes] | None = None
path
shape
fps
codec
bps
qp
crf
ffmpeg_args
input_format
dtype
encoded_format
sandbox_max_run_time_secs
audio
audio_sample_rate
audio_codec
def add_image(self, image: np.ndarray) -> None:
1855  def add_image(self, image: _NDArray) -> None:
1856    """Writes a video frame.
1857
1858    Args:
1859      image: Array whose dtype and first two dimensions must match the `dtype`
1860        and `shape` specified in `VideoWriter` initialization.  If
1861        `input_format` is 'gray', the image must be 2D.  For the 'rgb'
1862        input_format, the image may be either 2D (interpreted as grayscale) or
1863        3D with three (R, G, B) channels.  For the 'yuv' input_format, the image
1864        must be 3D with three (Y, U, V) channels.
1865
1866    Raises:
1867      RuntimeError: If there is an error writing to the output file.
1868    """
1869    assert self._proc, 'Error: writing to an already closed context.'
1870    if issubclass(image.dtype.type, (np.floating, np.bool_)):
1871      image = to_type(image, self.dtype)
1872    if image.dtype != self.dtype:
1873      raise ValueError(f'Image type {image.dtype} != {self.dtype}.')
1874    if self.input_format == 'gray':
1875      if image.ndim != 2:
1876        raise ValueError(f'Image dimensions {image.shape} are not 2D.')
1877    else:
1878      if image.ndim == 2 and self.input_format == 'rgb':
1879        image = np.dstack((image, image, image))
1880      if not (image.ndim == 3 and image.shape[2] == 3):
1881        raise ValueError(f'Image dimensions {image.shape} are invalid.')
1882    if image.shape[:2] != self.shape:
1883      raise ValueError(
1884          f'Image dimensions {image.shape[:2]} do not match'
1885          f' those of the initialized video {self.shape}.'
1886      )
1887    if self.input_format == 'yuv':  # Convert from per-pixel YUV to planar YUV.
1888      image = np.moveaxis(image, 2, 0)
1889    data = image.tobytes()
1890    stdin = self._proc.stdin
1891    assert stdin is not None
1892    if stdin.write(data) != len(data):
1893      self._proc.wait()
1894      stderr = self._proc.stderr
1895      assert stderr is not None
1896      s = stderr.read().decode('utf-8')
1897      raise RuntimeError(f"Error writing '{self.path}': {s}")

Writes a video frame.

Arguments:
  • image: Array whose dtype and first two dimensions must match the dtype and shape specified in VideoWriter initialization. If input_format is 'gray', the image must be 2D. For the 'rgb' input_format, the image may be either 2D (interpreted as grayscale) or 3D with three (R, G, B) channels. For the 'yuv' input_format, the image must be 3D with three (Y, U, V) channels.
Raises:
  • RuntimeError: If there is an error writing to the output file.
def close(self) -> None:
1899  def close(self) -> None:
1900    """Finishes writing the video.  (Called automatically at end of context.)"""
1901    if self._exit_stack is None:
1902      return
1903    # Unwinding the stack terminates the `ffmpeg` process, removes the
1904    # temporary audio file, and copies the encoded video to a remote `path`.
1905    # Raising within the `with` propagates the error into those contexts, so
1906    # an incomplete video is discarded rather than copied.
1907    with self._exit_stack:
1908      self._exit_stack = None
1909      if self._proc is not None:
1910        proc, self._proc = self._proc, None
1911        stdin = proc.stdin
1912        assert stdin is not None
1913        stdin.close()
1914        if proc.wait():
1915          stderr = proc.stderr
1916          assert stderr is not None
1917          s = stderr.read().decode('utf-8')
1918          raise RuntimeError(f"Error writing '{self.path}': {s}")

Finishes writing the video. (Called automatically at end of context.)

class VideoMetadata(typing.NamedTuple):
1302class VideoMetadata(NamedTuple):
1303  """Represents the data stored in a video container header.
1304
1305  Attributes:
1306    num_images: Number of frames that is expected from the video stream.  This
1307      is estimated from the framerate and the duration stored in the video
1308      header, so it might be inexact.  We set the value to -1 if number of
1309      frames is not found in the header.
1310    shape: The dimensions (height, width) of each video frame.
1311    fps: The framerate in frames per second.
1312    bps: The estimated bitrate of the video stream in bits per second, retrieved
1313      from the video header.
1314  """
1315
1316  num_images: int
1317  shape: tuple[int, int]
1318  fps: float
1319  bps: int | None

Represents the data stored in a video container header.

Attributes:
  • num_images: Number of frames that is expected from the video stream. This is estimated from the framerate and the duration stored in the video header, so it might be inexact. We set the value to -1 if number of frames is not found in the header.
  • shape: The dimensions (height, width) of each video frame.
  • fps: The framerate in frames per second.
  • bps: The estimated bitrate of the video stream in bits per second, retrieved from the video header.
def compress_image(image: ArrayLike, *, fmt: str = 'png', **kwargs: Any) -> bytes:
887def compress_image(
888    image: _ArrayLike, *, fmt: str = 'png', **kwargs: Any
889) -> bytes:
890  """Returns a buffer containing a compressed image.
891
892  Args:
893    image: Array in a format supported by `PIL`, e.g. np.uint8 or np.uint16.
894    fmt: Desired compression encoding, e.g. 'png'.
895    **kwargs: Options for `PIL.save()`, e.g. `optimize=True` for greater
896      compression.
897  """
898  image = _as_valid_media_array(image)
899  with io.BytesIO() as output:
900    _pil_image(image).save(output, format=fmt, **kwargs)
901    return output.getvalue()

Returns a buffer containing a compressed image.

Arguments:
  • image: Array in a format supported by PIL, e.g. np.uint8 or np.uint16.
  • fmt: Desired compression encoding, e.g. 'png'.
  • **kwargs: Options for PIL.save(), e.g. optimize=True for greater compression.
def decompress_image( data: bytes, dtype: DTypeLike = None, apply_exif_transpose: bool = True) -> np.ndarray:
904def decompress_image(
905    data: bytes, dtype: _DTypeLike = None, apply_exif_transpose: bool = True  # pyrefly: ignore[bad-function-definition]
906) -> _NDArray:
907  """Returns an image from a compressed data buffer.
908
909  Decoding is performed using `PIL`, which supports `uint8` images with 1, 3,
910  or 4 channels and `uint16` images with a single channel.
911
912  Args:
913    data: Buffer containing compressed image.
914    dtype: Data type of the returned array.  If None, `np.uint8` or `np.uint16`
915      is inferred automatically.
916    apply_exif_transpose: If True, rotate image according to EXIF orientation.
917  """
918  pil_image: PIL.Image.Image = PIL.Image.open(io.BytesIO(data))
919  if apply_exif_transpose:
920    tmp_image = PIL.ImageOps.exif_transpose(pil_image)  # Future: in_place=True.
921    assert tmp_image
922    pil_image = tmp_image
923  if dtype is None:
924    dtype = np.uint16 if pil_image.mode.startswith('I') else np.uint8
925  return np.array(pil_image, dtype=dtype)

Returns an image from a compressed data buffer.

Decoding is performed using PIL, which supports uint8 images with 1, 3, or 4 channels and uint16 images with a single channel.

Arguments:
  • data: Buffer containing compressed image.
  • dtype: Data type of the returned array. If None, np.uint8 or np.uint16 is inferred automatically.
  • apply_exif_transpose: If True, rotate image according to EXIF orientation.
def compress_video( images: Iterable[np.ndarray], *, codec: str = 'h264', **kwargs: Any) -> bytes:
1992def compress_video(
1993    images: Iterable[_NDArray], *, codec: str = 'h264', **kwargs: Any
1994) -> bytes:
1995  """Returns a buffer containing a compressed video.
1996
1997  The video container is 'gif' for 'gif' codec, 'webm' for 'vp9' codec,
1998  and mp4 otherwise.
1999
2000  >>> video = read_video('/tmp/river.mp4')
2001  >>> data = compress_video(video, bps=10_000_000)
2002  >>> print(len(data))
2003
2004  >>> data = compress_video(moving_circle((100, 100), num_images=10), fps=10)
2005
2006  Args:
2007    images: Iterable over video frames.
2008    codec: Compression algorithm as defined by `ffmpeg -codecs` (e.g., 'h264',
2009      'hevc', 'vp9', or 'gif').
2010    **kwargs: Additional parameters for `VideoWriter`.
2011
2012  Returns:
2013    A bytes buffer containing the compressed video.
2014  """
2015  suffix = _filename_suffix_from_codec(codec)
2016  with tempfile.TemporaryDirectory() as directory_name:
2017    tmp_path = pathlib.Path(directory_name) / f'file{suffix}'
2018    write_video(tmp_path, images, codec=codec, **kwargs)
2019    return tmp_path.read_bytes()

Returns a buffer containing a compressed video.

The video container is 'gif' for 'gif' codec, 'webm' for 'vp9' codec, and mp4 otherwise.

>>> video = read_video('/tmp/river.mp4')
>>> data = compress_video(video, bps=10_000_000)
>>> print(len(data))
>>> data = compress_video(moving_circle((100, 100), num_images=10), fps=10)
Arguments:
  • images: Iterable over video frames.
  • codec: Compression algorithm as defined by ffmpeg -codecs (e.g., 'h264', 'hevc', 'vp9', or 'gif').
  • **kwargs: Additional parameters for VideoWriter.
Returns:

A bytes buffer containing the compressed video.

def decompress_video(data: bytes, **kwargs: Any) -> np.ndarray:
2022def decompress_video(data: bytes, **kwargs: Any) -> _NDArray:
2023  """Returns video images from an MP4-compressed data buffer."""
2024  with tempfile.TemporaryDirectory() as directory_name:
2025    tmp_path = pathlib.Path(directory_name) / 'file.mp4'
2026    tmp_path.write_bytes(data)
2027    return read_video(tmp_path, **kwargs)

Returns video images from an MP4-compressed data buffer.

def html_from_compressed_image( data: bytes, width: int, height: int, *, title: str | None = None, border: bool | str = False, pixelated: bool = True, fmt: str = 'png') -> str:
928def html_from_compressed_image(
929    data: bytes,
930    width: int,
931    height: int,
932    *,
933    title: str | None = None,
934    border: bool | str = False,
935    pixelated: bool = True,
936    fmt: str = 'png',
937) -> str:
938  """Returns an HTML string with an image tag containing encoded data.
939
940  Args:
941    data: Compressed image bytes.
942    width: Width of HTML image in pixels.
943    height: Height of HTML image in pixels.
944    title: Optional text shown centered above image.
945    border: If `bool`, whether to place a black boundary around the image, or if
946      `str`, the boundary CSS style.
947    pixelated: If True, sets the CSS style to 'image-rendering: pixelated;'.
948    fmt: Compression encoding.
949  """
950  b64 = base64.b64encode(data).decode('utf-8')
951  if isinstance(border, str):
952    border = f'{border}; '
953  elif border:
954    border = 'border:1px solid black; '
955  else:
956    border = ''
957  s_pixelated = 'pixelated' if pixelated else 'auto'
958  s = (
959      f'<img width="{width}" height="{height}"'
960      f' style="{border}image-rendering:{s_pixelated}; object-fit:cover;"'
961      f' src="data:image/{fmt};base64,{b64}"/>'
962  )
963  if title is not None:
964    s = f"""<div style="display:flex; align-items:left;">
965      <div style="display:flex; flex-direction:column; align-items:center;">
966      <div>{title}</div><div>{s}</div></div></div>"""
967  return s

Returns an HTML string with an image tag containing encoded data.

Arguments:
  • data: Compressed image bytes.
  • width: Width of HTML image in pixels.
  • height: Height of HTML image in pixels.
  • title: Optional text shown centered above image.
  • border: If bool, whether to place a black boundary around the image, or if str, the boundary CSS style.
  • pixelated: If True, sets the CSS style to 'image-rendering: pixelated;'.
  • fmt: Compression encoding.
def html_from_compressed_video( data: bytes, width: int, height: int, *, title: str | None = None, border: bool | str = False, loop: bool = True, autoplay: bool = True) -> str:
2030def html_from_compressed_video(
2031    data: bytes,
2032    width: int,
2033    height: int,
2034    *,
2035    title: str | None = None,
2036    border: bool | str = False,
2037    loop: bool = True,
2038    autoplay: bool = True,
2039) -> str:
2040  """Returns an HTML string with a video tag containing H264-encoded data.
2041
2042  Args:
2043    data: MP4-compressed video bytes.
2044    width: Width of HTML video in pixels.
2045    height: Height of HTML video in pixels.
2046    title: Optional text shown centered above the video.
2047    border: If `bool`, whether to place a black boundary around the image, or if
2048      `str`, the boundary CSS style.
2049    loop: If True, the playback repeats forever.
2050    autoplay: If True, video playback starts without having to click.
2051  """
2052  b64 = base64.b64encode(data).decode('utf-8')
2053  if isinstance(border, str):
2054    border = f'{border}; '
2055  elif border:
2056    border = 'border:1px solid black; '
2057  else:
2058    border = ''
2059  options = (
2060      f'controls width="{width}" height="{height}"'
2061      f' style="{border}object-fit:cover;"'
2062      f'{" loop" if loop else ""}'
2063      f'{" autoplay muted" if autoplay else ""}'
2064  )
2065  s = f"""<video {options}>
2066      <source src="data:video/mp4;base64,{b64}" type="video/mp4"/>
2067      This browser does not support the video tag.
2068      </video>"""
2069  if title is not None:
2070    s = f"""<div style="display:flex; align-items:left;">
2071      <div style="display:flex; flex-direction:column; align-items:center;">
2072      <div>{title}</div><div>{s}</div></div></div>"""
2073  return s

Returns an HTML string with a video tag containing H264-encoded data.

Arguments:
  • data: MP4-compressed video bytes.
  • width: Width of HTML video in pixels.
  • height: Height of HTML video in pixels.
  • title: Optional text shown centered above the video.
  • border: If bool, whether to place a black boundary around the image, or if str, the boundary CSS style.
  • loop: If True, the playback repeats forever.
  • autoplay: If True, video playback starts without having to click.
def resize_image(image: ArrayLike, shape: tuple[int, int]) -> np.ndarray:
626def resize_image(image: _ArrayLike, shape: tuple[int, int]) -> _NDArray:
627  """Resizes image to specified spatial dimensions using a Lanczos filter.
628
629  Args:
630    image: Array-like 2D or 3D object, where dtype is uint or floating-point.
631    shape: 2D spatial dimensions (height, width) of output image.
632
633  Returns:
634    A resampled image whose spatial dimensions match `shape`.
635  """
636  image = _as_valid_media_array(image)
637  if image.ndim not in (2, 3):
638    raise ValueError(f'Image shape {image.shape} is neither 2D nor 3D.')
639  _check_2d_shape(shape)
640
641  # A PIL image can be multichannel only if it has 3 or 4 uint8 channels,
642  # and it can be resized only if it is uint8 or float32.
643  supported_single_channel = (
644      np.issubdtype(image.dtype, np.floating) or image.dtype == np.uint8
645  ) and image.ndim == 2
646  supported_multichannel = (
647      image.dtype == np.uint8 and image.ndim == 3 and image.shape[2] in (3, 4)
648  )
649  if supported_single_channel or supported_multichannel:
650    return np.array(
651        _pil_image(image).resize(
652            shape[::-1], resample=PIL.Image.Resampling.LANCZOS
653        ),
654        dtype=image.dtype,
655    )
656  if image.ndim == 2:
657    # We convert to floating-point for resizing and convert back.
658    return to_type(resize_image(to_float01(image), shape), image.dtype)
659  # We resize each image channel individually.
660  return np.dstack(
661      [resize_image(channel, shape) for channel in np.moveaxis(image, -1, 0)]
662  )

Resizes image to specified spatial dimensions using a Lanczos filter.

Arguments:
  • image: Array-like 2D or 3D object, where dtype is uint or floating-point.
  • shape: 2D spatial dimensions (height, width) of output image.
Returns:

A resampled image whose spatial dimensions match shape.

def resize_video(video: Iterable[np.ndarray], shape: tuple[int, int]) -> np.ndarray:
668def resize_video(video: Iterable[_NDArray], shape: tuple[int, int]) -> _NDArray:
669  """Resizes `video` to specified spatial dimensions using a Lanczos filter.
670
671  Args:
672    video: Iterable of images.
673    shape: 2D spatial dimensions (height, width) of output video.
674
675  Returns:
676    A resampled video whose spatial dimensions match `shape`.
677  """
678  _check_2d_shape(shape)
679  return np.array([resize_image(image, shape) for image in video])

Resizes video to specified spatial dimensions using a Lanczos filter.

Arguments:
  • video: Iterable of images.
  • shape: 2D spatial dimensions (height, width) of output video.
Returns:

A resampled video whose spatial dimensions match shape.

def to_rgb( array: ArrayLike, *, vmin: float | None = None, vmax: float | None = None, cmap: str | Callable[[ArrayLike], np.ndarray] = 'gray') -> np.ndarray:
843def to_rgb(
844    array: _ArrayLike,
845    *,
846    vmin: float | None = None,
847    vmax: float | None = None,
848    cmap: str | Callable[[_ArrayLike], _NDArray] = 'gray',
849) -> _NDArray:
850  """Maps scalar values to RGB using value bounds and a color map.
851
852  Args:
853    array: Scalar values, with arbitrary shape.
854    vmin: Explicit min value for remapping; if None, it is obtained as the
855      minimum finite value of `array`.
856    vmax: Explicit max value for remapping; if None, it is obtained as the
857      maximum finite value of `array`.
858    cmap: A `pyplot` color map or callable, to map from 1D value to 3D or 4D
859      color.
860
861  Returns:
862    A new array in which each element is affinely mapped from [vmin, vmax]
863    to [0.0, 1.0] and then color-mapped.
864  """
865  a = _as_valid_media_array(array)
866  del array
867  # For future numpy version 1.7.0:
868  # vmin = np.amin(a, where=np.isfinite(a)) if vmin is None else vmin
869  # vmax = np.amax(a, where=np.isfinite(a)) if vmax is None else vmax
870  vmin = np.amin(np.where(np.isfinite(a), a, np.inf)) if vmin is None else vmin
871  vmax = np.amax(np.where(np.isfinite(a), a, -np.inf)) if vmax is None else vmax
872  a = (a.astype('float') - vmin) / (vmax - vmin + np.finfo(float).eps)
873  if isinstance(cmap, str):
874    if hasattr(matplotlib, 'colormaps'):
875      rgb_from_scalar: Any = matplotlib.colormaps[cmap]  # Newer version.
876    else:
877      rgb_from_scalar = matplotlib.pyplot.cm.get_cmap(cmap)  # pylint: disable=no-member
878  else:
879    rgb_from_scalar = cmap
880  a = cast(_NDArray, rgb_from_scalar(a))
881  # If there is a fully opaque alpha channel, remove it.
882  if a.shape[-1] == 4 and np.all(to_float01(a[..., 3])) == 1.0:
883    a = a[..., :3]
884  return a

Maps scalar values to RGB using value bounds and a color map.

Arguments:
  • array: Scalar values, with arbitrary shape.
  • vmin: Explicit min value for remapping; if None, it is obtained as the minimum finite value of array.
  • vmax: Explicit max value for remapping; if None, it is obtained as the maximum finite value of array.
  • cmap: A pyplot color map or callable, to map from 1D value to 3D or 4D color.
Returns:

A new array in which each element is affinely mapped from [vmin, vmax] to [0.0, 1.0] and then color-mapped.

def to_type(array: ArrayLike, dtype: DTypeLike) -> np.ndarray:
387def to_type(array: _ArrayLike, dtype: _DTypeLike) -> _NDArray:
388  """Returns media array converted to specified type.
389
390  A "media array" is one in which the dtype is either a floating-point type
391  (np.float32 or np.float64) or an unsigned integer type.  The array values are
392  assumed to lie in the range [0.0, 1.0] for floating-point values, and in the
393  full range for unsigned integers, e.g. [0, 255] for np.uint8.
394
395  Conversion between integers and floats maps uint(0) to 0.0 and uint(MAX) to
396  1.0.  The input array may also be of type bool, whereby True maps to
397  uint(MAX) or 1.0.  The values are scaled and clamped as appropriate during
398  type conversions.
399
400  Args:
401    array: Input array-like object (floating-point, unsigned int, or bool).
402    dtype: Desired output type (floating-point or unsigned int).
403
404  Returns:
405    Array `a` if it is already of the specified dtype, else a converted array.
406  """
407  a = np.asarray(array)
408  dtype = np.dtype(dtype)
409  del array
410  if a.dtype != bool:
411    _as_valid_media_type(a.dtype)  # Verify that 'a' has a valid dtype.
412  if a.dtype == bool:
413    result = a.astype(dtype)
414    if np.issubdtype(dtype, np.unsignedinteger):
415      result = result * dtype.type(np.iinfo(dtype).max)  # pyrefly: ignore[no-matching-overload]
416  elif a.dtype == dtype:
417    result = a
418  elif np.issubdtype(dtype, np.unsignedinteger):
419    if np.issubdtype(a.dtype, np.unsignedinteger):
420      src_max: float = np.iinfo(a.dtype).max
421    else:
422      a = np.clip(a, 0.0, 1.0)
423      src_max = 1.0
424    dst_max = np.iinfo(dtype).max  # pyrefly: ignore[no-matching-overload]
425    if dst_max <= np.iinfo(np.uint16).max:
426      scale = np.array(dst_max / src_max, dtype=np.float32)
427      result = (a * scale + 0.5).astype(dtype)
428    elif dst_max <= np.iinfo(np.uint32).max:
429      result = (a.astype(np.float64) * (dst_max / src_max) + 0.5).astype(dtype)
430    else:
431      # https://stackoverflow.com/a/66306123/
432      a = a.astype(np.float64) * (dst_max / src_max) + 0.5
433      dst = np.atleast_1d(a)
434      values_too_large = dst >= np.float64(dst_max)
435      with np.errstate(invalid='ignore'):
436        dst = dst.astype(dtype)
437      dst[values_too_large] = dst_max
438      result = dst if a.ndim > 0 else dst[0]
439  else:
440    assert np.issubdtype(dtype, np.floating)
441    result = a.astype(dtype)
442    if np.issubdtype(a.dtype, np.unsignedinteger):
443      result = result / dtype.type(np.iinfo(a.dtype).max)
444  return result

Returns media array converted to specified type.

A "media array" is one in which the dtype is either a floating-point type (np.float32 or np.float64) or an unsigned integer type. The array values are assumed to lie in the range [0.0, 1.0] for floating-point values, and in the full range for unsigned integers, e.g. [0, 255] for np.uint8.

Conversion between integers and floats maps uint(0) to 0.0 and uint(MAX) to 1.0. The input array may also be of type bool, whereby True maps to uint(MAX) or 1.0. The values are scaled and clamped as appropriate during type conversions.

Arguments:
  • array: Input array-like object (floating-point, unsigned int, or bool).
  • dtype: Desired output type (floating-point or unsigned int).
Returns:

Array a if it is already of the specified dtype, else a converted array.

def to_float01( a: ArrayLike, dtype: DTypeLike = <class 'numpy.float32'>) -> np.ndarray:
447def to_float01(a: _ArrayLike, dtype: _DTypeLike = np.float32) -> _NDArray:
448  """If array has unsigned integers, rescales them to the range [0.0, 1.0].
449
450  Scaling is such that uint(0) maps to 0.0 and uint(MAX) maps to 1.0.  See
451  `to_type`.
452
453  Args:
454    a: Input array.
455    dtype: Desired floating-point type if rescaling occurs.
456
457  Returns:
458    A new array of dtype values in the range [0.0, 1.0] if the input array `a`
459    contains unsigned integers; otherwise, array `a` is returned unchanged.
460  """
461  a = np.asarray(a)
462  dtype = np.dtype(dtype)
463  if not np.issubdtype(dtype, np.floating):
464    raise ValueError(f'Type {dtype} is not floating-point.')
465  if np.issubdtype(a.dtype, np.floating):
466    return a
467  return to_type(a, dtype)

If array has unsigned integers, rescales them to the range [0.0, 1.0].

Scaling is such that uint(0) maps to 0.0 and uint(MAX) maps to 1.0. See to_type.

Arguments:
  • a: Input array.
  • dtype: Desired floating-point type if rescaling occurs.
Returns:

A new array of dtype values in the range [0.0, 1.0] if the input array a contains unsigned integers; otherwise, array a is returned unchanged.

def to_uint8(a: ArrayLike) -> np.ndarray:
470def to_uint8(a: _ArrayLike) -> _NDArray:
471  """Returns array converted to uint8 values; see `to_type`."""
472  return to_type(a, np.uint8)

Returns array converted to uint8 values; see to_type.

def set_output_height(num_pixels: int) -> None:
340def set_output_height(num_pixels: int) -> None:
341  """Overrides the height of the current output cell, if using Colab."""
342  try:
343    # We want to fail gracefully for non-Colab IPython notebooks.
344    output = importlib.import_module('google.colab.output')
345    s = f'google.colab.output.setIframeHeight("{num_pixels}px")'
346    output.eval_js(s)
347  except (ModuleNotFoundError, AttributeError):
348    pass

Overrides the height of the current output cell, if using Colab.

def set_max_output_height(num_pixels: int) -> None:
351def set_max_output_height(num_pixels: int) -> None:
352  """Sets the maximum height of the current output cell, if using Colab."""
353  try:
354    # We want to fail gracefully for non-Colab IPython notebooks.
355    output = importlib.import_module('google.colab.output')
356    s = (
357        'google.colab.output.setIframeHeight('
358        f'0, true, {{maxHeight: {num_pixels}}})'
359    )
360    output.eval_js(s)
361  except (ModuleNotFoundError, AttributeError):
362    pass

Sets the maximum height of the current output cell, if using Colab.

def color_ramp( shape: tuple[int, int] = (64, 64), *, dtype: DTypeLike = <class 'numpy.float32'>) -> np.ndarray:
478def color_ramp(
479    shape: tuple[int, int] = (64, 64), *, dtype: _DTypeLike = np.float32
480) -> _NDArray:
481  """Returns an image of a red-green color gradient.
482
483  This is useful for quick experimentation and testing.  See also
484  `moving_circle` to generate a sample video.
485
486  Args:
487    shape: 2D spatial dimensions (height, width) of generated image.
488    dtype: Type (uint or floating) of resulting pixel values.
489  """
490  _check_2d_shape(shape)
491  dtype = _as_valid_media_type(dtype)
492  yx = (np.moveaxis(np.indices(shape), 0, -1) + 0.5) / shape
493  image = np.insert(yx, 2, 0.0, axis=-1)
494  return to_type(image, dtype)

Returns an image of a red-green color gradient.

This is useful for quick experimentation and testing. See also moving_circle to generate a sample video.

Arguments:
  • shape: 2D spatial dimensions (height, width) of generated image.
  • dtype: Type (uint or floating) of resulting pixel values.
def moving_circle( shape: tuple[int, int] = (256, 256), num_images: int = 10, *, dtype: DTypeLike = <class 'numpy.float32'>) -> np.ndarray:
497def moving_circle(
498    shape: tuple[int, int] = (256, 256),
499    num_images: int = 10,
500    *,
501    dtype: _DTypeLike = np.float32,
502) -> _NDArray:
503  """Returns a video of a circle moving in front of a color ramp.
504
505  This is useful for quick experimentation and testing.  See also `color_ramp`
506  to generate a sample image.
507
508  >>> show_video(moving_circle((480, 640), 60), fps=60)
509
510  Args:
511    shape: 2D spatial dimensions (height, width) of generated video.
512    num_images: Number of video frames.
513    dtype: Type (uint or floating) of resulting pixel values.
514  """
515  _check_2d_shape(shape)
516  dtype = np.dtype(dtype)
517
518  def generate_image(image_index: int) -> _NDArray:
519    """Returns a video frame image."""
520    image = color_ramp(shape, dtype=dtype)
521    yx = np.moveaxis(np.indices(shape), 0, -1)
522    center = shape[0] * 0.6, shape[1] * (image_index + 0.5) / num_images
523    radius_squared = (min(shape) * 0.1) ** 2
524    inside = np.sum((yx - center) ** 2, axis=-1) < radius_squared
525    white_circle_color = 1.0, 1.0, 1.0
526    if np.issubdtype(dtype, np.unsignedinteger):
527      white_circle_color = to_type([white_circle_color], dtype)[0]
528    image[inside] = white_circle_color
529    return image
530
531  return np.array([generate_image(i) for i in range(num_images)])

Returns a video of a circle moving in front of a color ramp.

This is useful for quick experimentation and testing. See also color_ramp to generate a sample image.

>>> show_video(moving_circle((480, 640), 60), fps=60)
Arguments:
  • shape: 2D spatial dimensions (height, width) of generated video.
  • num_images: Number of video frames.
  • dtype: Type (uint or floating) of resulting pixel values.
class set_show_save_dir:
764class set_show_save_dir:  # pylint: disable=invalid-name
765  """Save all titled output from `show_*()` calls into files.
766
767  If the specified `directory` is not None, all titled images and videos
768  displayed by `show_image`, `show_images`, `show_video`, and `show_videos` are
769  also saved as files within the directory.
770
771  It can be used either to set the state or as a context manager:
772
773  >>> set_show_save_dir('/tmp')
774  >>> show_image(color_ramp(), title='image1')  # Creates /tmp/image1.png.
775  >>> show_video(moving_circle(), title='video2')  # Creates /tmp/video2.mp4.
776  >>> set_show_save_dir(None)
777
778  >>> with set_show_save_dir('/tmp'):
779  ...   show_image(color_ramp(), title='image1')  # Creates /tmp/image1.png.
780  ...   show_video(moving_circle(), title='video2')  # Creates /tmp/video2.mp4.
781  """
782
783  def __init__(self, directory: _Path | None):
784    self._old_show_save_dir = _config.show_save_dir
785    _config.show_save_dir = directory
786
787  def __enter__(self) -> None:
788    pass
789
790  def __exit__(self, *_: Any) -> None:
791    _config.show_save_dir = self._old_show_save_dir

Save all titled output from show_*() calls into files.

If the specified directory is not None, all titled images and videos displayed by show_image, show_images, show_video, and show_videos are also saved as files within the directory.

It can be used either to set the state or as a context manager:

>>> set_show_save_dir('/tmp')
>>> show_image(color_ramp(), title='image1')  # Creates /tmp/image1.png.
>>> show_video(moving_circle(), title='video2')  # Creates /tmp/video2.mp4.
>>> set_show_save_dir(None)
>>> with set_show_save_dir('/tmp'):
...   show_image(color_ramp(), title='image1')  # Creates /tmp/image1.png.
...   show_video(moving_circle(), title='video2')  # Creates /tmp/video2.mp4.
set_show_save_dir(directory: str | os.PathLike[str] | None)
783  def __init__(self, directory: _Path | None):
784    self._old_show_save_dir = _config.show_save_dir
785    _config.show_save_dir = directory
def set_ffmpeg(name_or_path: str | os.PathLike[str]) -> None:
326def set_ffmpeg(name_or_path: _Path) -> None:
327  """Specifies the name or path for the `ffmpeg` external program.
328
329  The `ffmpeg` program is required for compressing and decompressing video.
330  (It is used in `read_video`, `write_video`, `show_video`, `show_videos`,
331  etc.)
332
333  Args:
334    name_or_path: Either a filename within a directory of `os.environ['PATH']`
335      or a filepath.  The default setting is 'ffmpeg'.
336  """
337  _config.ffmpeg_name_or_path = name_or_path

Specifies the name or path for the ffmpeg external program.

The ffmpeg program is required for compressing and decompressing video. (It is used in read_video, write_video, show_video, show_videos, etc.)

Arguments:
  • name_or_path: Either a filename within a directory of os.environ['PATH'] or a filepath. The default setting is 'ffmpeg'.
def video_is_available() -> bool:
1294def video_is_available() -> bool:
1295  """Returns True if the program `ffmpeg` is found.
1296
1297  See also `set_ffmpeg`.
1298  """
1299  return _search_for_ffmpeg_path() is not None

Returns True if the program ffmpeg is found.

See also set_ffmpeg.