deepface-emotion-imx / evaluate.py
NXP
Release 2.0
b2e5bfb
Raw History Blame Contribute Delete
2.45 kB
#!/usr/bin/env python3
# Copyright 2022-2024,2026 NXP
# SPDX-License-Identifier: MIT
#
# Evaluate a Deepface emotion TFLite model on the FER2013 PrivateTest split.
#
# Usage:
# python evaluate.py \
# --model original_model/emotion_uint8_float32.tflite \
# --fer2013_csv PATH/TO/fer2013/fer2013.csv
import argparse
import sys
import numpy as np
import tensorflow as tf
LABELS = ['angry', 'disgust', 'fear', 'happy', 'sad', 'surprise', 'neutral']
parser = argparse.ArgumentParser(
description="Evaluate Deepface emotion TFLite model on FER2013 PrivateTest split."
)
parser.add_argument("-m", "--model",
default="original_model/emotion_uint8_float32.tflite",
type=str,
help="Path to the TFLite model file.")
parser.add_argument("--fer2013_csv", default="fer2013/fer2013.csv", type=str,
help="Path to the FER2013 CSV file.")
args = parser.parse_args()
interpreter = tf.lite.Interpreter(args.model)
interpreter.allocate_tensors()
input_details = interpreter.get_input_details()
output_details = interpreter.get_output_details()
input_scale, input_zero_point = input_details[0]["quantization"]
print("Loaded model:", args.model)
print(f"Reading FER2013 dataset from {args.fer2013_csv} ...")
with open(args.fer2013_csv, 'r') as f:
lines = f.readlines()
lines = [li.strip().split(',') for li in lines[1:]]
lines = [li for li in lines if li[2] == "PrivateTest"]
lines = [(int(li[0]), np.array(li[1].split(' ')).reshape((48, 48, 1)).astype(np.float32))
for li in lines]
print(f" {len(lines)} PrivateTest images found.")
n_correct = 0
total = len(lines)
for i, (label, pixels) in enumerate(lines):
# Normalize to [0, 1] and add batch dimension.
im = pixels[None, ...] / 255.0
# Quantize to uint8 using model quantization params.
im_q = (im / input_scale + input_zero_point).astype(np.uint8)
interpreter.set_tensor(input_details[0]['index'], im_q)
interpreter.invoke()
out = interpreter.get_tensor(output_details[0]['index'])
if int(out.argmax()) == label:
n_correct += 1
if (i + 1) % 500 == 0:
print(f" [{i + 1}/{total}] accuracy so far: {n_correct / (i + 1) * 100:.2f}%")
sys.stdout.flush()
accuracy = n_correct / total * 100 if total > 0 else 0.0
print(f"\nResults for {args.model}")
print(f" Images evaluated : {total}")
print(f" Accuracy : {accuracy:.2f}%")