(self, net, param_init_net, param_info)
| 1497 | self.engine = engine |
| 1498 | |
| 1499 | def _run(self, net, param_init_net, param_info): |
| 1500 | param = param_info.blob |
| 1501 | grad = param_info.grad |
| 1502 | |
| 1503 | if self.alpha <= 0: |
| 1504 | return |
| 1505 | |
| 1506 | nz = param_init_net.ConstantFill( |
| 1507 | [param], str(param) + "_ftrl_nz", extra_shape=[2], value=0.0 |
| 1508 | ) |
| 1509 | self._aux_params.local.append(nz) |
| 1510 | if isinstance(grad, core.GradientSlice): |
| 1511 | grad = self.dedup(net, self.sparse_dedup_aggregator, grad) |
| 1512 | net.SparseFtrl( |
| 1513 | [param, nz, grad.indices, grad.values], |
| 1514 | [param, nz], |
| 1515 | engine=self.engine, |
| 1516 | alpha=self.alpha, |
| 1517 | beta=self.beta, |
| 1518 | lambda1=self.lambda1, |
| 1519 | lambda2=self.lambda2, |
| 1520 | ) |
| 1521 | else: |
| 1522 | net.Ftrl( |
| 1523 | [param, nz, grad], |
| 1524 | [param, nz], |
| 1525 | engine=self.engine, |
| 1526 | alpha=self.alpha, |
| 1527 | beta=self.beta, |
| 1528 | lambda1=self.lambda1, |
| 1529 | lambda2=self.lambda2, |
| 1530 | ) |
| 1531 | |
| 1532 | def scale_learning_rate(self, scale): |
| 1533 | self.alpha *= scale |
nothing calls this directly
no test coverage detected