aboutsummaryrefslogtreecommitdiff
path: root/src/nn_train.py
diff options
context:
space:
mode:
Diffstat (limited to 'src/nn_train.py')
-rw-r--r--src/nn_train.py31
1 files changed, 16 insertions, 15 deletions
diff --git a/src/nn_train.py b/src/nn_train.py
index f2550b4..03852f5 100644
--- a/src/nn_train.py
+++ b/src/nn_train.py
@@ -56,7 +56,7 @@ def train(size, model, data, iterations=1, epochs=10, batch=None):
tra_res = model.evaluate(tra_input, tra_outcome, verbose=False)
results.append((tra_res, val_res))
print(val_res)
- write_weights(model, str(i+1).zfill(len(str(iterations))), (tra_res, val_res))
+ write_weights(size, model, str(i+1).zfill(len(str(iterations))), (tra_res, val_res))
print("\nScores")
for i, data in enumerate(results):
print(f"Iteration {i+1}: {data}")
@@ -74,33 +74,34 @@ def make_model(size, magic):
return model
-def write_weights(model, iteration, performance):
+def write_weights(size, model, iteration, performance):
def fix(val):
string = str(np.array(val).tolist())
string = string.replace("[", "{").replace("]", "}")
return string
- dense1_weights = transpose(model.trainable_variables[0], perm=[1, 0])
- dense1_biases = model.trainable_variables[1]
+ size_str = "five" if size==5 else "six"
+ dense_weights = transpose(model.trainable_variables[0], perm=[1, 0])
+ dense_biases = model.trainable_variables[1]
output_weights = transpose(model.trainable_variables[2], perm=[1, 0])
output_bias = model.trainable_variables[3]
# Prepare output
- to_output = [("dense1_weights[DENSE_NUM][INP_NUM]", dense1_weights),
- ("dense1_biases[DENSE_NUM]", dense1_biases),
- ("output_weights[2][DENSE_NUM]", output_weights),
- ("output_bias[2]", output_bias)]
+ to_output = [(f"{size_str}_dense_weights[{size_str.upper()}_DENSE_NUM][{size_str.upper()}_INP_NUM]", dense_weights),
+ (f"{size_str}_dense_biases[{size_str.upper()}_DENSE_NUM]", dense_biases),
+ (f"{size_str}_output_weights[2][{size_str.upper()}_DENSE_NUM]", output_weights),
+ (f"{size_str}_output_bias[2]", output_bias)]
# Write to file
- f = open("weights-"+iteration+".txt", "w")
+ f = open(f"{size_str}_weights-"+iteration+".txt", "w")
f.write("/*\n")
model.summary(print_fn=lambda l: f.write(" * "+l+"\n"))
f.write(" * "+str(performance)+"\n*/\n\n")
- f.write("#include \"weights.h\"\n\n")
+ f.write(f"#include \"weights_{size}.h\"\n\n")
for (name, val) in to_output:
f.write("const float "+name+" =\n"+fix(val)+";\n\n")
f.close()
-data = load_data(5)
-model = make_model(5, 64)
-
-print("Before training", model.evaluate(data[1][0], data[1][1], verbose=False, batch_size=16))
-results = train(5, model, data, iterations=5, epochs=10, batch=128)
+for size in (5,6):
+ data = load_data(size)
+ model = make_model(size, 64)
+ print("Before training", model.evaluate(data[1][0], data[1][1], verbose=False, batch_size=16))
+ results = train(size, model, data, iterations=5, epochs=10, batch=128)