- add RbmParams.from_dict
- added stack construction - added empty implementation of Deep Stack and Rnn Stack
This commit is contained in:
@@ -26,6 +26,27 @@ class RbmParams(Params):
|
||||
self.do_gibbs_sample_hidden = False
|
||||
self.do_batch_sample = False
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, params: dict, version: str = '0'):
|
||||
obj = RbmParams()
|
||||
|
||||
if "0" in version or "1" in version:
|
||||
obj.learning_rate = params["learningRate"]
|
||||
obj.momentum = params["momentum"]
|
||||
obj.weight_decay = params["weightDecay"]
|
||||
obj.num_epochs = params["numEpochs"]
|
||||
obj.num_gibbs_samples = params["numGibbs"]
|
||||
obj.mini_batch_size = params["miniBatchSize"]
|
||||
obj.do_rao_blackwell = params["doRaoBlackwell"]
|
||||
obj.do_gibbs_sample_visible = params["gibbsDoSampleVisible"]
|
||||
obj.do_gibbs_sample_hidden = params["gibbsDoSampleHidden"]
|
||||
obj.do_batch_sample = params["doSampleBatch"]
|
||||
|
||||
if "1" in version:
|
||||
obj.do_gaussian_visible = params["doGaussianVisible"]
|
||||
obj.do_gaussian_hidden = params["doGaussianHidden"]
|
||||
|
||||
return obj
|
||||
|
||||
class LayerParams(Params):
|
||||
def __init__(self, num_visible, num_hidden):
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
from stack import Stack, StackType
|
||||
|
||||
class StackDeep(Stack):
|
||||
def __init__(self, name: str):
|
||||
Stack.__init__(self, StackType.Deep, name)
|
||||
@@ -0,0 +1,74 @@
|
||||
import json
|
||||
from collections.abc import Callable
|
||||
from stack import StackType
|
||||
from layer import Layer
|
||||
from params import RbmParams
|
||||
from stack_deep import StackDeep
|
||||
from stack_rnn import StackRnn
|
||||
|
||||
|
||||
class StackFactory:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, project: dict, layer_constructor: Callable = None):
|
||||
name = project["stack"]["name"]
|
||||
layers = project["stack"]["layers"]
|
||||
|
||||
try:
|
||||
stack_type = StackType[project["stack"]["type_string"]]
|
||||
except KeyError:
|
||||
stack_type = StackType.Deep
|
||||
|
||||
obj = None
|
||||
if stack_type == StackType.Deep:
|
||||
obj = StackDeep(name)
|
||||
if stack_type == StackType.Rnn:
|
||||
obj = StackRnn(name)
|
||||
|
||||
if obj is None:
|
||||
return obj
|
||||
|
||||
for layer in layers:
|
||||
layer_id = layer["id"]
|
||||
layer_name = layer["name"]
|
||||
num_visible_x = layer["numVisibleX"]
|
||||
num_visible_y = layer["numVisibleY"]
|
||||
num_hidden = layer["numHidden"]
|
||||
try:
|
||||
num_context = layer["numContext"]
|
||||
except KeyError:
|
||||
num_context = 0
|
||||
|
||||
# Determine version by existence of keys
|
||||
layer_params = layer["rbm"]["params"]
|
||||
params_version = '0'
|
||||
if "doGaussianHidden" in layer_params and "doGaussianVisible" in layer_params:
|
||||
params_version = '1'
|
||||
|
||||
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)
|
||||
|
||||
# Add layer to stack
|
||||
obj.layer_add(layer_obj)
|
||||
|
||||
return obj
|
||||
|
||||
@classmethod
|
||||
def from_file(cls, filename: str, layer_constructor: Callable = None):
|
||||
obj = None
|
||||
with open(filename, "r") as fp:
|
||||
prj = json.load(fp)
|
||||
obj = StackFactory.from_dict(prj)
|
||||
|
||||
return obj
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
stack = StackFactory.from_file("/home/jens/work/repos/Rbm/many.prj")
|
||||
|
||||
print("Test: [passed]")
|
||||
@@ -0,0 +1,5 @@
|
||||
from stack import Stack, StackType
|
||||
|
||||
class StackRnn(Stack):
|
||||
def __init__(self, name: str):
|
||||
Stack.__init__(self, StackType.Rnn, name)
|
||||
Reference in New Issue
Block a user