From 453c0636397019b7b59db0ee6412f92bd88fe82c Mon Sep 17 00:00:00 2001 From: wang-xinyu Date: Thu, 28 May 2020 17:46:29 +0800 Subject: [PATCH] add arcface/gen_wts --- arcface/gen_wts.py | 37 +++++++++++++++++++++++++++++++++++++ 1 file changed, 37 insertions(+) create mode 100644 arcface/gen_wts.py diff --git a/arcface/gen_wts.py b/arcface/gen_wts.py new file mode 100644 index 0000000..852cde0 --- /dev/null +++ b/arcface/gen_wts.py @@ -0,0 +1,37 @@ +import struct +import sys +import argparse +import face_model +import cv2 +import numpy as np + +parser = argparse.ArgumentParser(description='face model test') +# general +parser.add_argument('--image-size', default='112,112', help='') +parser.add_argument('--model', default='model-r50-am-lfw/model,1', help='path to load model.') +parser.add_argument('--ga-model', default='', help='path to load model.') +parser.add_argument('--gpu', default=0, type=int, help='gpu id') +parser.add_argument('--det', default=0, type=int, help='mtcnn option, 1 means using R+O, 0 means detect from begining') +parser.add_argument('--flip', default=0, type=int, help='whether do lr flip aug') +parser.add_argument('--threshold', default=1.24, type=float, help='ver dist threshold') +args = parser.parse_args() + +model = face_model.FaceModel(args) + +f = open('arcface-r50.wts', 'w') +f.write('{}\n'.format(len(model.model.get_params()[0].keys()) + len(model.model.get_params()[1].keys()))) +for k, v in model.model.get_params()[0].items(): + vr = v.reshape(-1).asnumpy() + f.write('{} {} '.format(k, len(vr))) + for vv in vr: + f.write(' ') + f.write(struct.pack('>f',float(vv)).hex()) + f.write('\n') +for k, v in model.model.get_params()[1].items(): + vr = v.reshape(-1).asnumpy() + f.write('{} {} '.format(k, len(vr))) + for vv in vr: + f.write(' ') + f.write(struct.pack('>f',float(vv)).hex()) + f.write('\n') +