faceage-onnx / scripts /validate_fp16.py
imbcmdth's picture
README, the DINOv3 license, and the conversion scripts
ae2b7b0 verified
Raw
History Blame Contribute Delete
4.46 kB
"""Compare faceage-dino-fp16.onnx against faceage_dino_fp32.onnx on real face crops.
Crops come from the same WIDER FACE validation images sampled for Task 1, using the
model card's rule: 10 percent proportional padding on each side, then resize to
224x224 bicubic, rescale 1/255, ImageNet mean/std, RGB, NCHW.
"""
import json
import sys
import numpy as np
import onnxruntime as ort
from PIL import Image
FP32 = r"E:/projects/faceage-onnx/source/faceage_dino_fp32.onnx"
FP16 = sys.argv[1] if len(sys.argv) > 1 else r"E:/projects/faceage-onnx/faceage-dino-fp16.onnx"
SAMPLE = r"E:/projects/model-work/wider/sample.json"
N_FACES = 200
MEAN = np.array([0.485, 0.456, 0.406], dtype=np.float32)
STD = np.array([0.229, 0.224, 0.225], dtype=np.float32)
def crop_face(image_rgb, x0, y0, x1, y1, pad=0.10):
h, w = image_rgb.shape[:2]
pw, ph = (x1 - x0) * pad, (y1 - y0) * pad
x0 = max(0, int(x0 - pw)); y0 = max(0, int(y0 - ph))
x1 = min(w, int(x1 + pw)); y1 = min(h, int(y1 + ph))
return image_rgb[y0:y1, x0:x1]
def preprocess(img_rgb):
pil = Image.fromarray(img_rgb).resize((224, 224), Image.BICUBIC)
arr = np.asarray(pil, dtype=np.float32) / 255.0
arr = (arr - MEAN) / STD
return arr.transpose(2, 0, 1)
def main():
sample = json.load(open(SAMPLE))
crops = []
for e in sample:
img = None
for (x, y, w, h) in e["boxes"]:
if np.sqrt(w * h) < 32: # skip tiny boxes; the crop would be mush
continue
if img is None:
img = np.array(Image.open(e["path"]).convert("RGB"))
c = crop_face(img, x, y, x + w, y + h)
if c.shape[0] < 8 or c.shape[1] < 8:
continue
crops.append(preprocess(c))
if len(crops) >= N_FACES:
break
if len(crops) >= N_FACES:
break
batch = np.stack(crops).astype(np.float32)
print("face crops:", batch.shape)
def run(path):
s = ort.InferenceSession(path, providers=["CPUExecutionProvider"])
name = s.get_inputs()[0].name
outs = [o.name for o in s.get_outputs()]
ages, genders = [], []
for i in range(0, len(batch), 16):
a, g = s.run(outs, {name: batch[i:i + 16]})
ages.append(a); genders.append(g)
return np.concatenate(ages), np.concatenate(genders)
a32, g32 = run(FP32)
print("fp32 done")
a16, g16 = run(FP16)
print("fp16 done")
print("age_logits dtype fp32:", a32.dtype, " fp16 model returns:", a16.dtype)
age32 = (1.0 / (1.0 + np.exp(-a32.astype(np.float64)))).sum(axis=1)
age16 = (1.0 / (1.0 + np.exp(-a16.astype(np.float64)))).sum(axis=1)
diff = np.abs(age16 - age32)
print("=== age ===")
print("mean |fp16-fp32| age years:", float(diff.mean()))
print("max |fp16-fp32| age years:", float(diff.max()))
print("p95 |fp16-fp32| age years:", float(np.percentile(diff, 95)))
print("NaN/inf in fp16 age_logits:",
int(np.isnan(a16).sum()), int(np.isinf(a16).sum()))
print("NaN/inf in fp16 gender_logits:",
int(np.isnan(g16).sum()), int(np.isinf(g16).sum()))
print("max |fp16-fp32| age_logits:", float(np.abs(a16.astype(np.float64) - a32.astype(np.float64)).max()))
print("max |fp16-fp32| gender_logits:", float(np.abs(g16.astype(np.float64) - g32.astype(np.float64)).max()))
cls32 = g32.argmax(axis=1)
cls16 = g16.argmax(axis=1)
agree = float((cls32 == cls16).mean())
print("=== gender ===")
print("gender agreement fp16 vs fp32:", agree, f"({int((cls32==cls16).sum())}/{len(cls32)})")
print("fp32 gender split female/male:", int((cls32 == 0).sum()), int((cls32 == 1).sum()))
print("fp16 gender split female/male:", int((cls16 == 0).sum()), int((cls16 == 1).sum()))
print("=== predicted age distribution (fp32 | fp16) ===")
edges = [0, 10, 20, 30, 40, 50, 60, 70, 80, 101]
h32, _ = np.histogram(age32, bins=edges)
h16, _ = np.histogram(age16, bins=edges)
for i in range(len(edges) - 1):
print(f" {edges[i]:>3}-{edges[i+1]-1:<3} {h32[i]:>4} | {h16[i]:>4} {'#' * int(h32[i])}")
print("fp32 age mean/min/max:", round(float(age32.mean()), 3),
round(float(age32.min()), 3), round(float(age32.max()), 3))
print("fp16 age mean/min/max:", round(float(age16.mean()), 3),
round(float(age16.min()), 3), round(float(age16.max()), 3))
if __name__ == "__main__":
main()