"""Command-line interface for local RMBG 2.0 inference."""
import argparse
from collections.abc import Sequence
from pathlib import Path
import sys
from rmbg_mask.inference import RmbgError, generateMask
DEFAULT_PROCESSING_RESOLUTION = 1024
def positiveInt(value: str) -> int:
"""Parse a positive integer for an argparse option."""
try:
parsed_value = int(value)
except ValueError as error:
raise argparse.ArgumentTypeError(
f"expected a positive integer, got {value!r}"
) from error
if parsed_value <= 0:
raise argparse.ArgumentTypeError(
f"expected a positive integer, got {value!r}"
)
return parsed_value
def buildParser() -> argparse.ArgumentParser:
"""Build and return the command-line argument parser."""
parser = argparse.ArgumentParser(
prog="rmbg-mask",
description=(
"Generate an original-size grayscale mask using a local "
"RMBG 2.0 model."
),
)
parser.add_argument(
"input",
type=Path,
help="input image path",
)
parser.add_argument(
"output",
type=Path,
help="output grayscale mask path (PNG recommended)",
)
parser.add_argument(
"--model",
type=Path,
required=True,
help="directory containing the already-downloaded model",
)
parser.add_argument(
"-r",
"--processing-resolution",
type=positiveInt,
default=DEFAULT_PROCESSING_RESOLUTION,
metavar="PIXELS",
help=(
"square inference resolution "
f"(default: {DEFAULT_PROCESSING_RESOLUTION})"
),
)
parser.add_argument(
"--device",
choices=("auto", "cpu", "cuda", "mps"),
default="auto",
help="inference device (default: auto)",
)
return parser
def main(argv: Sequence[str] | None = None) -> int:
"""Run the command-line program and return its process exit status."""
arguments = buildParser().parse_args(argv)
try:
generateMask(
model_path=arguments.model,
input_path=arguments.input,
output_path=arguments.output,
processing_resolution=arguments.processing_resolution,
device_name=arguments.device,
)
except RmbgError as error:
print(f"rmbg-mask: error: {error}", file=sys.stderr)
return 1
return 0