Class: DNN::Layers::RNN
- Inherits:
-
Connection
- Object
- Layer
- HasParamLayer
- Connection
- DNN::Layers::RNN
- Includes:
- Initializers
- Defined in:
- lib/dnn/core/rnn_layers.rb
Overview
Super class of all RNN classes.
Instance Attribute Summary collapse
-
#num_nodes ⇒ Integer
readonly
Number of nodes.
-
#return_sequences ⇒ Bool
readonly
Set the false, only the last of each cell of RNN is left.
-
#stateful ⇒ Bool
readonly
Maintain state between batches.
Attributes inherited from Connection
#bias_initializer, #l1_lambda, #l2_lambda, #weight_initializer
Attributes inherited from HasParamLayer
Attributes inherited from Layer
Instance Method Summary collapse
- #backward(dh2s) ⇒ Object
- #forward(xs) ⇒ Object
-
#initialize(num_nodes, stateful: false, return_sequences: true, weight_initializer: RandomNormal.new, bias_initializer: Zeros.new, l1_lambda: 0, l2_lambda: 0, use_bias: true) ⇒ RNN
constructor
A new instance of RNN.
- #output_shape ⇒ Object
- #regularizers ⇒ Object
-
#reset_state ⇒ Object
Reset the state of RNN.
- #to_hash(merge_hash = nil) ⇒ Object
Methods inherited from Connection
Methods inherited from HasParamLayer
Methods inherited from Layer
Constructor Details
#initialize(num_nodes, stateful: false, return_sequences: true, weight_initializer: RandomNormal.new, bias_initializer: Zeros.new, l1_lambda: 0, l2_lambda: 0, use_bias: true) ⇒ RNN
Returns a new instance of RNN.
15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 |
# File 'lib/dnn/core/rnn_layers.rb', line 15 def initialize(num_nodes, stateful: false, return_sequences: true, weight_initializer: RandomNormal.new, bias_initializer: Zeros.new, l1_lambda: 0, l2_lambda: 0, use_bias: true) super(weight_initializer: weight_initializer, bias_initializer: bias_initializer, l1_lambda: l1_lambda, l2_lambda: l2_lambda, use_bias: use_bias) @num_nodes = num_nodes @stateful = stateful @return_sequences = return_sequences @layers = [] @hidden = @params[:h] = Param.new # TODO # Change to a good name. @params[:weight2] = @weight2 = Param.new end |
Instance Attribute Details
#num_nodes ⇒ Integer (readonly)
Returns number of nodes.
9 10 11 |
# File 'lib/dnn/core/rnn_layers.rb', line 9 def num_nodes @num_nodes end |
#return_sequences ⇒ Bool (readonly)
Returns Set the false, only the last of each cell of RNN is left.
13 14 15 |
# File 'lib/dnn/core/rnn_layers.rb', line 13 def return_sequences @return_sequences end |
#stateful ⇒ Bool (readonly)
Returns Maintain state between batches.
11 12 13 |
# File 'lib/dnn/core/rnn_layers.rb', line 11 def stateful @stateful end |
Instance Method Details
#backward(dh2s) ⇒ Object
48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 |
# File 'lib/dnn/core/rnn_layers.rb', line 48 def backward(dh2s) @weight.grad = Xumo::SFloat.zeros(*@weight.data.shape) @weight2.grad = Xumo::SFloat.zeros(*@weight2.data.shape) @bias.grad = Xumo::SFloat.zeros(*@bias.data.shape) if @bias unless @return_sequences dh = dh2s dh2s = Xumo::SFloat.zeros(dh.shape[0], @time_length, dh.shape[1]) dh2s[true, -1, false] = dh end dxs = Xumo::SFloat.zeros(@xs_shape) dh = 0 (0...dh2s.shape[1]).to_a.reverse.each do |t| dh2 = dh2s[true, t, false] dx, dh = @layers[t].backward(dh2 + dh) dxs[true, t, false] = dx end dxs end |
#forward(xs) ⇒ Object
35 36 37 38 39 40 41 42 43 44 45 46 |
# File 'lib/dnn/core/rnn_layers.rb', line 35 def forward(xs) @xs_shape = xs.shape hs = Xumo::SFloat.zeros(xs.shape[0], @time_length, @num_nodes) h = (@stateful && @hidden.data) ? @hidden.data : Xumo::SFloat.zeros(xs.shape[0], @num_nodes) xs.shape[1].times do |t| x = xs[true, t, false] h = @layers[t].forward(x, h) hs[true, t, false] = h end @hidden.data = h @return_sequences ? hs : h end |
#output_shape ⇒ Object
67 68 69 |
# File 'lib/dnn/core/rnn_layers.rb', line 67 def output_shape @return_sequences ? [@time_length, @num_nodes] : [@num_nodes] end |
#regularizers ⇒ Object
86 87 88 89 90 91 92 93 94 95 96 97 |
# File 'lib/dnn/core/rnn_layers.rb', line 86 def regularizers regularizers = [] if @l1_lambda > 0 regularizers << Lasso.new(@l1_lambda, @weight) regularizers << Lasso.new(@l1_lambda, @weight2) end if @l2_lambda > 0 regularizers << Ridge.new(@l2_lambda, @weight) regularizers << Ridge.new(@l2_lambda, @weight2) end regularizers end |
#reset_state ⇒ Object
Reset the state of RNN.
82 83 84 |
# File 'lib/dnn/core/rnn_layers.rb', line 82 def reset_state @hidden.data = @hidden.data.fill(0) if @hidden.data end |
#to_hash(merge_hash = nil) ⇒ Object
71 72 73 74 75 76 77 78 79 |
# File 'lib/dnn/core/rnn_layers.rb', line 71 def to_hash(merge_hash = nil) hash = { num_nodes: @num_nodes, stateful: @stateful, return_sequences: @return_sequences } hash.merge!(merge_hash) if merge_hash super(hash) end |