Class: DNN::Layers::RNN
Overview
Super class of all RNN classes.
Instance Attribute Summary collapse
Attributes inherited from Connection
#l1_lambda, #l2_lambda
#grads, #params, #trainable
Instance Method Summary
collapse
#build, #update
Methods inherited from Layer
#build, #built?, #prev_layer
Constructor Details
#initialize(num_nodes, stateful: false, return_sequences: true, weight_initializer: nil, bias_initializer: nil, l1_lambda: 0, l2_lambda: 0) ⇒ RNN
Returns a new instance of RNN.
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
|
# File 'lib/dnn/core/rnn_layers.rb', line 12
def initialize(num_nodes,
stateful: false,
return_sequences: true,
weight_initializer: nil,
bias_initializer: nil,
l1_lambda: 0,
l2_lambda: 0)
super(weight_initializer: weight_initializer, bias_initializer: bias_initializer,
l1_lambda: l1_lambda, l2_lambda: l2_lambda)
@num_nodes = num_nodes
@stateful = stateful
@return_sequences = return_sequences
@layers = []
@h = nil
end
|
Instance Attribute Details
#h ⇒ Object
Returns the value of attribute h.
8
9
10
|
# File 'lib/dnn/core/rnn_layers.rb', line 8
def h
@h
end
|
#num_nodes ⇒ Object
Returns the value of attribute num_nodes.
9
10
11
|
# File 'lib/dnn/core/rnn_layers.rb', line 9
def num_nodes
@num_nodes
end
|
#stateful ⇒ Object
Returns the value of attribute stateful.
10
11
12
|
# File 'lib/dnn/core/rnn_layers.rb', line 10
def stateful
@stateful
end
|
Instance Method Details
#backward(dh2s) ⇒ Object
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
|
# File 'lib/dnn/core/rnn_layers.rb', line 41
def backward(dh2s)
@grads[:weight] = Xumo::SFloat.zeros(*@params[:weight].shape)
@grads[:weight2] = Xumo::SFloat.zeros(*@params[:weight2].shape)
@grads[:bias] = Xumo::SFloat.zeros(*@params[:bias].shape)
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
|
#dlasso ⇒ Object
95
96
97
98
99
|
# File 'lib/dnn/core/rnn_layers.rb', line 95
def dlasso
dlasso = Xumo::SFloat.ones(*@params[:weight].shape)
dlasso[@params[:weight] < 0] = -1
@l1_lambda * dlasso
end
|
#dlasso2 ⇒ Object
105
106
107
108
109
|
# File 'lib/dnn/core/rnn_layers.rb', line 105
def dlasso2
dlasso = Xumo::SFloat.ones(*@params[:weight2].shape)
dlasso[@params[:weight2] < 0] = -1
@l1_lambda * dlasso
end
|
#dridge ⇒ Object
101
102
103
|
# File 'lib/dnn/core/rnn_layers.rb', line 101
def dridge
@l2_lambda * @params[:weight]
end
|
#dridge2 ⇒ Object
111
112
113
|
# File 'lib/dnn/core/rnn_layers.rb', line 111
def dridge2
@l2_lambda * @params[:weight2]
end
|
#forward(xs) ⇒ Object
28
29
30
31
32
33
34
35
36
37
38
39
|
# File 'lib/dnn/core/rnn_layers.rb', line 28
def forward(xs)
@xs_shape = xs.shape
hs = Xumo::SFloat.zeros(xs.shape[0], @time_length, @num_nodes)
h = (@stateful && @h) ? @h : 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
@h = h
@return_sequences ? hs : h
end
|
#lasso ⇒ Object
79
80
81
82
83
84
85
|
# File 'lib/dnn/core/rnn_layers.rb', line 79
def lasso
if @l1_lambda > 0
@l1_lambda * (@params[:weight].abs.sum + @params[:weight2].abs.sum)
else
0
end
end
|
#reset_state ⇒ Object
75
76
77
|
# File 'lib/dnn/core/rnn_layers.rb', line 75
def reset_state
@h = @h.fill(0) if @h
end
|
#ridge ⇒ Object
87
88
89
90
91
92
93
|
# File 'lib/dnn/core/rnn_layers.rb', line 87
def ridge
if @l2_lambda > 0
0.5 * (@l2_lambda * ((@params[:weight]**2).sum + (@params[:weight2]**2).sum))
else
0
end
end
|
#shape ⇒ Object
71
72
73
|
# File 'lib/dnn/core/rnn_layers.rb', line 71
def shape
@return_sequences ? [@time_length, @num_nodes] : [@num_nodes]
end
|
#to_hash(merge_hash = nil) ⇒ Object
60
61
62
63
64
65
66
67
68
69
|
# File 'lib/dnn/core/rnn_layers.rb', line 60
def to_hash(merge_hash = nil)
hash = {
num_nodes: @num_nodes,
stateful: @stateful,
return_sequences: @return_sequences,
h: @h.to_a
}
hash.merge!(merge_hash) if merge_hash
super(hash)
end
|