From 91640db661fdcc04960d1f048ac60c2c4b2c5319 Mon Sep 17 00:00:00 2001 From: tslil Date: Tue, 12 Jan 2021 17:01:29 -0500 Subject: Auto-generating weights --- src/train5.py | 29 +++++++++++++++++++++-------- 1 file changed, 21 insertions(+), 8 deletions(-) (limited to 'src/train5.py') diff --git a/src/train5.py b/src/train5.py index adb81a5..71a8652 100644 --- a/src/train5.py +++ b/src/train5.py @@ -119,20 +119,33 @@ def write_weights(model): string = re.sub(r'([0-9]+)\n', r'\1,\n', string) string = re.sub(r'([0-9]+) ', r'\1, ', string) return string - f = open("weights.txt", "w") + # Prepare everything in a sane memory order conv2d_weights = np.array(transpose(model.trainable_variables[0], perm=[3, 1, 0, 2])) conv2d_biases = np.array(model.trainable_variables[1]) dense1_weights = np.array(transpose(model.trainable_variables[2], perm=[1, 0])) dense1_biases = np.array(model.trainable_variables[3]) dense2_weights = np.array(transpose(model.trainable_variables[4], perm=[1, 0])) dense2_biases = np.array(model.trainable_variables[5]) - output_weights = np.array(transpose(model.trainable_variables[6], perm=[1, 0])) - output_biases = np.array(model.trainable_variables[7]) - for v in [conv2d_weights, conv2d_biases, - dense1_weights, dense1_biases, - dense2_weights, dense2_biases, - output_weights, output_biases]: - f.write(fix(str(v))+"\n\n") + output_weights = np.array(transpose(model.trainable_variables[6], perm=[1, 0])[0]) + output_bias = np.array(model.trainable_variables[7][0]) + # Prepare formatting + names = ["conv2d_weights[KERN_NUM][KERN_SIZE][KERN_SIZE][KERN_CHAN]", + "conv2d_biases[KERN_NUM]", + "dense1_weights[DENSE1_NUM][CONV_NUM+2]", + "dense1_biases[DENSE1_NUM]", + "dense2_weights[DENSE2_NUM][DENSE1_NUM]", + "dense2_biases[DENSE2_NUM]", + "output_weights[DENSE2_NUM]", + "output_bias"] + variables = [conv2d_weights, conv2d_biases, + dense1_weights, dense1_biases, + dense2_weights, dense2_biases, + output_weights, output_bias] + # Write to file + f = open("weights.c", "w") + f.write("#include \"weights.h\"\n\n") + for (name, val) in zip(names, variables): + f.write("const float "+name+" =\n"+fix(str(val))+";\n\n") f.close() -- cgit v1.3.1