(self, net, param_init_net, param_info)
| 1430 | self.init_kwargs = kwargs |
| 1431 | |
| 1432 | def _run(self, net, param_init_net, param_info): |
| 1433 | param = param_info.blob |
| 1434 | grad = param_info.grad |
| 1435 | |
| 1436 | if self.alpha <= 0: |
| 1437 | return |
| 1438 | |
| 1439 | lr, _ = self.build_lr( |
| 1440 | net, |
| 1441 | param_init_net, |
| 1442 | base_learning_rate=self.alpha, |
| 1443 | policy=self.policy, |
| 1444 | **(self.init_kwargs) |
| 1445 | ) |
| 1446 | |
| 1447 | moment = param_init_net.ConstantFill( |
| 1448 | [param], str(param) + "_squared_moment", value=0.0 |
| 1449 | ) |
| 1450 | |
| 1451 | moment_update = param_init_net.ConstantFill( |
| 1452 | [param], str(param) + "_squared_moment_update", value=0.0 |
| 1453 | ) |
| 1454 | |
| 1455 | self._aux_params.local.append(moment) |
| 1456 | self._aux_params.local.append(moment_update) |
| 1457 | |
| 1458 | if isinstance(grad, core.GradientSlice): |
| 1459 | grad = self.dedup(net, self.sparse_dedup_aggregator, grad) |
| 1460 | net.SparseAdadelta( |
| 1461 | [param, moment, moment_update, grad.indices, grad.values, lr], |
| 1462 | [param, moment, moment_update], |
| 1463 | epsilon=self.epsilon, |
| 1464 | decay=self.decay, |
| 1465 | engine=self.engine, |
| 1466 | ) |
| 1467 | else: |
| 1468 | net.Adadelta( |
| 1469 | [param, moment, moment_update, grad, lr], |
| 1470 | [param, moment, moment_update], |
| 1471 | epsilon=self.epsilon, |
| 1472 | decay=self.decay, |
| 1473 | engine=self.engine, |
| 1474 | ) |
| 1475 | |
| 1476 | def scale_learning_rate(self, scale): |
| 1477 | self.alpha *= scale |
nothing calls this directly
no test coverage detected