"""Standalone inference for g_render2 — identity-preserving latent-fingerprint enhancement. Pipeline (matches the paper's `g2fix_clan`): latent --classic_lan--> build 6ch cond --RenderGenerator--> sigmoid --coverage-gate--> classic_lan post-proc Usage: python inference.py [--roi roi_mask.png] [--minu x_y_per_line.txt] [--device cpu|cuda] Notes: - `--roi` : optional foreground/ROI mask PNG (white=print). If omitted, it is derived from ridge coherence (works, but a curated ROI — e.g. an EFS quality grid or a learned segmenter — gives the cleanest result; the coverage-gate blanks everything outside the ROI). - `--minu` : optional minutiae file, one `x y` (pixel coords) per line. If omitted, channel 5 = 0. """ import argparse import numpy as np import torch import torch.nn.functional as F from PIL import Image from g_render.models.render_generator import RenderGenerator from g_render.models.structure.heatmap import compute_orientation_field, render_gaussian_points from g_render.models.frequency.ridge_freq import RidgeFrequencyNormalizer RN = RidgeFrequencyNormalizer() K, TS, GMIN, GMAX = 13, 0.18, 0.5, 4.0 def classic_lan(x): """analytic local adaptive contrast norm (0-param): out = lm + gain*(x-lm).""" p = K // 2 lm = F.avg_pool2d(F.pad(x, (p,) * 4, mode="replicate"), K, 1) ex2 = F.avg_pool2d(F.pad(x * x, (p,) * 4, mode="replicate"), K, 1) lstd = (ex2 - lm * lm).clamp_min(0).add(1e-6).sqrt() return (lm + (TS / lstd).clamp(GMIN, GMAX) * (x - lm)).clamp(0, 1) def robust_of(x): of = compute_orientation_field(x) coh = (of.pow(2).sum(1, keepdim=True) + 1e-12).sqrt().clamp(0, 1) ofu = of / coh.clamp_min(1e-4) num = F.avg_pool2d(F.pad(ofu * coh, (12,) * 4, mode="replicate"), 25, 1) den = F.avg_pool2d(F.pad(coh, (12,) * 4, mode="replicate"), 25, 1).clamp_min(1e-4) ofs = num / den return ofs / (ofs.pow(2).sum(1, keepdim=True) + 1e-12).sqrt().clamp_min(1e-4), coh def build_cond(img_np, size, roi_np=None, minu_pts=None, WH=None): x = torch.from_numpy(img_np)[None, None] # already resized [0,1] ofs, coh = robust_of(x) if roi_np is not None: roi = torch.from_numpy(roi_np)[None, None] else: # derive ROI from coherence roi = F.avg_pool2d(F.pad(coh, (24,) * 4, mode="replicate"), 49, 1) roi = (roi / 0.08).clamp(0, 1) cover = F.avg_pool2d(F.pad(roi, (8,) * 4, mode="replicate"), 17, 1).clamp(0, 1) fmap, _ = RN.estimate_frequency_map(x) fmap = F.interpolate(fmap, size=(size, size), mode="bilinear", align_corners=False).clamp(0, 1) if minu_pts: W, H = WH minu = render_gaussian_points([(mx * size / W, my * size / H) for mx, my in minu_pts], out_size=size, src_size=(float(size), float(size)), sigma=4.0)[None] else: minu = torch.zeros(1, 1, size, size) return torch.cat([x, ofs * cover, cover, fmap * cover, minu], dim=1) @torch.no_grad() def enhance(model, path, roi_path=None, minu=None, device="cpu", size=256, post=True): img = Image.open(path).convert("L"); W, H = img.size xin = classic_lan(torch.from_numpy(np.asarray(img.resize((size, size)), np.float32) / 255.)[None, None]) xin = xin[0, 0].numpy() roi_np = None if roi_path: roi_np = np.asarray(Image.open(roi_path).convert("L").resize((size, size)), np.float32) / 255. cond = build_cond(xin, size, roi_np, minu, (W, H)).to(device) out = model(cond) # sigmoid + coverage-gate inside out = out.cpu() if post: out = classic_lan(out) # analytic post-proc (+identity) y = out[0, 0].numpy() return Image.fromarray(np.clip(y * 255 + 0.5, 0, 255).astype(np.uint8), "L").resize((W, H), Image.BILINEAR) def main(): ap = argparse.ArgumentParser() ap.add_argument("input"); ap.add_argument("output") ap.add_argument("--roi", default=None); ap.add_argument("--minu", default=None) ap.add_argument("--device", default="cpu"); ap.add_argument("--no-post", action="store_true") a = ap.parse_args() model = RenderGenerator(in_ch=6).eval().to(a.device) model.load_state_dict(torch.load("pytorch_model.bin", map_location="cpu"), strict=False) minu = None if a.minu: minu = [tuple(map(float, l.split()[:2])) for l in open(a.minu) if len(l.split()) >= 2] out = enhance(model, a.input, a.roi, minu, a.device, post=not a.no_post) out.save(a.output) print(f"saved -> {a.output}") if __name__ == "__main__": main()