LFR-GAN / g_render /models /structure /heatmap.py
anbinh39's picture
Add g_render2 (LFR-GAN) weights + code + model card
dfc2650 verified
Raw
History Blame Contribute Delete
12.4 kB
from __future__ import annotations
from collections.abc import Sequence
import torch
import torch.nn.functional as F
Point = Sequence[float] # (x, y) or (x, y, theta_degrees)
RS = "\x1e"
# Pattern-class vocabulary (EFS 9.307): whorl / loops / arches / unclassifiable.
# Mapped to a fixed index for global one-hot conditioning. Sub-codes (e.g. "LS",
# "RS", "WU", "AU", "TA", "PA") collapse to their family by first letter group.
PATTERN_CLASSES = ["whorl", "loop", "arch", "other"]
def pattern_to_index(code: str) -> int:
"""Map an EFS 9.307 pattern code to a PATTERN_CLASSES index (whorl/loop/arch/other)."""
c = (code or "").strip().upper().split("\x1f")[0].split("\x1e")[0][:2]
if c.startswith("W"):
return 0 # whorl (WU, etc.)
if c in ("LS", "RS") or c.startswith("L") or c.startswith("R"):
return 1 # loop (left/right slant)
if c.startswith("A") or c in ("TA", "PA"):
return 2 # arch (tented/plain)
return 3 # UC / unclassifiable / other
def pattern_onehot(code: str) -> list[float]:
v = [0.0] * len(PATTERN_CLASSES)
v[pattern_to_index(code)] = 1.0
return v
def compute_orientation_field(
image: torch.Tensor, block: int = 8, sigma_smooth: int = 9
) -> torch.Tensor:
"""Ridge orientation field via the gradient structure tensor.
Returns ``(B, 2, H, W)`` = ``[cos(2θ), sin(2θ)] * coherence``. Using the
double-angle (2θ) encoding avoids the orientation wrap-around at 0/π, and
weighting by coherence down-weights noisy/blank regions (low coherence → ~0),
so it is a clean dense conditioning/loss signal that is registered to the image
(computed from the image itself, no annotation needed). Identity-bearing: it
captures ridge flow, the basis of pattern class and singular points.
"""
if image.ndim != 4:
raise ValueError(f"expected (B,C,H,W), got {tuple(image.shape)}")
g = image if image.shape[1] == 1 else image.mean(1, keepdim=True)
kx = torch.tensor([[-1, 0, 1], [-2, 0, 2], [-1, 0, 1]], dtype=g.dtype, device=g.device).view(1, 1, 3, 3)
ky = kx.transpose(-1, -2)
# Replicate-pad before Sobel: zero padding fabricates strong gradients along
# the image border (a flat/blank image otherwise gets a high-coherence ring,
# which leaks through the foreground gate and pollutes orientation losses).
g = F.pad(g, (1, 1, 1, 1), mode="replicate")
gx = F.conv2d(g, kx)
gy = F.conv2d(g, ky)
gxx, gyy, gxy = gx * gx, gy * gy, gx * gy
pad = sigma_smooth // 2
sm = lambda t: F.avg_pool2d(t, sigma_smooth, stride=1, padding=pad)
gxx, gyy, gxy = sm(gxx), sm(gyy), sm(gxy)
cos2 = gxx - gyy
sin2 = 2 * gxy
# Epsilon INSIDE the sqrt: d/dx sqrt(x) -> inf at x=0, so sqrt(0)+eps still
# backprops NaN on exactly-flat regions. This function runs on generator
# OUTPUTS (regen orientation loss, equalized critic cond), and the foreground
# gate makes outputs exactly equal flat inputs there — sqrt(0) killed a full
# training run (every weight NaN within epoch 0) before this guard.
mag = torch.sqrt(cos2 * cos2 + sin2 * sin2 + 1e-12) + 1e-6
coh = mag / (gxx + gyy + 1e-6) # coherence in [0,1]
return torch.cat([cos2 / mag * coh, sin2 / mag * coh], dim=1)
def parse_quality_grid(value: str) -> list[list[int]]:
"""Parse a Type-9 ``9.308`` ridge-quality grid (RS-separated digit rows).
Returns a rectangular int grid (rows of equal length); ragged input is
right-padded with zeros. Values are quality levels (0=background .. 5).
"""
rows = [row for row in value.split(RS) if row]
grid = [[int(c) for c in row if c.isdigit()] for row in rows]
grid = [r for r in grid if r]
if not grid:
return []
width = max(len(r) for r in grid)
return [r + [0] * (width - len(r)) for r in grid]
def render_quality_grid(
grid: str | list[list[int]],
out_size: int,
max_level: int = 5,
binarize: bool = False,
device: torch.device | str = "cpu",
dtype: torch.dtype = torch.float32,
) -> torch.Tensor:
"""Render a ``9.308`` quality grid as a dense ``(1, out, out)`` channel.
The grid spans the full annotation image frame, so it is spatially registered
to the (uniformly-scaled) latent image and is upsampled to ``out_size``. With
``binarize=True`` it returns a foreground mask (level>0), suitable as the
``M_fg`` for ``L_seg``; otherwise quality is normalised to ``[0,1]``.
Empty grid -> all-ones (no annotation = treat everything as valid, the safe
default for a mask) when binarized, else zeros.
"""
if isinstance(grid, str):
grid = parse_quality_grid(grid)
if not grid:
fill = 1.0 if binarize else 0.0
return torch.full((1, out_size, out_size), fill, device=device, dtype=dtype)
t = torch.tensor(grid, device=device, dtype=dtype).unsqueeze(0).unsqueeze(0)
if binarize:
t = (t > 0).to(dtype)
else:
t = t.clamp(0, max_level) / float(max_level)
# pad_square-aware (match resize_pad_square): the grid spans the frame, so its aspect
# = frame aspect -> uniform-resize the longer side then center-pad, NOT anisotropic
# stretch (which misregistered the quality/ROI mask vs the padded latent image).
hg, wg = t.shape[-2], t.shape[-1]
s = out_size / max(hg, wg)
nh, nw = max(1, round(hg * s)), max(1, round(wg * s))
t = F.interpolate(t, size=(nh, nw), mode="bilinear", align_corners=False)
canvas = torch.zeros(1, 1, out_size, out_size, device=device, dtype=dtype)
oy, ox = (out_size - nh) // 2, (out_size - nw) // 2
canvas[:, :, oy:oy + nh, ox:ox + nw] = t
return canvas[0]
def structural_channel_names(orientation: bool = False) -> list[str]:
"""Channel order produced by :func:`build_structural_conditioning`."""
names = ["minutiae", "core", "delta"]
if orientation:
names += ["minutiae_sin", "minutiae_cos"]
return names
def render_gaussian_points(
points: Sequence[Point],
out_size: int,
src_size: tuple[float, float] | None = None,
sigma: float = 3.0,
weights: Sequence[float] | None = None,
device: torch.device | str = "cpu",
dtype: torch.dtype = torch.float32,
) -> torch.Tensor:
"""Render points as a single-channel Gaussian heatmap ``(1, out, out)``.
Coordinates are in the pixel system of the source annotation image. Pass
``src_size=(W, H)`` (the EFS record's image size, e.g. 9.300) to rescale onto
the ``out_size`` training grid; omit it if points are already in ``out_size``
pixels. Peaks are 1.0 at each point (or ``weights[i]`` if given). An empty
point list yields an all-zero map, which is the correct conditioning for a
latent that genuinely has no annotated core/delta.
"""
hm = torch.zeros(1, out_size, out_size, device=device, dtype=dtype)
if not points:
return hm
# pad_square-aware mapping: MATCH FingerprintTransform.resize_pad_square (uniform
# scale of the longer side + centered pad), NOT an anisotropic stretch. The old
# anisotropic (out/W, out/H) misregistered minutiae vs the padded image by mean ~8px
# (up to 38px, 14% of pairs >15px tol) on non-square latents -> the minutiae anchor
# rewarded ridge energy at the WRONG pixels.
if src_size:
s = out_size / max(float(src_size[0]), float(src_size[1]))
sx = sy = s
ox = (out_size - float(src_size[0]) * s) / 2.0
oy = (out_size - float(src_size[1]) * s) / 2.0
else:
sx = sy = 1.0
ox = oy = 0.0
yy, xx = torch.meshgrid(
torch.arange(out_size, device=device, dtype=dtype),
torch.arange(out_size, device=device, dtype=dtype),
indexing="ij",
)
two_s2 = 2.0 * float(sigma) * float(sigma)
for i, point in enumerate(points):
x = float(point[0]) * sx + ox
y = float(point[1]) * sy + oy
if not (0 <= x < out_size and 0 <= y < out_size):
continue
w = float(weights[i]) if weights is not None else 1.0
hm[0] += w * torch.exp(-((xx - x) ** 2 + (yy - y) ** 2) / two_s2)
return hm.clamp(0.0, 1.0)
def conditioning_from_sidecar(
sidecar: dict,
out_size: int,
channels: Sequence[str] = ("quality", "minutiae"),
prefix: str = "latent",
sigma_minutiae: float = 3.0,
sigma_singular: float = 6.0,
image: torch.Tensor | None = None,
device: torch.device | str = "cpu",
dtype: torch.dtype = torch.float32,
) -> torch.Tensor:
"""Build a registered conditioning tensor ``(len(channels), out, out)`` from a
pair-annotation sidecar (see scripts/extract_pair_annotations.py).
Channels are registered to the image of the chosen side via its frame:
``prefix="latent"`` (default) uses ``latent_frame`` / ``latent_minutiae`` ...;
``prefix="mate"`` uses ``mate_frame`` / ``mate_minutiae`` ... so the mate
conditioning is registered to the *mate* image. ``channels`` may include
``"quality"`` (dense 9.308 grid, latent-only -> zeros for mate), ``"minutiae"``,
``"core"``, ``"delta"``.
"""
frame = sidecar.get(f"{prefix}_frame")
src = (float(frame[0]), float(frame[1])) if frame else None
kw = dict(out_size=out_size, device=device, dtype=dtype)
field = {
"minutiae": f"{prefix}_minutiae",
"core": f"{prefix}_cores",
"delta": f"{prefix}_deltas",
}
maps = []
for ch in channels:
if ch == "quality":
maps.append(render_quality_grid(sidecar.get(f"{prefix}_quality_grid", ""), **kw))
elif ch == "orientation":
# 2-channel dense ridge-flow field, computed from the image itself
# (registered). Requires `image` (B=1,1,H,W) in [0,1].
if image is None:
maps.append(torch.zeros(2, out_size, out_size, device=device, dtype=dtype))
else:
of = compute_orientation_field(image.to(device=device, dtype=dtype))
of = F.interpolate(of, size=(out_size, out_size), mode="bilinear", align_corners=False)
maps.append(of[0])
elif ch in field:
sigma = sigma_minutiae if ch == "minutiae" else sigma_singular
maps.append(render_gaussian_points(
sidecar.get(field[ch], []), src_size=src, sigma=sigma, **kw))
else:
raise ValueError(f"Unknown conditioning channel: {ch}")
return torch.cat(maps, dim=0)
def cond_channel_count(channels: Sequence[str]) -> int:
"""Number of tensor channels produced for a channel list (orientation = 2)."""
return sum(2 if c == "orientation" else 1 for c in channels)
def build_structural_conditioning(
minutiae: Sequence[Point],
cores: Sequence[Point],
deltas: Sequence[Point],
out_size: int,
src_size: tuple[float, float] | None = None,
sigma_minutiae: float = 3.0,
sigma_singular: float = 6.0,
orientation: bool = False,
device: torch.device | str = "cpu",
dtype: torch.dtype = torch.float32,
) -> torch.Tensor:
"""Stack annotation channels into a conditioning tensor ``(C, out, out)``.
Base channels: ``[minutiae, core, delta]``. Singular points use a larger
sigma because they are coarse, low-count topological landmarks. With
``orientation=True`` two extra channels encode minutiae direction as
``sin``/``cos`` weighted heatmaps (theta in degrees from the 3rd column),
giving the critic ridge-flow context. Channel order matches
:func:`structural_channel_names`.
"""
kw = dict(out_size=out_size, src_size=src_size, device=device, dtype=dtype)
channels = [
render_gaussian_points(minutiae, sigma=sigma_minutiae, **kw),
render_gaussian_points(cores, sigma=sigma_singular, **kw),
render_gaussian_points(deltas, sigma=sigma_singular, **kw),
]
if orientation:
import math
sin_w = [math.sin(math.radians(float(p[2]))) if len(p) > 2 else 0.0 for p in minutiae]
cos_w = [math.cos(math.radians(float(p[2]))) if len(p) > 2 else 0.0 for p in minutiae]
channels.append(
render_gaussian_points(minutiae, sigma=sigma_minutiae, weights=sin_w, **kw)
)
channels.append(
render_gaussian_points(minutiae, sigma=sigma_minutiae, weights=cos_w, **kw)
)
return torch.cat(channels, dim=0)