refactored
This commit is contained in:
Binary file not shown.
@@ -10,19 +10,14 @@ class Layer:
|
|||||||
self.name = name
|
self.name = name
|
||||||
self.shape = shape
|
self.shape = shape
|
||||||
self.entity = Entity((shape[0]*shape[1]+shape[2], shape[3]), params)
|
self.entity = Entity((shape[0]*shape[1]+shape[2], shape[3]), params)
|
||||||
self.state_filename = f"{self.name}_state.npz"
|
|
||||||
|
|
||||||
def init(self, std: float):
|
def init(self, std: float):
|
||||||
self.entity.state.init(mu=0, std=std)
|
self.entity.state.init(mu=0, std=std)
|
||||||
|
|
||||||
def save(self, filename: str = None):
|
def save(self, filename: str = None):
|
||||||
if filename is None:
|
|
||||||
filename = self.state_filename
|
|
||||||
self.entity.state.to_file(filename)
|
self.entity.state.to_file(filename)
|
||||||
|
|
||||||
def load(self, filename: str = None):
|
def load(self, filename: str = None):
|
||||||
if filename is None:
|
|
||||||
filename = self.state_filename
|
|
||||||
state = RbmState.from_file(filename)
|
state = RbmState.from_file(filename)
|
||||||
if state is not None:
|
if state is not None:
|
||||||
self.entity.state = state
|
self.entity.state = state
|
||||||
|
|||||||
@@ -29,4 +29,5 @@ class Status:
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
def on_report(self, status: dict) -> bool:
|
def on_report(self, status: dict) -> bool:
|
||||||
|
Status.print_status(status)
|
||||||
return True
|
return True
|
||||||
@@ -1,10 +1,10 @@
|
|||||||
import os.path
|
import os.path
|
||||||
import cv2 as cv
|
import cv2 as cv
|
||||||
from argparse import ArgumentParser
|
from argparse import ArgumentParser
|
||||||
from .stack_factory import StackFactory
|
from rbm.stack_factory import StackFactory
|
||||||
from .status import Status
|
from rbm.status import Status
|
||||||
from .stack_deep import StackDeep
|
from rbm.stack_deep import StackDeep
|
||||||
from .matrix import Mat, np, convert
|
from rbm.matrix import Mat, np, convert
|
||||||
|
|
||||||
def cv_show(name: str, vec: Mat, shape):
|
def cv_show(name: str, vec: Mat, shape):
|
||||||
img = cv.Mat(convert(np.resize(vec, shape)))
|
img = cv.Mat(convert(np.resize(vec, shape)))
|
||||||
@@ -1,9 +1,12 @@
|
|||||||
|
import os.path
|
||||||
|
|
||||||
from rbm.params import EntityParams
|
from rbm.params import EntityParams
|
||||||
from rbm.layer import Layer
|
from rbm.layer import Layer
|
||||||
from rbm.status import Status
|
from rbm.status import Status
|
||||||
from rbm.train import train
|
from rbm.train import train
|
||||||
from rbm.matrix import Mat, np
|
from rbm.matrix import Mat, np
|
||||||
|
|
||||||
|
work_dir = "../../results"
|
||||||
def xor():
|
def xor():
|
||||||
# Create params
|
# Create params
|
||||||
params = EntityParams()
|
params = EntityParams()
|
||||||
@@ -17,7 +20,7 @@ def xor():
|
|||||||
layer.init(0.01)
|
layer.init(0.01)
|
||||||
|
|
||||||
# Load weights (if exists)
|
# Load weights (if exists)
|
||||||
layer.load()
|
layer.load(os.path.join(work_dir, "xor_layer0_state.npz"))
|
||||||
|
|
||||||
# Prepare training data
|
# Prepare training data
|
||||||
training_batch = Mat([[0,1,1], [0,0,0], [1,1,0], [1,0,1]], dtype=np.float64)
|
training_batch = Mat([[0,1,1], [0,0,0], [1,1,0], [1,0,1]], dtype=np.float64)
|
||||||
@@ -26,7 +29,7 @@ def xor():
|
|||||||
train(layer.entity, training_batch, Status())
|
train(layer.entity, training_batch, Status())
|
||||||
|
|
||||||
# Save weights
|
# Save weights
|
||||||
layer.save()
|
layer.save(os.path.join(work_dir, "xor_layer0_state.npz"))
|
||||||
|
|
||||||
# Test with test data
|
# Test with test data
|
||||||
test_batch = Mat([[0,0,0], [0,1,0], [1,0,0], [1,1,0]], dtype=np.float64)
|
test_batch = Mat([[0,0,0], [0,1,0], [1,0,0], [1,1,0]], dtype=np.float64)
|
||||||
Reference in New Issue
Block a user