refactored
This commit is contained in:
@@ -37,17 +37,19 @@ if __name__ == "__main__":
|
||||
model.load()
|
||||
|
||||
# Load train data
|
||||
batch = read_armadillo(os.path.join(prj_root, f"{prj_name}.training.dat"))
|
||||
train_batch = read_armadillo(os.path.join(prj_root, f"{prj_name}.training.dat"))
|
||||
|
||||
# Load test data
|
||||
test_batch = read_armadillo(os.path.join(prj_root, f"{prj_name}.test.dat"))
|
||||
|
||||
# Train
|
||||
model.train(batch, TrainingParams(learning_rate=0.00001, momentum=0.9, do_rao_blackwell=True, num_epochs=100, num_gibbs_samples=3))
|
||||
model.train(train_batch, TrainingParams(learning_rate=0.00001, momentum=0.9, do_rao_blackwell=True, num_epochs=100, num_gibbs_samples=3))
|
||||
|
||||
# save state
|
||||
model.save()
|
||||
|
||||
num_patterns = len(batch)
|
||||
fig, axes = plt.subplots(1, num_patterns, figsize=(12, 3))
|
||||
for index, inp in enumerate(batch):
|
||||
fig, axes = plt.subplots(1, len(test_batch), figsize=(12, 3))
|
||||
for index, inp in enumerate(test_batch):
|
||||
out_normalized = model.backward(model.forward(inp))
|
||||
img = 2*(out_normalized + 0.5)
|
||||
img = np.reshape(img, (96, 96))
|
||||
|
||||
Reference in New Issue
Block a user