tonigi's picture
Add complete Gradio interface
3c2cd23 verified
Raw History Blame Contribute Delete
7.78 kB
from __future__ import annotations
import argparse
import importlib.util
import json
import sys
from pathlib import Path
from .image_io import InvalidImage
from .inpaint import InpaintUnavailable, install_lama, lama_installed, resolve_torch_device
from .models import BitmapExtractor, CleanupMode, DEFAULT_SNAP_ANGLES, Device, ProcessOptions
from .ocr import OcrUnavailable, is_ocr_available
from .pipeline import export_result, process_path
from .sam2 import (
ShapeExtractionUnavailable,
install_sam2,
sam2_dependencies_installed,
sam2_installed,
)
IMAGE_SUFFIXES = {".png", ".jpg", ".jpeg", ".webp"}
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(prog="editable-image")
subparsers = parser.add_subparsers(dest="command", required=True)
convert = subparsers.add_parser("convert", help="convert bitmap images")
convert.add_argument("inputs", nargs="+", type=Path)
convert.add_argument("--output", "-o", type=Path, required=True)
convert.add_argument("--recursive", action="store_true")
convert.add_argument("--cleanup", choices=[item.value for item in CleanupMode], default="opencv")
convert.add_argument("--device", choices=[item.value for item in Device], default="auto")
convert.add_argument("--confidence", type=float, default=0.5)
convert.add_argument(
"--snap-angles",
type=parse_angles,
default=list(DEFAULT_SNAP_ANGLES),
metavar="DEGREES",
help="comma-separated target angles (default: 0,45,90,-45,-90)",
)
convert.add_argument(
"--snap-tolerance",
type=float,
default=6.0,
metavar="DEGREES",
help="maximum distance from a target angle (default: 6)",
)
convert.add_argument(
"--snap-font-sizes",
action="store_true",
help="snap estimated font sizes to peaks in the image's size distribution",
)
convert.add_argument(
"--embed-fonts",
action="store_true",
help="embed used font faces in the SVG (default: link system fonts)",
)
convert.add_argument(
"--vectorize-shapes",
action="store_true",
help="promote confident diagram primitives into editable SVG shapes",
)
convert.add_argument(
"--bitmap-extractor",
choices=[item.value for item in BitmapExtractor],
default=BitmapExtractor.OPENCV.value,
help="bitmap-layer mask extractor used with --vectorize-shapes (default: opencv)",
)
convert.add_argument("--overwrite", action="store_true")
models = subparsers.add_parser("models", help="manage optional model files")
model_commands = models.add_subparsers(dest="model_command", required=True)
install = model_commands.add_parser("install")
install.add_argument("model", choices=["ocr", "lama", "sam2"])
model_commands.add_parser("status")
return parser
def main(argv: list[str] | None = None) -> int:
args = build_parser().parse_args(argv)
if args.command == "models":
return model_command(args)
return convert_command(args)
def model_command(args: argparse.Namespace) -> int:
if args.model_command == "install":
if args.model == "lama":
path = install_lama()
print(f"Installed LaMa at {path}")
return 0
if args.model == "sam2":
path = install_sam2()
print(f"Installed SAM 2 at {path}")
return 0
if not is_ocr_available():
print("Install the cpu or cuda project extra before installing OCR models", file=sys.stderr)
return 2
from .ocr import RapidOcrEngine
RapidOcrEngine(Device.CPU)
print("OCR models are ready")
return 0
print(
json.dumps(
{
"ocr": is_ocr_available(),
"lama": lama_installed(),
"sam2": sam2_installed(),
"sam2_dependencies": sam2_dependencies_installed(),
"devices": available_devices(),
},
indent=2,
)
)
return 0
def convert_command(args: argparse.Namespace) -> int:
if args.bitmap_extractor != BitmapExtractor.OPENCV.value and not args.vectorize_shapes:
print("--bitmap-extractor requires --vectorize-shapes", file=sys.stderr)
return 2
options = ProcessOptions(
confidence=args.confidence,
cleanup=CleanupMode(args.cleanup),
device=Device(args.device),
snap_angles=args.snap_angles,
snap_tolerance=args.snap_tolerance,
snap_font_sizes=args.snap_font_sizes,
vectorize_shapes=args.vectorize_shapes,
bitmap_extractor=BitmapExtractor(args.bitmap_extractor),
)
inputs = collect_inputs(args.inputs, args.recursive)
if not inputs:
print("No supported images found", file=sys.stderr)
return 2
failures = 0
for source in inputs:
destination = args.output / source.stem
expected = [destination / f"{source.stem}-background.png", destination / f"{source.stem}-overlay.svg"]
if not args.overwrite and any(path.exists() for path in expected):
failures += 1
print(f"skip {source}: output exists (use --overwrite)", file=sys.stderr)
continue
try:
result = process_path(source, options)
background, overlay, assets = export_result(
result,
destination,
source.stem,
embed_fonts=args.embed_fonts,
)
print(
json.dumps(
{
"source": str(source),
"background": str(background),
"overlay": str(overlay),
"assets": [str(path) for path in assets],
}
)
)
except (
InvalidImage,
OcrUnavailable,
InpaintUnavailable,
ShapeExtractionUnavailable,
OSError,
ValueError,
) as exc:
failures += 1
print(f"failed {source}: {exc}", file=sys.stderr)
return 1 if failures else 0
def collect_inputs(paths: list[Path], recursive: bool) -> list[Path]:
output: list[Path] = []
for path in paths:
if path.is_file() and path.suffix.lower() in IMAGE_SUFFIXES:
output.append(path)
elif path.is_dir():
iterator = path.rglob("*") if recursive else path.glob("*")
output.extend(item for item in iterator if item.is_file() and item.suffix.lower() in IMAGE_SUFFIXES)
return sorted(set(output))
def parse_angles(value: str) -> list[float]:
try:
angles = [float(item.strip()) for item in value.split(",") if item.strip()]
except ValueError as exc:
raise argparse.ArgumentTypeError("angles must be comma-separated numbers") from exc
if not angles:
raise argparse.ArgumentTypeError("at least one snap angle is required")
if any(angle < -180 or angle > 180 for angle in angles):
raise argparse.ArgumentTypeError("snap angles must be between -180 and 180")
return angles
def available_devices() -> list[str]:
devices = ["cpu"]
try:
import onnxruntime as ort
if "CUDAExecutionProvider" in ort.get_available_providers():
devices.append("cuda")
except ImportError:
pass
if importlib.util.find_spec("torch") is not None:
try:
if resolve_torch_device(Device.AUTO) == "mps":
devices.append("mps")
except InpaintUnavailable:
pass
return devices
if __name__ == "__main__":
raise SystemExit(main())