-
Notifications
You must be signed in to change notification settings - Fork 85
[feat] support ftrl sparse optimizer for dynamicemb #682
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -15,20 +15,32 @@ | |
| from parameterized import param, parameterized | ||
| from torch import Tensor, nn | ||
| from torch.optim import Optimizer | ||
| from torchrec.distributed.utils import _OPTIMIZER_CLASS_TO_EMB_OPT_TYPE | ||
| from torchrec.optim.keyed import KeyedOptimizerWrapper | ||
|
|
||
| from tzrec.optim import optimizer_builder | ||
| from tzrec.optim.optimizer import ( | ||
| FTRL, | ||
| set_sparse_init_accumulator_value, | ||
| sparse_init_accumulator_value, | ||
| ) | ||
| from tzrec.protos import optimizer_pb2 | ||
| from tzrec.utils.test_util import parameterized_name_func | ||
| from tzrec.utils.test_util import mark_ci_scope, parameterized_name_func | ||
|
|
||
| try: | ||
| from dynamicemb import DynamicEmbOptimType | ||
|
|
||
| has_dynamicemb = True | ||
| except ImportError: | ||
| has_dynamicemb = False | ||
|
|
||
|
|
||
| class OpimizerBuilderTest(unittest.TestCase): | ||
| def tearDown(self): | ||
| set_sparse_init_accumulator_value(0.0) | ||
| # The whole suite runs in one process, so leaving FTRL registered would | ||
| # mask a missing register_ftrl_emb_opt_type() call in a later test. | ||
| _OPTIMIZER_CLASS_TO_EMB_OPT_TYPE.pop(FTRL, None) | ||
|
|
||
| @parameterized.expand( | ||
| [ | ||
|
|
@@ -73,6 +85,40 @@ def test_create_sparse_optimizer_init_accumulator_value( | |
| self.assertNotIn("initial_accumulator_value", kwargs) | ||
| self.assertAlmostEqual(sparse_init_accumulator_value(), expected) | ||
|
|
||
| @unittest.skipUnless( | ||
| has_dynamicemb, "dynamicemb with FTRL support is not installed." | ||
| ) | ||
| @mark_ci_scope("gpu") | ||
| def test_create_sparse_optimizer_ftrl(self): | ||
| from torchrec.distributed.utils import optimizer_type_to_emb_opt_type | ||
|
|
||
| optimizer_config = optimizer_pb2.SparseOptimizer( | ||
| ftrl_optimizer=optimizer_pb2.FusedFTRLOptimizer( | ||
| lr=0.01, | ||
| learning_rate_power=-0.4, | ||
| ftrl_beta=1.0, | ||
| l1_reg=0.01, | ||
| l2_reg=0.02, | ||
| initial_accumulator_value=0.1, | ||
| ), | ||
| constant_learning_rate=optimizer_pb2.ConstantLR(), | ||
| ) | ||
| optim_cls, kwargs = optimizer_builder.create_sparse_optimizer(optimizer_config) | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Minor test hygiene: this call runs |
||
| self.assertIs(optim_cls, FTRL) | ||
| # torchrec turns the class into the fused optimizer param of the table. | ||
| self.assertIs( | ||
| optimizer_type_to_emb_opt_type(optim_cls), DynamicEmbOptimType.FTRL | ||
| ) | ||
| self.assertNotIn("initial_accumulator_value", kwargs) | ||
| self.assertAlmostEqual(sparse_init_accumulator_value(), 0.1) | ||
| self.assertAlmostEqual(kwargs["lr"], 0.01) | ||
| self.assertAlmostEqual(kwargs["learning_rate_power"], -0.4) | ||
| self.assertAlmostEqual(kwargs["ftrl_beta"], 1.0) | ||
| self.assertAlmostEqual(kwargs["l1_reg"], 0.01) | ||
| self.assertAlmostEqual(kwargs["l2_reg"], 0.02) | ||
| # apply_optimizer_in_backward constructs the class with these kwargs. | ||
| optim_cls([nn.Parameter(torch.zeros(4, 8))], **kwargs) | ||
|
|
||
| def test_create_part_optimizer(self): | ||
| pattern1 = "model.dbmtl.task(.*)" | ||
| pattern2 = "model.dbmtl.mmoe(.*)" | ||
|
|
||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -13,6 +13,7 @@ message SparseOptimizer { | |
| FusedRowWiseAdagradOptimizer rowwise_adagrad_optimizer = 8; | ||
| FusedAdadeltaOptimizer adadelta_optimizer = 9; | ||
| FusedRMSpropOptimizer rmsprop_optimizer = 10; | ||
| FusedFTRLOptimizer ftrl_optimizer = 11; | ||
| } | ||
| oneof learning_rate { | ||
| ConstantLR constant_learning_rate = 101; | ||
|
|
@@ -160,6 +161,26 @@ message FusedRMSpropOptimizer { | |
| optional float max_gradient = 6 [default = 1.0]; | ||
| } | ||
|
|
||
| // FTRL-Proximal, Algorithm 1 of McMahan et al. 2013, with the paper's fixed | ||
| // square root generalized to an arbitrary exponent. dynamicemb tables only. | ||
| message FusedFTRLOptimizer { | ||
| // alpha in the paper, must be > 0 because the update divides by it. | ||
| required float lr = 1 [default = 0.002]; | ||
| // exponent on the accumulator in the learning rate, must be <= 0. | ||
| // -0.5 recovers the paper's alpha / (beta + sqrt(n)). | ||
| optional float learning_rate_power = 2 [default = -0.5]; | ||
| // beta in the paper, bounds the learning rate while the accumulator is small. | ||
| optional float ftrl_beta = 3 [default = 0.0]; | ||
| // lambda1 in the paper, L1 regularization. | ||
| optional float l1_reg = 4 [default = 0.0]; | ||
| // lambda2 in the paper, L2 regularization. | ||
| optional float l2_reg = 5 [default = 0.0]; | ||
| optional bool gradient_clipping = 6 [default = false]; | ||
| optional float max_gradient = 7 [default = 1.0]; | ||
|
Comment on lines
+178
to
+179
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
Unlike the sibling Since the message is new in this PR, consider dropping both fields, or keeping them with a comment marking them unsupported and warning in |
||
| // initial value of the squared-gradient accumulator. | ||
| optional float initial_accumulator_value = 8 [default = 0.0]; | ||
| } | ||
|
|
||
| message SGDOptimizer { | ||
| required float lr = 1 [default = 0.002]; | ||
| optional float momentum = 2 [default = 0.9]; | ||
|
|
||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Nit: this setter docstring now says the switch is used "by Adagrad and, on dynamicemb tables, by FTRL", but the paired getter
sparse_init_accumulator_value()(line 141) still reads "Sparse Adagrad accumulator initial value". Worth syncing the one-liner.