Image Classification
timm
LiteRT
Safetensors
coffee
plant-disease
agriculture
efficientnet
supervised-fine-tuning
ConnorLee08's picture
Add B2 coffee leaf model and B0–B3 comparison package
31a46d7 verified
Raw History Blame Contribute Delete
6.82 kB
"""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()