Files
pyRBM/src/rbm/rbm.py
T

106 lines
2.7 KiB
Python

import os.path
import numpy as np
import cv2 as cv
from argparse import ArgumentParser
from stack_factory import StackFactory
from status import Status
from stack_deep import StackDeep
def cv_show(name: str, vec: np.array, shape):
img = cv.Mat(np.resize(vec, shape))
img_n = cv.normalize(src=img, dst=None, alpha=255, beta=0, norm_type=cv.NORM_MINMAX, dtype=cv.CV_8U)
cv.imshow(f"{name}", img_n)
class MyStatus(Status):
def __init__(self, _stack: StackDeep, _batch: np.ndarray):
Status.__init__(self, update_interval=2)
self.stack = _stack
self.batch = _batch
self.index = 0
def on_report(self, status: dict) -> bool:
do_continue = True
# print status values
Status.print_status(status)
# Shape of training vector
shape = self.stack.from_index(0).shape[0:2] + (1,)
# User input
max_index = self.batch.shape[0] - 1
key = cv.waitKeyEx(1)
if key == ord('q'):
do_continue = False
if key == ord('-'):
self.index = max(0, self.index-1)
if key == ord('+'):
self.index = min(max_index, self.index+1)
# Show reconstruction
img = self.stack.pass_down_up(self.batch[self.index,:])
cv_show("img", img, shape)
return do_continue
def read_armadillo(filename: str) -> np.ndarray:
result = None
with open(filename) as fp:
identifier = fp.readline().replace("\n", '')
if "ARMA_MAT_TXT_FN008" not in identifier:
raise Exception("Not a armadillo data file!")
line = fp.readline().replace("\n", '').split(' ')
shape = [int(s) for s in line]
print(f"shape: {shape}")
result = np.zeros(shape=shape, dtype=np.float64)
for row in range(shape[0]):
line = fp.readline().replace("\n", '').split(' ')
line = line[1:]
data = [float(s) for s in line]
result[row, :] = data
return result
def main(prj_name: str = "test"):
work_dir = "../../results"
prj_root = "/home/jens/work/repos/Rbm"
prj_path = os.path.join(prj_root, f"{prj_name}.prj")
# Create stack from project file
stack = StackFactory.from_file(prj_path, work_dir=work_dir)
# Init state
stack.state_init(0.01)
# Load state
stack.state_load()
# Load train data
training_data = read_armadillo(os.path.join(prj_root, f"{prj_name}.training.dat"))
try:
test_data = read_armadillo(os.path.join(prj_root, f"{prj_name}.test.dat"))
except FileNotFoundError:
test_data = training_data
# Prepare status listener
my_status = MyStatus(stack, test_data)
# Train
stack.train(training_data, status=my_status)
# Save state
stack.state_save()
if __name__ == "__main__":
ap = ArgumentParser()
ap.add_argument("name", type=str,
default='default',
help="Name of project")
args = ap.parse_args()
var_args = vars(args)
main(var_args["name"])
cv.destroyAllWindows()
print("Test: [passed]")