| """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 <input.png> <output.png> [--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] |
| ofs, coh = robust_of(x) |
| if roi_np is not None: |
| roi = torch.from_numpy(roi_np)[None, None] |
| else: |
| 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) |
| out = out.cpu() |
| if post: |
| out = classic_lan(out) |
| 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() |
|
|