"""Run one packaged EfficientNet variant with the documented desktop app policy.""" import argparse import hashlib import json from pathlib import Path import numpy as np from PIL import Image from torchvision import transforms ROOT = Path(__file__).resolve().parent INPUT_SIZE = 224 HIGHLIGHT = 250 def read_json(path): return json.loads(Path(path).read_text(encoding="utf-8")) def digest(path): return hashlib.sha256(Path(path).read_bytes()).hexdigest() def raw_tensor(path, config): transform = transforms.Compose( [ transforms.Resize( int(INPUT_SIZE / config["pretrained_cfg"]["crop_pct"]), interpolation=transforms.InterpolationMode.BICUBIC, ), transforms.CenterCrop(INPUT_SIZE), ] ) with Image.open(path) as image: return np.asarray(transform(image.convert("RGB")), dtype=np.float32)[None] def crop_rgba(image, short_side): width, height = image.size scale = short_side / min(width, height) width, height = int(width * scale), int(height * scale) resized = image.convert("RGBA").resize((width, height), Image.Resampling.BICUBIC) left = int((width - INPUT_SIZE) / 2 + 0.5) top = int((height - INPUT_SIZE) / 2 + 0.5) return np.asarray(resized.crop((left, top, left + INPUT_SIZE, top + INPUT_SIZE))) def exposure_features(rgba): rgb = rgba[:, :, :3].astype(np.uint32) gray = ( 19595 * rgb[:, :, 0] + 38470 * rgb[:, :, 1] + 7471 * rgb[:, :, 2] + 32768 ) // 65536 alpha = rgba[:, :, 3].astype(np.float64) total = float(alpha.sum()) if total == 0: raise ValueError("Cannot measure an entirely transparent image") return { "mean_luminance": float((gray * alpha).sum() / total), "highlight_fraction": float(((gray >= HIGHLIGHT) * alpha).sum() / total), } def exposure_state(features, config): if features["mean_luminance"] < config["minimum_mean_luminance"]: return "too_dark" if ( features["mean_luminance"] > config["maximum_mean_luminance"] and features["highlight_fraction"] > config["maximum_highlight_fraction"] ): return "too_bright" return "acceptable" def image_quality(path, raw, calibration, brightness): gray = np.asarray( Image.fromarray(raw[0].astype(np.uint8)).convert("L"), dtype=np.float32 ) laplacian = ( gray[1:-1, 1:-1] * 4 - gray[2:, 1:-1] - gray[:-2, 1:-1] - gray[1:-1, 2:] - gray[1:-1, :-2] ) edge_variance = float(laplacian.var()) with Image.open(path) as image: exposure = exposure_features( crop_rgba(image, brightness["preprocessing"]["resize_short_side"]) ) state = exposure_state(exposure, brightness) return dict( edge_variance=edge_variance, blur_passed=edge_variance >= calibration["quality"]["minimum_edge_variance"], exposure_state=state, **exposure, ) def probabilities_tflite(folder, raw, calibration, threads): from ai_edge_litert.interpreter import Interpreter path = folder / "model.tflite" if digest(path) != calibration["artifact_sha256"]: raise ValueError("TFLite weights do not match their calibration") interpreter = Interpreter(model_path=str(path), num_threads=threads) interpreter.allocate_tensors() inp = interpreter.get_input_details()[0] out = interpreter.get_output_details()[0] assert inp["shape"].tolist() == [1, INPUT_SIZE, INPUT_SIZE, 3] interpreter.set_tensor(inp["index"], raw) interpreter.invoke() return interpreter.get_tensor(out["index"])[0] def probabilities_pytorch(folder, raw, config, calibration, threads): import timm import torch from safetensors.torch import load_file torch.set_num_threads(threads) path = folder / "model.safetensors" if digest(path) != config["checkpoint_sha256"]: raise ValueError("PyTorch checkpoint does not match its configuration") model = timm.create_model( config["architecture"], pretrained=False, num_classes=len(config["label_names"]) ) model.load_state_dict(load_file(str(path)), strict=True) model.eval() cfg = config["pretrained_cfg"] tensor = torch.from_numpy(raw).permute(0, 3, 1, 2) / 255.0 tensor = (tensor - torch.tensor(cfg["mean"])[None, :, None, None]) / torch.tensor( cfg["std"] )[None, :, None, None] with torch.inference_mode(): return (model(tensor) / calibration["temperature"]).softmax(1)[0].numpy() def predict(image, variant="b2", backend="tflite", threads=4): folder = ROOT if variant == "b2" else ROOT / "models" / variant config = read_json(folder / "config.json") calibration = read_json(folder / "model-config.json")["calibration"] brightness = read_json(ROOT / "brightness-config.json") raw = raw_tensor(image, config) quality = image_quality(image, raw, calibration, brightness) if backend == "tflite": probabilities = probabilities_tflite(folder, raw, calibration, threads) else: probabilities = probabilities_pytorch(folder, raw, config, calibration, threads) assert np.isfinite(probabilities).all() np.testing.assert_allclose(probabilities.sum(), 1.0, atol=1e-4) index = int(probabilities.argmax()) score = float(probabilities[index]) thresholds = calibration["class_thresholds"] + [1.01] confident = calibration["confident_thresholds"] + [1.01] accepted = ( index < 7 and score >= thresholds[index] and quality["blur_passed"] and quality["exposure_state"] == "acceptable" ) state = ( "unclear" if not accepted else "confident" if score >= confident[index] else "possible" ) return dict( model=variant, backend=backend, predicted_label=config["label_names"][index], score=score, accepted=bool(accepted), state=state, installed_in_original_app=config["installed_in_original_app"], quality=quality, probabilities=dict(zip(config["label_names"], map(float, probabilities))), ) def main(): parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("image", type=Path) parser.add_argument("--model", choices=["b0", "b1", "b2", "b3"], default="b2") parser.add_argument("--backend", choices=["tflite", "pytorch"], default="tflite") parser.add_argument("--threads", type=int, default=4) args = parser.parse_args() if args.threads < 1: parser.error("--threads must be positive") print( json.dumps( predict(args.image, args.model, args.backend, args.threads), indent=2 ) ) if __name__ == "__main__": main()