#!/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}%")