| |
| """Single-image FFA-Net axmodel inference.""" |
|
|
| import argparse |
| import sys |
| from pathlib import Path |
|
|
| import numpy as np |
| import axengine as axe |
| from PIL import Image |
|
|
|
|
| FILE = Path(__file__).resolve() |
| NET_DIR = FILE.parent |
| ROOT_DIR = NET_DIR.parent |
| sys.path.insert(0, str(NET_DIR)) |
|
|
|
|
| def parse_args(): |
| parser = argparse.ArgumentParser(description="FFA-Net axmodel single-image inference.") |
| parser.add_argument("--axmodel", default='./FFANet.axmodel', help="axmodel model path.") |
| parser.add_argument("--input", default='outdoor_natural/nh(4).jpg', help="Path to input hazy image.") |
| parser.add_argument("--output", default="axmodel_result.png", help="Output image path.") |
| parser.add_argument("--height", type=int, default=512, help="axmodel input height.") |
| parser.add_argument("--width", type=int, default=512, help="axmodel input width.") |
| return parser.parse_args() |
|
|
| |
| |
|
|
| def preprocess(image_path, height, width): |
| image = Image.open(image_path).convert("RGB") |
| image = image.resize((width, height), Image.BICUBIC) |
| arr = np.asarray(image).astype(np.float32) |
| arr = arr.transpose(2, 0, 1) |
| |
| return arr[None, ...].astype(np.uint8) |
|
|
|
|
| def postprocess(output): |
| arr = np.squeeze(output, axis=0).transpose(1, 2, 0) |
| arr = np.clip(arr, 0.0, 1.0) |
| return Image.fromarray((arr * 255.0 + 0.5).astype(np.uint8)) |
|
|
|
|
| def main(): |
| args = parse_args() |
|
|
| inp = preprocess(args.input, args.height, args.width) |
| session = axe.InferenceSession(args.axmodel, providers=["AxEngineExecutionProvider"]) |
| input_name = session.get_inputs()[0].name |
| out = session.run(None, {input_name: inp})[0] |
|
|
| output_path = Path(args.output) |
| output_path.parent.mkdir(parents=True, exist_ok=True) |
| result = postprocess(out) |
| result.save(str(output_path)) |
| print(f"Saved: {output_path}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|