(self)
| 222 | return net |
| 223 | |
| 224 | def optimizer(self): |
| 225 | loss_func = tf.losses.mean_squared_error |
| 226 | self.predict = tf.squeeze(self.predict) |
| 227 | loss = tf.math.reduce_mean(loss_func(self.label, self.predict)) |
| 228 | tf.summary.scalar('loss', loss) |
| 229 | |
| 230 | self.global_step = tf.train.get_or_create_global_step() |
| 231 | if self.optimizer_type == 'adam': |
| 232 | optimizer = tf.train.AdamOptimizer( |
| 233 | learning_rate=self.learning_rate, |
| 234 | beta1=0.9, |
| 235 | beta2=0.999, |
| 236 | epsilon=1e-8) |
| 237 | elif self.optimizer_type == 'adagrad': |
| 238 | optimizer = tf.train.AdagradOptimizer( |
| 239 | learning_rate=self.learning_rate, |
| 240 | initial_accumulator_value=1e-8) |
| 241 | elif self.optimizer_type == 'adamasync': |
| 242 | optimizer = tf.train.AdamAsyncOptimizer( |
| 243 | learning_rate=self.learning_rate, |
| 244 | beta1=0.9, |
| 245 | beta2=0.999, |
| 246 | epsilon=1e-8) |
| 247 | else: |
| 248 | raise ValueError('Optimizer type error.') |
| 249 | |
| 250 | update_ops = tf.get_collection(tf.GraphKeys.UPDATE_OPS) |
| 251 | with tf.control_dependencies(update_ops): |
| 252 | train_op = optimizer.minimize(loss, global_step=self.global_step) |
| 253 | |
| 254 | return train_op, loss |
| 255 | |
| 256 | |
| 257 | def get_arg_parser(): |
no test coverage detected