Class: DNN::Layers::SimpleRNN

Inherits:
HasParamLayer show all
Includes:
Activations, Initializers
Defined in:
lib/dnn/core/rnn_layers.rb

Constant Summary

Constants included from Activations

Activations::Layer

Instance Attribute Summary collapse

Attributes inherited from HasParamLayer

#grads, #params

Class Method Summary collapse

Instance Method Summary collapse

Methods inherited from HasParamLayer

#build, #update

Methods inherited from Layer

#build, #built?, #prev_layer

Constructor Details

#initialize(num_nodes, stateful: false, activation: nil, weight_initializer: nil, bias_initializer: nil, weight_decay: 0) ⇒ SimpleRNN

Returns a new instance of SimpleRNN.



21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
# File 'lib/dnn/core/rnn_layers.rb', line 21

def initialize(num_nodes,
               stateful: false,
               activation: nil,
               weight_initializer: nil,
               bias_initializer: nil,
               weight_decay: 0)
  super()
  @num_nodes = num_nodes
  @stateful = stateful
  @activation = (activation || Tanh.new)
  @weight_initializer = (weight_initializer || RandomNormal.new)
  @bias_initializer = (bias_initializer || Zeros.new)
  @weight_decay = weight_decay
  @h = nil
end

Instance Attribute Details

#num_nodes ⇒ Object (readonly)

Returns the value of attribute num_nodes.



8
9
10
# File 'lib/dnn/core/rnn_layers.rb', line 8

def num_nodes
  @num_nodes
end

#stateful ⇒ Object (readonly)

Returns the value of attribute stateful.



9
10
11
# File 'lib/dnn/core/rnn_layers.rb', line 9

def stateful
  @stateful
end

#weight_decay ⇒ Object (readonly)

Returns the value of attribute weight_decay.



10
11
12
# File 'lib/dnn/core/rnn_layers.rb', line 10

def weight_decay
  @weight_decay
end

Class Method Details

.load_hash(hash) ⇒ Object



12
13
14
15
16
17
18
19
# File 'lib/dnn/core/rnn_layers.rb', line 12

def self.load_hash(hash)
  self.new(hash[:num_nodes],
           stateful: hash[:stateful],
           activation: Util.load_hash(hash[:activation]),
           weight_initializer: Util.load_hash(hash[:weight_initializer]),
           bias_initializer: Util.load_hash(hash[:bias_initializer]),
           weight_decay: hash[:weight_decay])
end

Instance Method Details

#backward(douts) ⇒ Object



51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
# File 'lib/dnn/core/rnn_layers.rb', line 51

def backward(douts)
  @grads[:weight] = SFloat.zeros(*@params[:weight].shape)
  @grads[:weight2] = SFloat.zeros(*@params[:weight2].shape)
  dxs = SFloat.zeros(@xs.shape)
  (0...douts.shape[1]).to_a.reverse.each do |t|
    dout = douts[true, t, false]
    x = @xs[true, t, false]
    h = @hs[true, t, false]
    dout = @activation.backward(dout)
    @grads[:weight] += x.transpose.dot(dout)
    @grads[:weight2] += h.transpose.dot(dout)
    dxs[true, t, false] = dout.dot(@params[:weight].transpose)
  end
  @grads[:bias] = douts.sum(0).sum(0)
  dxs
end

#forward(xs) ⇒ Object



37
38
39
40
41
42
43
44
45
46
47
48
49
# File 'lib/dnn/core/rnn_layers.rb', line 37

def forward(xs)
  @xs = xs
  @hs = SFloat.zeros(xs.shape[0], *shape)
  h = (@stateful && @h) ? @h : SFloat.zeros(xs.shape[0], @num_nodes)
  xs.shape[1].times do |t|
    x = xs[true, t, false]
    h = x.dot(@params[:weight]) + h.dot(@params[:weight2]) + @params[:bias]
    h = @activation.forward(h)
    @hs[true, t, false] = h
  end
  @h = h
  @hs
end

#ridge ⇒ Object



72
73
74
75
76
77
78
# File 'lib/dnn/core/rnn_layers.rb', line 72

def ridge
  if @weight_decay > 0
    0.5 * (@weight_decay * (@params[:weight]**2).sum + @weight_decay * (@params[:weight]**2).sum)
  else
    0
  end
end

#shape ⇒ Object



68
69
70
# File 'lib/dnn/core/rnn_layers.rb', line 68

def shape
  [@time_length, @num_nodes]
end

#to_hash ⇒ Object



80
81
82
83
84
85
86
87
# File 'lib/dnn/core/rnn_layers.rb', line 80

def to_hash
  super({num_nodes: @num_nodes,
         stateful: @stateful,
         activation: @activation.to_hash,
         weight_initializer: @weight_initializer.to_hash,
         bias_initializer: @bias_initializer.to_hash,
         weight_decay: @weight_decay})
end