File size: 5,199 Bytes
7615398 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 | """Remove examiner red ink from exam-script images, self-contained.
Two stages: a U-Net segmenter locates the red, a fine-tuned DeepEraser removes it and
reconstructs the black underneath. All weights and model code are bundled and referenced
by paths relative to this file, so it runs from anywhere.
python code/remove_red.py <image_or_dir> [image2 ...] [--out OUTDIR]
Examples:
python code/remove_red.py page.jpg
python code/remove_red.py scans/ --out cleaned/
"""
import os, sys, glob, argparse, numpy as np, torch
from PIL import Image
HERE = os.path.dirname(os.path.abspath(__file__))
BUNDLE = os.path.dirname(HERE)
MODELS = os.path.join(BUNDLE, "models")
sys.path.insert(0, HERE)
sys.path.insert(0, os.path.join(BUNDLE, "deeperaser"))
from train_redseg import UNet as SegUNet
from model import DeepEraser
TILE, OV = 768, 96
BIG = 1400
def _load_weights(dev, path):
ext = os.path.splitext(path)[1]
if ext == ".safetensors":
from safetensors.torch import load_file
return load_file(path, device=dev)
sd = torch.load(path, map_location=dev, weights_only=True)
if any(k.startswith("module.") for k in sd):
sd = {k[7:]: v for k, v in sd.items()}
return sd
def _find_weights(name):
"""Try .safetensors first, then fall back to .pt/.pth."""
for ext in (".safetensors", ".pt", ".pth"):
p = os.path.join(MODELS, name + ext)
if os.path.exists(p):
return p
raise FileNotFoundError(f"no weights found for {name} in {MODELS}")
def load_models(dev):
seg = SegUNet().to(dev)
seg.load_state_dict(_load_weights(dev, _find_weights("redseg_best"))); seg.eval()
de = DeepEraser().to(dev)
sd = _load_weights(dev, _find_weights("deeperaser_ft_real"))
de.load_state_dict({k: v for k, v in sd.items() if k in de.state_dict()}); de.eval()
return seg, de
def red_mask(seg, a, dev):
H, W = a.shape[:2]; ph, pw = (16 - H % 16) % 16, (16 - W % 16) % 16
x = np.pad(a, ((0, ph), (0, pw), (0, 0)), mode="reflect")
t = torch.from_numpy(x.transpose(2, 0, 1) / 255.).float()[None].to(dev)
with torch.no_grad():
p = torch.sigmoid(seg(t))[0, 0].cpu().numpy()
return p[:H, :W] > 0.5
def _erase_full(de, a, m, dev):
H, W = a.shape[:2]; ph, pw = (16 - H % 16) % 16, (16 - W % 16) % 16
img = np.pad(a, ((0, ph), (0, pw), (0, 0)), mode="reflect")
mk = np.pad(m.astype(np.float32), ((0, ph), (0, pw)))
im = torch.from_numpy(img / 255.).permute(2, 0, 1).float()[None].to(dev)
mask = torch.from_numpy(mk)[None, None].float().to(dev)
with torch.no_grad():
o = torch.clamp(de(im, mask)[-1], 0, 1)[0].permute(1, 2, 0).cpu().numpy()[:H, :W]
return (o * 255).astype(np.uint8)
def _erase_tiled(de, a, m, dev):
H, W = a.shape[:2]; acc = np.zeros((H, W, 3), np.float32); wsum = np.zeros((H, W, 1), np.float32)
step = TILE - OV
ys = list(range(0, max(1, H - OV), step)); xs = list(range(0, max(1, W - OV), step))
if ys[-1] + TILE < H: ys.append(H - TILE)
if xs[-1] + TILE < W: xs.append(W - TILE)
def ramp(n):
w = np.ones(n, np.float32); r = min(OV, n // 2)
if r > 0: w[:r] = np.linspace(0.05, 1, r); w[-r:] = np.linspace(1, 0.05, r)
return w
for y in ys:
for x in xs:
y0, x0 = max(0, min(y, H - TILE)), max(0, min(x, W - TILE))
y1, x1 = min(H, y0 + TILE), min(W, x0 + TILE)
th, tw = y1 - y0, x1 - x0
o = _erase_full(de, a[y0:y1, x0:x1], m[y0:y1, x0:x1], dev)
win = (ramp(th)[:, None] * ramp(tw)[None, :])[..., None]
acc[y0:y1, x0:x1] += o * win; wsum[y0:y1, x0:x1] += win
wsum[wsum == 0] = 1
return (acc / wsum).astype(np.uint8)
def clean_image(seg, de, path, dev):
a = np.asarray(Image.open(path).convert("RGB"))
m = red_mask(seg, a, dev)
if m.mean() < 0.0005: # no meaningful red -> return as-is
return a, m.mean()
out = _erase_tiled(de, a, m, dev) if max(a.shape[:2]) > BIG else _erase_full(de, a, m, dev)
return out, m.mean()
def main():
ap = argparse.ArgumentParser()
ap.add_argument("inputs", nargs="+", help="image file(s) or directory")
ap.add_argument("--out", default="cleaned", help="output directory")
ap.add_argument("--cpu", action="store_true", help="force CPU")
args = ap.parse_args()
dev = "cpu" if args.cpu or not torch.cuda.is_available() else "cuda"
files = []
for p in args.inputs:
if os.path.isdir(p):
files += [f for f in sorted(glob.glob(os.path.join(p, "*")))
if f.lower().endswith((".jpg", ".jpeg", ".png"))]
else:
files.append(p)
os.makedirs(args.out, exist_ok=True)
print(f"device={dev} | {len(files)} image(s) -> {args.out}", flush=True)
seg, de = load_models(dev)
for p in files:
base = os.path.splitext(os.path.basename(p))[0]
out, red = clean_image(seg, de, p, dev)
outp = os.path.join(args.out, base + ".jpg")
Image.fromarray(out).save(outp, quality=95)
print(f" {base}: red={red*100:.2f}% -> {outp}", flush=True)
print("done", flush=True)
if __name__ == "__main__":
main()
|