MCPcopy Create free account
hub / github.com/apache/singa / __init__

Method __init__

examples/cnn_ms/train_ms_model.py:90–141  ·  view source on GitHub ↗
(self,
                 lr=0.1,
                 momentum=0,
                 dampening=0,
                 weight_decay=0,
                 nesterov=False,
                 dtype=tensor.float32)

Source from the content-addressed store, hash-verified

88 """
89
90 def __init__(self,
91 lr=0.1,
92 momentum=0,
93 dampening=0,
94 weight_decay=0,
95 nesterov=False,
96 dtype=tensor.float32):
97 super(MSSGD, self).__init__(lr, dtype)
98
99 # init momentum
100 if type(momentum) == float or type(momentum) == int:
101 if momentum < 0.0:
102 raise ValueError("Invalid momentum value: {}".format(momentum))
103 self.momentum = Constant(momentum)
104 elif isinstance(momentum, DecayScheduler):
105 self.momentum = momentum
106 momentum = momentum.init_value
107 else:
108 raise TypeError("Wrong momentum type")
109 self.mom_value = self.momentum(self.step_counter).as_type(self.dtype)
110
111 # init dampening
112 if type(dampening) == float or type(dampening) == int:
113 self.dampening = Constant(dampening)
114 elif isinstance(dampening, DecayScheduler):
115 self.dampening = dampening
116 dampening = dampening.init_value
117 else:
118 raise TypeError("Wrong dampening type")
119 self.dam_value = self.dampening(self.step_counter).as_type(self.dtype)
120
121 # init weight_decay
122 if type(weight_decay) == float or type(weight_decay) == int:
123 if weight_decay < 0.0:
124 raise ValueError(
125 "Invalid weight_decay value: {}".format(weight_decay))
126 self.weight_decay = Constant(weight_decay)
127 elif isinstance(weight_decay, DecayScheduler):
128 self.weight_decay = weight_decay
129 else:
130 raise TypeError("Wrong weight_decay type")
131 self.decay_value = self.weight_decay(self.step_counter).as_type(
132 self.dtype)
133
134 # init other params
135 self.nesterov = nesterov
136 self.moments = dict()
137
138 # check value
139 if nesterov and (momentum <= 0 or dampening != 0):
140 raise ValueError(
141 "Nesterov momentum requires a momentum and zero dampening")
142
143 def apply(self, param_name, param_value, param_grad):
144 """Performs a single optimization step.

Callers

nothing calls this directly

Calls 3

ConstantClass · 0.90
typeFunction · 0.85
as_typeMethod · 0.45

Tested by

no test coverage detected