Group Lasso FTRL Optimizer.
| 1535 | |
| 1536 | |
| 1537 | class GFtrlOptimizer(Optimizer): |
| 1538 | """Group Lasso FTRL Optimizer.""" |
| 1539 | |
| 1540 | def __init__( |
| 1541 | self, |
| 1542 | alpha=0.01, |
| 1543 | beta=1e-4, |
| 1544 | lambda1=0, |
| 1545 | lambda2=0, |
| 1546 | sparse_dedup_aggregator=None, |
| 1547 | engine="", |
| 1548 | ): |
| 1549 | super().__init__() |
| 1550 | self.alpha = alpha |
| 1551 | self.beta = beta |
| 1552 | self.lambda1 = lambda1 |
| 1553 | self.lambda2 = lambda2 |
| 1554 | self.sparse_dedup_aggregator = sparse_dedup_aggregator |
| 1555 | self.engine = engine |
| 1556 | |
| 1557 | def _run(self, net, param_init_net, param_info): |
| 1558 | param = param_info.blob |
| 1559 | grad = param_info.grad |
| 1560 | |
| 1561 | if self.alpha <= 0: |
| 1562 | return |
| 1563 | |
| 1564 | nz = param_init_net.ConstantFill( |
| 1565 | [param], str(param) + "_gftrl_nz", extra_shape=[2], value=0.0 |
| 1566 | ) |
| 1567 | self._aux_params.local.append(nz) |
| 1568 | net.GFtrl( |
| 1569 | [param, nz, grad], |
| 1570 | [param, nz], |
| 1571 | engine=self.engine, |
| 1572 | alpha=self.alpha, |
| 1573 | beta=self.beta, |
| 1574 | lambda1=self.lambda1, |
| 1575 | lambda2=self.lambda2, |
| 1576 | ) |
| 1577 | |
| 1578 | def scale_learning_rate(self, scale): |
| 1579 | self.alpha *= scale |
| 1580 | return |
| 1581 | |
| 1582 | |
| 1583 | class AdamOptimizer(Optimizer): |
no outgoing calls
no test coverage detected
searching dependent graphs…