#!/usr/bin/env python3 """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() # MEAN = np.array([0.64, 0.6, 0.58], dtype=np.float32).reshape(3, 1, 1) # STD = np.array([0.14, 0.15, 0.152], dtype=np.float32).reshape(3, 1, 1) 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) # arr = (arr - MEAN) / STD # 训练同款归一化 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()