diff --git a/js/chooselattice/index.tsx b/js/chooselattice/index.tsx index 88cdcb1a..6773617c 100644 --- a/js/chooselattice/index.tsx +++ b/js/chooselattice/index.tsx @@ -14,11 +14,19 @@ import Typography from "@mui/material/Typography"; import Stack from "@mui/material/Stack"; import Button from "@mui/material/Button"; import { useTheme } from "../theme"; -import { extractBytes, preserveRestoredWidgetModelsOnSave } from "../format"; +import { preserveRestoredWidgetModelsOnSave } from "../format"; import { useHideStaticFallback } from "../staticFallback"; +import { + canvasPoint, + clamp, + drawImage, + imageToScreen, + screenToImage, + usePngBitmap, + zoomAt, + type ImageViewport, +} from "../imageView"; -const MIN_ZOOM = 0.5; -const MAX_ZOOM = 20; const CANVAS_SIZE = 512; const CANVAS_BORDER_PX = 1; const HIT_PX = 10; @@ -37,10 +45,6 @@ const compactButton = { type Point = [number, number]; // [row, col] in original image pixels type DragMode = "none" | "pan" | "point"; -function clamp(value: number, lo: number, hi: number): number { - return Math.min(hi, Math.max(lo, value)); -} - function ChooseLattice() { const model = useModel(); const rootRef = React.useRef(null); @@ -56,38 +60,19 @@ function ChooseLattice() { const [pointLabels] = useModelState("point_labels"); const [points, setPoints] = useModelState("points"); - // Decode the PNG payload once per change into a drawable bitmap. - const [image, setImage] = React.useState(null); - React.useEffect(() => { - const bytes = extractBytes(frameBytes); - if (bytes.length === 0) { - setImage(null); - return; - } - let cancelled = false; - const blob = new Blob([bytes as unknown as BlobPart], { type: "image/png" }); - if (typeof createImageBitmap === "function") { - createImageBitmap(blob).then((bmp) => { if (!cancelled) setImage(bmp); }); - } else { - const url = URL.createObjectURL(blob); - const img = new Image(); - img.onload = () => { if (!cancelled) setImage(img); URL.revokeObjectURL(url); }; - img.src = url; - } - return () => { cancelled = true; }; - }, [frameBytes]); + const image = usePngBitmap(frameBytes); // View state: zoom + pan (CSS px, canvas-centered). const [zoom, setZoom] = React.useState(1); const [panX, setPanX] = React.useState(0); const [panY, setPanY] = React.useState(0); - // displayScale maps original image pixels -> CSS px at zoom=1. - const displayScale = height > 0 && width > 0 - ? CANVAS_SIZE / Math.max(height, width) - : 1; const canvasW = CANVAS_SIZE; const canvasH = CANVAS_SIZE; + const viewport: ImageViewport = React.useMemo( + () => ({ height, width, canvas: CANVAS_SIZE, zoom, panX, panY }), + [height, width, zoom, panX, panY], + ); const canvasRef = React.useRef(null); const uiRef = React.useRef(null); @@ -116,13 +101,7 @@ function ChooseLattice() { ctx.fillStyle = themeColors.bg; ctx.fillRect(0, 0, canvasW, canvasH); if (!image || !width || !height) return; - const cx = canvasW / 2; - const cy = canvasH / 2; - const drawW = width * displayScale * zoom; - const drawH = height * displayScale * zoom; - const x = cx - drawW / 2 + panX; - const y = cy - drawH / 2 + panY; - ctx.drawImage(image, x, y, drawW, drawH); + drawImage(ctx, image, viewport); // Confirm the decoded bitmap on the next compositing frame. Static docs // can mount while Chrome is still promoting the canvas layer; without a // second paint that one-shot draw can remain a white presentation frame. @@ -132,42 +111,33 @@ function ChooseLattice() { if (!current || !currentCtx) return; currentCtx.fillStyle = themeColors.bg; currentCtx.fillRect(0, 0, canvasW, canvasH); - currentCtx.drawImage(image, x, y, drawW, drawH); + drawImage(currentCtx, image, viewport); }); return () => window.cancelAnimationFrame(confirmFrame); - }, [image, width, height, displayScale, zoom, panX, panY, canvasW, canvasH, themeColors.bg]); + }, [image, width, height, viewport, canvasW, canvasH, themeColors.bg]); // Convert a mouse event to original-image (row, col) coordinates. const screenToImg = React.useCallback((e: { clientX: number; clientY: number }): Point => { const canvas = canvasRef.current; if (!canvas) return [0, 0]; - const rect = canvas.getBoundingClientRect(); - const mouseCanvasX = (e.clientX - rect.left) * (canvas.width / rect.width); - const mouseCanvasY = (e.clientY - rect.top) * (canvas.height / rect.height); - const cx = canvasW / 2; - const cy = canvasH / 2; - const col = (mouseCanvasX - cx - panX) / (displayScale * zoom) + width / 2; - const row = (mouseCanvasY - cy - panY) / (displayScale * zoom) + height / 2; - return [row, col]; - }, [canvasW, canvasH, panX, panY, displayScale, zoom, width, height]); - - const imgToScreen = React.useCallback((row: number, col: number): [number, number] => { - const cx = canvasW / 2; - const cy = canvasH / 2; - const x = cx + (col - width / 2) * displayScale * zoom + panX; - const y = cy + (row - height / 2) * displayScale * zoom + panY; - return [x, y]; - }, [canvasW, canvasH, panX, panY, displayScale, zoom, width, height]); + const [x, y] = canvasPoint(canvas, e); + return screenToImage(viewport, x, y); + }, [viewport]); + + const imgToScreen = React.useCallback( + (row: number, col: number): [number, number] => imageToScreen(viewport, row, col), + [viewport], + ); const hitTestPoint = React.useCallback((row: number, col: number): number => { - const hitArea = HIT_PX / (displayScale * zoom); + const hitArea = HIT_PX / ((canvasW / Math.max(height, width)) * zoom); const list = points || []; for (let i = list.length - 1; i >= 0; i--) { const [pr, pc] = list[i]; if (Math.hypot(row - pr, col - pc) <= hitArea) return i; } return -1; - }, [points, displayScale, zoom]); + }, [points, canvasW, height, width, zoom]); // Wheel: cursor-anchored zoom. Page-scroll prevention is handled by a // native non-passive listener below (React's synthetic onWheel is passive, @@ -175,18 +145,11 @@ function ChooseLattice() { const handleWheel = (e: React.WheelEvent) => { const canvas = canvasRef.current; if (!canvas) return; - const rect = canvas.getBoundingClientRect(); - const mouseCanvasX = (e.clientX - rect.left) * (canvas.width / rect.width); - const mouseCanvasY = (e.clientY - rect.top) * (canvas.height / rect.height); - const cx = canvasW / 2; - const cy = canvasH / 2; - const mouseImageX = (mouseCanvasX - cx - panX) / zoom + cx; - const mouseImageY = (mouseCanvasY - cy - panY) / zoom + cy; - const zoomFactor = e.deltaY > 0 ? 0.9 : 1.1; - const newZoom = clamp(zoom * zoomFactor, MIN_ZOOM, MAX_ZOOM); - setPanX(mouseCanvasX - (mouseImageX - cx) * newZoom - cx); - setPanY(mouseCanvasY - (mouseImageY - cy) * newZoom - cy); - setZoom(newZoom); + const [x, y] = canvasPoint(canvas, e); + const next = zoomAt(viewport, x, y, e.deltaY); + setPanX(next.panX); + setPanY(next.panY); + setZoom(next.zoom); }; const resetView = React.useCallback(() => { diff --git a/js/imageView.test.ts b/js/imageView.test.ts new file mode 100644 index 00000000..47115dee --- /dev/null +++ b/js/imageView.test.ts @@ -0,0 +1,22 @@ +import { describe, expect, it } from "vitest"; +import { imageToScreen, screenToImage, zoomAt, type ImageViewport } from "./imageView"; + +const view: ImageViewport = { height: 48, width: 64, canvas: 512, zoom: 1, panX: 0, panY: 0 }; + +describe("imageView", () => { + it("round-trips image and screen coordinates under pan and zoom", () => { + const panned: ImageViewport = { ...view, zoom: 3.2, panX: -40, panY: 17 }; + const [x, y] = imageToScreen(panned, 12.5, 30.25); + const [row, col] = screenToImage(panned, x, y); + expect(row).toBeCloseTo(12.5); + expect(col).toBeCloseTo(30.25); + }); + + it("keeps the point under the cursor fixed while zooming", () => { + const before = screenToImage(view, 300, 200); + const zoomed: ImageViewport = { ...view, ...zoomAt(view, 300, 200, -1) }; + const after = screenToImage(zoomed, 300, 200); + expect(after[0]).toBeCloseTo(before[0]); + expect(after[1]).toBeCloseTo(before[1]); + }); +}); diff --git a/js/imageView.ts b/js/imageView.ts new file mode 100644 index 00000000..3515d62e --- /dev/null +++ b/js/imageView.ts @@ -0,0 +1,135 @@ +/** Shared pan/zoom transform for the single-canvas image widgets. */ + +import { useEffect, useState } from "react"; +import { extractBytes } from "./format"; + +export const MIN_ZOOM = 0.5; +export const MAX_ZOOM = 20; + +export interface ImageViewport { + /** Original image rows and columns. */ + height: number; + width: number; + /** Square canvas edge, in CSS px. */ + canvas: number; + zoom: number; + panX: number; + panY: number; +} + +export type ViewTransform = Pick; + +export const IDENTITY_VIEW: ViewTransform = { zoom: 1, panX: 0, panY: 0 }; + +export function clamp(value: number, lo: number, hi: number): number { + return Math.min(hi, Math.max(lo, value)); +} + +/** Image pixels per CSS px at zoom 1, fitting the longest side to the canvas. */ +export function displayScale(v: ImageViewport): number { + const longest = Math.max(v.height, v.width); + return longest > 0 ? v.canvas / longest : 1; +} + +/** Canvas coordinates of a mouse event, corrected for CSS scaling. */ +export function canvasPoint( + canvas: HTMLCanvasElement, + e: { clientX: number; clientY: number }, +): [number, number] { + const rect = canvas.getBoundingClientRect(); + return [ + (e.clientX - rect.left) * (canvas.width / rect.width), + (e.clientY - rect.top) * (canvas.height / rect.height), + ]; +} + +export function screenToImage( + v: ImageViewport, + canvasX: number, + canvasY: number, +): [number, number] { + const scale = displayScale(v) * v.zoom; + const center = v.canvas / 2; + const col = (canvasX - center - v.panX) / scale + v.width / 2; + const row = (canvasY - center - v.panY) / scale + v.height / 2; + return [row, col]; +} + +export function imageToScreen( + v: ImageViewport, + row: number, + col: number, +): [number, number] { + const scale = displayScale(v) * v.zoom; + const center = v.canvas / 2; + return [ + center + (col - v.width / 2) * scale + v.panX, + center + (row - v.height / 2) * scale + v.panY, + ]; +} + +/** Wheel zoom that keeps the image point under the cursor fixed. */ +export function zoomAt( + v: ImageViewport, + canvasX: number, + canvasY: number, + deltaY: number, +): ViewTransform { + const center = v.canvas / 2; + const anchorX = (canvasX - center - v.panX) / v.zoom + center; + const anchorY = (canvasY - center - v.panY) / v.zoom + center; + const zoom = clamp(v.zoom * (deltaY > 0 ? 0.9 : 1.1), MIN_ZOOM, MAX_ZOOM); + return { + zoom, + panX: canvasX - (anchorX - center) * zoom - center, + panY: canvasY - (anchorY - center) * zoom - center, + }; +} + +/** Paint the image into the canvas under the current view. */ +export function drawImage( + ctx: CanvasRenderingContext2D, + image: CanvasImageSource, + v: ImageViewport, +): void { + const scale = displayScale(v) * v.zoom; + const center = v.canvas / 2; + const drawW = v.width * scale; + const drawH = v.height * scale; + ctx.drawImage(image, center - drawW / 2 + v.panX, center - drawH / 2 + v.panY, drawW, drawH); +} + +/** Decode PNG bytes from a synced trait into a drawable bitmap. */ +export function usePngBitmap( + bytes: DataView | Uint8Array | null | undefined, +): ImageBitmap | HTMLImageElement | null { + const [image, setImage] = useState(null); + + useEffect(() => { + const raw = bytes ? extractBytes(bytes) : new Uint8Array(0); + if (raw.length === 0) { + setImage(null); + return; + } + let cancelled = false; + const blob = new Blob([raw as unknown as BlobPart], { type: "image/png" }); + if (typeof createImageBitmap === "function") { + createImageBitmap(blob).then((bitmap) => { + if (!cancelled) setImage(bitmap); + }); + } else { + const url = URL.createObjectURL(blob); + const element = new Image(); + element.onload = () => { + if (!cancelled) setImage(element); + URL.revokeObjectURL(url); + }; + element.src = url; + } + return () => { + cancelled = true; + }; + }, [bytes]); + + return image; +} diff --git a/src/quantem/widget/choose_lattice.py b/src/quantem/widget/choose_lattice.py index ec4a44a3..378d6ed1 100644 --- a/src/quantem/widget/choose_lattice.py +++ b/src/quantem/widget/choose_lattice.py @@ -12,9 +12,10 @@ from typing import Any, Sequence import anywidget -import matplotlib import numpy as np import traitlets +from quantem.widget.render import frame_to_rgb, rgb_to_png_bytes +from quantem.widget.utils.traits import reject_unknown_kwargs from quantem.widget.utils.array import to_numpy from quantem.widget.utils.static_fallback import StaticFallbackMixin @@ -36,45 +37,6 @@ def _core_image_dataset_types() -> tuple[type[Any], ...]: return _CORE_IMAGE_DATASET_TYPES -def _reject_unknown_kwargs(cls, kwargs: dict) -> None: - """Raise TypeError for any kwarg that isn't a declared trait (catches typos).""" - traits = set(cls.class_trait_names()) - unknown = [k for k in kwargs if k not in traits] - if unknown: - key = sorted(unknown)[0] - raise TypeError(f"{cls.__name__}() got unexpected keyword argument {key!r}.") - - -def _frame_to_rgb( - frame: np.ndarray, - *, - cmap: str, - vmin: float | None, - vmax: float | None, - log_scale: bool, -) -> np.ndarray: - """Colormap a 2D float frame into a uint8 (H, W, 3) RGB array.""" - values = frame.astype(np.float64, copy=False) - if log_scale: - values = np.log1p(np.clip(values - np.nanmin(values), 0, None)) - lo = float(np.nanpercentile(values, 1)) if vmin is None else float(vmin) - hi = float(np.nanpercentile(values, 99)) if vmax is None else float(vmax) - if hi <= lo: - hi = lo + 1.0 - normalized = np.clip((values - lo) / (hi - lo), 0.0, 1.0) - colormap = matplotlib.colormaps[cmap] - rgba = colormap(normalized) - return (rgba[..., :3] * 255).astype(np.uint8) - - -def _rgb_to_png_bytes(rgb: np.ndarray) -> bytes: - import io - from PIL import Image - buf = io.BytesIO() - Image.fromarray(rgb, mode="RGB").save(buf, format="PNG") - return buf.getvalue() - - class ChooseLattice(StaticFallbackMixin, anywidget.AnyWidget): """Interactive picker for an ordered origin + two lattice-vector points. @@ -155,7 +117,7 @@ def __init__( notebook_preview_max_px: int = 512, **kwargs, ) -> None: - _reject_unknown_kwargs(type(self), kwargs) + reject_unknown_kwargs(type(self), kwargs) super().__init__(**kwargs) core_image_dataset_types = _core_image_dataset_types() @@ -176,8 +138,8 @@ def __init__( ) self._data = frame - rgb = _frame_to_rgb(frame, cmap=cmap, vmin=vmin, vmax=vmax, log_scale=log_scale) - png_bytes = _rgb_to_png_bytes(rgb) + rgb = frame_to_rgb(frame, cmap=cmap, vmin=vmin, vmax=vmax, log_scale=log_scale) + png_bytes = rgb_to_png_bytes(rgb) self._configure_static_fallback( notebook_preview_format=notebook_preview_format, diff --git a/src/quantem/widget/render/__init__.py b/src/quantem/widget/render/__init__.py index 6b78a1be..31c6d6fe 100644 --- a/src/quantem/widget/render/__init__.py +++ b/src/quantem/widget/render/__init__.py @@ -1,5 +1,6 @@ """Rendering/export helpers shared by widget classes.""" +from quantem.widget.render.frame import frame_to_rgb, rgb_to_png_bytes from quantem.widget.render.thumbnail import ( save_thumbnail, thumbnail_bytes, @@ -8,6 +9,8 @@ ) __all__ = [ + "frame_to_rgb", + "rgb_to_png_bytes", "save_thumbnail", "thumbnail_bytes", "thumbnail_image", diff --git a/src/quantem/widget/render/frame.py b/src/quantem/widget/render/frame.py new file mode 100644 index 00000000..5db33728 --- /dev/null +++ b/src/quantem/widget/render/frame.py @@ -0,0 +1,59 @@ +"""Colormapping and PNG encoding for widget image panels.""" + +from __future__ import annotations + +import io + +import matplotlib +import numpy as np + +NO_DATA_COLOR = "#404040" + + +def frame_to_rgb( + frame: np.ndarray, + *, + cmap: str, + vmin: float | None = None, + vmax: float | None = None, + log_scale: bool = False, +) -> np.ndarray: + """Colormap a 2D float frame into a uint8 ``(H, W, 3)`` RGB array. + + Parameters + ---------- + frame : np.ndarray + ``(H, W)`` scalar image. + cmap : str + Matplotlib colormap name. + vmin, vmax : float, optional + Explicit display range. Defaults to a robust 1st/99th percentile + auto-contrast when not given. + log_scale : bool, default=False + Apply a ``log1p`` stretch before contrast scaling. + + Returns + ------- + np.ndarray + ``(H, W, 3)`` uint8 RGB image. + """ + values = frame.astype(np.float64, copy=False) + if log_scale: + values = np.log1p(np.clip(values - np.nanmin(values), 0, None)) + lo = float(np.nanpercentile(values, 1)) if vmin is None else float(vmin) + hi = float(np.nanpercentile(values, 99)) if vmax is None else float(vmax) + if hi <= lo: + hi = lo + 1.0 + normalized = np.clip((values - lo) / (hi - lo), 0.0, 1.0) + colormap = matplotlib.colormaps[cmap].with_extremes(bad=NO_DATA_COLOR) + rgba = colormap(normalized) + return (rgba[..., :3] * 255).astype(np.uint8) + + +def rgb_to_png_bytes(rgb: np.ndarray) -> bytes: + """Encode a uint8 ``(H, W, 3)`` RGB array as PNG bytes.""" + from PIL import Image + + buf = io.BytesIO() + Image.fromarray(rgb, mode="RGB").save(buf, format="PNG") + return buf.getvalue() diff --git a/src/quantem/widget/utils/traits.py b/src/quantem/widget/utils/traits.py new file mode 100644 index 00000000..841e29d3 --- /dev/null +++ b/src/quantem/widget/utils/traits.py @@ -0,0 +1,12 @@ +"""Trait-related helpers shared by widget constructors.""" + +from __future__ import annotations + + +def reject_unknown_kwargs(cls, kwargs: dict) -> None: + """Raise TypeError for any kwarg that isn't a declared trait (catches typos).""" + traits = set(cls.class_trait_names()) + unknown = [k for k in kwargs if k not in traits] + if unknown: + key = sorted(unknown)[0] + raise TypeError(f"{cls.__name__}() got unexpected keyword argument {key!r}.")