Image Classification
timm
LiteRT
Safetensors
coffee
plant-disease
agriculture
efficientnet
supervised-fine-tuning
Instructions to use ConnorLee08/coffee-leaf-efficientnet-b2 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- timm
How to use ConnorLee08/coffee-leaf-efficientnet-b2 with timm:
import timm model = timm.create_model("hf-hub:ConnorLee08/coffee-leaf-efficientnet-b2", pretrained=True) - Notebooks
- Google Colab
- Kaggle
Download predict.py from ConnorLee08/coffee-leaf-efficientnet-b2: direct link, hf CLI and curl.
- Browser
- Download file 6.82 kB
-
https://huggingface.co/ConnorLee08/coffee-leaf-efficientnet-b2/resolve/main/predict.py
- Command line
-
hf download hf://ConnorLee08/coffee-leaf-efficientnet-b2/predict.py
-
curl -L -o predict.py https://huggingface.co/ConnorLee08/coffee-leaf-efficientnet-b2/resolve/main/predict.py
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() | |