refactored
This commit is contained in:
@@ -37,17 +37,19 @@ if __name__ == "__main__":
|
|||||||
model.load()
|
model.load()
|
||||||
|
|
||||||
# Load train data
|
# 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
|
# 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
|
# save state
|
||||||
model.save()
|
model.save()
|
||||||
|
|
||||||
num_patterns = len(batch)
|
fig, axes = plt.subplots(1, len(test_batch), figsize=(12, 3))
|
||||||
fig, axes = plt.subplots(1, num_patterns, figsize=(12, 3))
|
for index, inp in enumerate(test_batch):
|
||||||
for index, inp in enumerate(batch):
|
|
||||||
out_normalized = model.backward(model.forward(inp))
|
out_normalized = model.backward(model.forward(inp))
|
||||||
img = 2*(out_normalized + 0.5)
|
img = 2*(out_normalized + 0.5)
|
||||||
img = np.reshape(img, (96, 96))
|
img = np.reshape(img, (96, 96))
|
||||||
|
|||||||
Reference in New Issue
Block a user