diff options
author | Determinant <[email protected]> | 2015-06-02 20:28:16 +0800 |
---|---|---|
committer | Determinant <[email protected]> | 2015-06-02 20:28:16 +0800 |
commit | 74d9e9e7371c80394698fb9805cbf0cbde67a8f3 (patch) | |
tree | 36b070f1fcfa2be8fc80c50b7a221862a0dfd14a /nn/layer_repo.lua | |
parent | 60083f2e51935ce55cec7a4c39d1724a16d9c769 (diff) |
add ParamRepo, LayerRepo, DAGLayer
Diffstat (limited to 'nn/layer_repo.lua')
-rw-r--r-- | nn/layer_repo.lua | 34 |
1 files changed, 34 insertions, 0 deletions
diff --git a/nn/layer_repo.lua b/nn/layer_repo.lua new file mode 100644 index 0000000..b1d2248 --- /dev/null +++ b/nn/layer_repo.lua @@ -0,0 +1,34 @@ +local LayerRepo = nerv.class("nerv.LayerRepo") + +function LayerRepo:__init(layer_spec, param_repo, global_conf) + local layers = {} + for ltype, llist in pairs(layer_spec) do + local layer_type = nerv.get_type(ltype) + for id, spec in pairs(llist) do + if layers[id] ~= nil then + nerv.error("a layer with id %s already exists", id) + end + nerv.utils.printf("id: %s\n", id) + if type(spec[2]) ~= "table" then + nerv.error("layer config table is need") + end + layer_config = spec[2] + if type(spec[1]) ~= "table" then + nerv.error("parameter description table is needed") + end + for pname, pid in pairs(spec[1]) do + layer_config[pname] = param_repo:get_param(pid, global_conf) + end + layers[id] = layer_type(id, global_conf, layer_config) + end + end + self.layers = layers +end + +function LayerRepo:get_layer(lid) + local layer = self.layers[lid] + if layer == nil then + nerv.error("layer with id %s not found", lid) + end + return layer +end |