- fixed armadillo data read
- added shape to layer - changed construction of layer from Stack factory - fixed rms_error_accu scaling - RBM: added image display using opencv - improved Status print
This commit is contained in:
+3
-1
@@ -12,7 +12,9 @@ authors = [
|
||||
{ name = "Jens Ahrensfeld" }
|
||||
]
|
||||
dependencies = [
|
||||
"numpy"
|
||||
"numpy",
|
||||
"opencv-python",
|
||||
"opencv-contrib-python"
|
||||
]
|
||||
readme = "README.md"
|
||||
|
||||
|
||||
+1
-1
@@ -21,6 +21,6 @@ def rms_error(d_err: np.ndarray):
|
||||
|
||||
def rms_error_accu(d_err: np.ndarray):
|
||||
d_err_squared = d_err * d_err
|
||||
s = np.sum(d_err_squared) / len(d_err)
|
||||
s = np.sum(d_err_squared) / d_err.size
|
||||
return s
|
||||
|
||||
|
||||
+3
-2
@@ -7,9 +7,10 @@ from status import Status
|
||||
from cd_train import cd_jens
|
||||
|
||||
class Layer:
|
||||
def __init__(self, name: str, dim: tuple[int, int], params: RbmParams):
|
||||
def __init__(self, name: str, shape: tuple[int, int, int, int], params: RbmParams):
|
||||
self.name = name
|
||||
self.state = RbmState.from_layer_params(dim)
|
||||
self.shape = shape
|
||||
self.state = RbmState.from_layer_params((shape[0]*shape[1]+shape[2], shape[3]))
|
||||
self.params = params
|
||||
self.state_filename = f"{self.name}_state.npz"
|
||||
|
||||
|
||||
+17
-26
@@ -1,29 +1,7 @@
|
||||
import os.path
|
||||
|
||||
import numpy as np
|
||||
from layer import Layer
|
||||
from stack import Stack, StackType, StackException
|
||||
import cv2 as cv
|
||||
from stack_factory import StackFactory
|
||||
from params import RbmParams
|
||||
|
||||
def test():
|
||||
params = RbmParams()
|
||||
stack = Stack(StackType.Deep, "Stack")
|
||||
dims = [(2, 3), (3, 4), (4, 5)]
|
||||
|
||||
for n, dim in enumerate(dims):
|
||||
layer = Layer(f"Layer-{n}", dim, params)
|
||||
stack.append(layer)
|
||||
|
||||
stack.init(std=0.1)
|
||||
stack.state_save()
|
||||
|
||||
lay0 = stack.layers[0]
|
||||
lay1 = stack.from_name("Layer-1")
|
||||
lay2 = stack.from_index(2)
|
||||
stack.remove(lay2)
|
||||
stack.remove(lay1)
|
||||
stack.remove(lay0)
|
||||
|
||||
def read_armadillo(filename: str) -> np.ndarray:
|
||||
result = None
|
||||
@@ -37,14 +15,18 @@ def read_armadillo(filename: str) -> np.ndarray:
|
||||
print(f"shape: {shape}")
|
||||
|
||||
result = np.zeros(shape=shape, dtype=np.float64)
|
||||
line = fp.readline().replace("\n", '').split(' ')
|
||||
line = line[1:]
|
||||
data = [float(s) for s in line]
|
||||
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 cv_show(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"training samples", img_n)
|
||||
|
||||
def main(prj_name: str = "test"):
|
||||
work_dir = "../../results"
|
||||
@@ -63,12 +45,21 @@ def main(prj_name: str = "test"):
|
||||
training_data_path = os.path.join(prj_root, f"{prj_name}.training.dat")
|
||||
batch = read_armadillo(training_data_path)
|
||||
|
||||
# Shape of training vector
|
||||
shape = stack.from_index(0).shape[0:2] + (1,)
|
||||
|
||||
# Train
|
||||
stack.train(batch)
|
||||
|
||||
# Save state
|
||||
stack.state_save()
|
||||
|
||||
for n in range(batch.shape[0]):
|
||||
cv_show(batch[n,:], shape)
|
||||
cv.waitKeyEx(200)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main("norb_small_16h")
|
||||
cv.destroyAllWindows()
|
||||
|
||||
print("Test: [passed]")
|
||||
|
||||
@@ -48,7 +48,7 @@ class StackFactory:
|
||||
params = RbmParams.from_dict(layer_params, params_version)
|
||||
|
||||
# Create layer
|
||||
layer_obj = Layer(f"{layer_name}-{layer_id}", (num_visible_x*num_visible_y+num_context, num_hidden), params)
|
||||
layer_obj = Layer(f"{layer_name}-{layer_id}", (num_visible_x, num_visible_y, num_context, num_hidden), params)
|
||||
|
||||
# Add layer to stack
|
||||
obj.append(layer_obj)
|
||||
|
||||
+4
-1
@@ -14,5 +14,8 @@ class Status:
|
||||
for key in status.keys():
|
||||
value = status[key]['value']
|
||||
unit = status[key]['unit']
|
||||
print(f"{key} : {value}{unit}")
|
||||
if isinstance(value, float):
|
||||
print(f"{key} : {value:0.6f}{unit}")
|
||||
else:
|
||||
print(f"{key} : {value}{unit}")
|
||||
return True
|
||||
|
||||
Reference in New Issue
Block a user