diff --git a/docs/source/feature/dynamicemb.md b/docs/source/feature/dynamicemb.md index 014d7398..22f0435b 100644 --- a/docs/source/feature/dynamicemb.md +++ b/docs/source/feature/dynamicemb.md @@ -9,10 +9,14 @@ cu126/cu129/cu130 镜像已预装 dynamicemb,其它环境需先安装如下whl pip install dynamicemb==0.1.0+20260920.9643985.${DEVICE} -f https://tzrec.oss-accelerate.aliyuncs.com/third_party/dynamicemb/${DEVICE}/repo.html ``` +支持配置 DynamicEmbedding 的特征类型为 `id_feature`、`combo_feature`、`lookup_feature`、`match_feature`、`regex_replace_feature`、`custom_feature`、`bool_mask_feature`,其余特征类型(如配置了 `boundaries` 的 `raw_feature`、`expr_feature`、`tokenize_feature`、`combine_feature`)暂不支持。 + 注:同一个 FeatureGroup 中若存在多个配置了 DynamicEmbedding 的特征,底层 dynamicemb 会自动将这些表融合到同一份存储里(table fusion),共享 cache/admission counter,降低显存占用并减少内存碎片,无需额外配置。 注:配置了 DynamicEmbedding 的模型导出时需设置环境变量 `USE_DISTRIBUTED_EMBEDDING=1`,使用分布式 embedding 导出模式,详见[模型导出](../usage/export.md)的环境变量章节。 +注:DynamicEmbedding 表额外支持 `ftrl_optimizer`(普通 Embedding 表不支持),配置方式见[优化器](../models/optimizer.md)文档。 + 以id_feature的配置为例,DynamicEmbedding 只需在id_feature新增一个dynamicemb的配置字段 ``` diff --git a/docs/source/models/optimizer.md b/docs/source/models/optimizer.md index 77853a3a..486427f0 100644 --- a/docs/source/models/optimizer.md +++ b/docs/source/models/optimizer.md @@ -52,7 +52,9 @@ train_config { **Note**: 被分片为`data_parallel`的Embedding表不受sparse_optimizer管理,实际由dense_optimizer更新,并跟随dense_optimizer的LR策略,详见[训练文档](../usage/train.md)的Embedding分片约束章节 - **Note**: `adagrad_optimizer`和`rowwise_adagrad_optimizer`的`initial_accumulator_value`对齐TensorFlow Adagrad的同名参数(TF默认0.1,TorchEasyRec默认0.0),对普通Embedding表和[dynamicemb](../feature/dynamicemb.md)表同时生效:普通Embedding表在建表时把整个accumulator初始化为该值,dynamicemb表则在key首次写入时把该key的accumulator初始化为该值 + **Note**: `adagrad_optimizer`和`rowwise_adagrad_optimizer`的`initial_accumulator_value`对齐TensorFlow Adagrad的同名参数(TF默认0.1,TorchEasyRec默认0.0),对普通Embedding表和[dynamicemb](../feature/dynamicemb.md)表同时生效:普通Embedding表在建表时把整个accumulator初始化为该值,dynamicemb表则在key首次写入时把该key的accumulator初始化为该值。`ftrl_optimizer`也用该字段初始化它的accumulator + + **Note**: `ftrl_optimizer`(FTRL-Proximal,McMahan et al. 2013)**只支持[dynamicemb](../feature/dynamicemb.md)表**,FBGEMM没有FTRL的embedding kernel,模型中只要还有一张非dynamicemb的sparse表,训练会在plan阶段直接报错并给出表名。可配置`dynamicemb`的特征类型见[dynamicemb文档](../feature/dynamicemb.md),配置了`boundaries`的`raw_feature`等不支持dynamicemb的特征,无法与`ftrl_optimizer`一起使用。被分片为`data_parallel`的表不受此限制(由dense_optimizer更新) - dense_optimizer diff --git a/tzrec/optim/optimizer.py b/tzrec/optim/optimizer.py index 074663fd..700bc9f3 100644 --- a/tzrec/optim/optimizer.py +++ b/tzrec/optim/optimizer.py @@ -9,7 +9,7 @@ # See the License for the specific language governing permissions and # limitations under the License. import logging -from typing import Any, Callable, Optional, Union +from typing import Any, Callable, Iterable, Optional, Union import torch from fbgemm_gpu import split_table_batched_embeddings_ops_training @@ -19,6 +19,8 @@ ) from torch import nn from torch.amp import GradScaler +from torch.optim.optimizer import Optimizer +from torchrec.distributed.utils import _OPTIMIZER_CLASS_TO_EMB_OPT_TYPE from torchrec.optim import KeyedOptimizer, OptimizerWrapper @@ -67,6 +69,51 @@ def step(self, closure: Any = None) -> None: self._optimizer.step(closure=closure) +class FTRL(Optimizer): + """Placeholder for the dynamicemb FTRL sparse embedding optimizer. + + FBGEMM has no FTRL kernel, so torchrec ships no FTRL wrapper to reuse. Like + torchrec's own placeholders this class never runs: it names the optimizer so + that the sharding plan can resolve it, and the update happens inside the + dynamicemb table. + + Args: + params (Iterable[nn.Parameter]): parameters to attach the optimizer to. + **kwargs: fused params, forwarded to the dynamicemb table. + """ + + def __init__(self, params: Iterable[nn.Parameter], **kwargs: Any) -> None: + self._params = params + self._kwargs = kwargs + + # pyrefly: ignore[bad-override] # matches torchrec's placeholder optimizers + def step(self, closure: Optional[Callable[[], float]] = None) -> None: + """Step, never reached, the dynamicemb table applies the update.""" + raise NotImplementedError + + +def register_ftrl_emb_opt_type() -> None: + """Map :class:`FTRL` to dynamicemb's optimizer type in torchrec's table. + + torchrec derives an embedding table's fused `optimizer` param from the + in-backward optimizer class, and does so unconditionally on the + EmbeddingCollection path, so FTRL has to be registered in that table rather + than injected into the fused params. Must be called before planning. + + Raises: + RuntimeError: dynamicemb is missing or predates its FTRL support. + """ + try: + from dynamicemb import DynamicEmbOptimType + except ImportError as e: + raise RuntimeError( + "sparse ftrl_optimizer requires dynamicemb >= " + "0.1.0+20260920.9643985; FBGEMM has no FTRL embedding kernel. " + "Please reinstall dynamicemb, see docs/source/feature/dynamicemb.md." + ) from e + _OPTIMIZER_CLASS_TO_EMB_OPT_TYPE[FTRL] = DynamicEmbOptimType.FTRL + + # The Adagrad optimizer in TensorFlow includes the parameter # `initial_accumulator_value`, with a default value of 0.1. # Here, we patch the fbgemm embedding optimizer state split helper @@ -75,11 +122,13 @@ def step(self, closure: Any = None) -> None: def set_sparse_init_accumulator_value(value: float) -> None: - """Record the Adagrad accumulator initial value for embedding tables built later. + """Record the accumulator initial value for embedding tables built later. Takes effect at table build time, in the ``apply_split_helper`` patch below and in ``dynamicemb_util``'s plan-time fused params, so it must be set before - planning; FBGEMM TBE has no such kwarg, hence this module-level switch. + planning; FBGEMM TBE has no such kwarg, hence this module-level switch. Used + by Adagrad and, on dynamicemb tables, by FTRL for its squared-gradient + accumulator. Args: value: accumulator initial value, 0.0 for optimizers without one. @@ -89,7 +138,7 @@ def set_sparse_init_accumulator_value(value: float) -> None: def sparse_init_accumulator_value() -> float: - """Sparse Adagrad accumulator initial value, 0.0 when not configured.""" + """Sparse accumulator initial value, 0.0 when not configured.""" return _sparse_init_accumulator_value diff --git a/tzrec/optim/optimizer_builder.py b/tzrec/optim/optimizer_builder.py index 5b7649c9..4b5f4c28 100644 --- a/tzrec/optim/optimizer_builder.py +++ b/tzrec/optim/optimizer_builder.py @@ -21,7 +21,11 @@ from torchrec.optim.keyed import KeyedOptimizerWrapper from tzrec.optim.lr_scheduler import BaseLR -from tzrec.optim.optimizer import set_sparse_init_accumulator_value +from tzrec.optim.optimizer import ( + FTRL, + register_ftrl_emb_opt_type, + set_sparse_init_accumulator_value, +) from tzrec.protos import optimizer_pb2 from tzrec.utils.config_util import config_to_kwargs from tzrec.utils.logging_util import logger @@ -92,6 +96,9 @@ def create_sparse_optimizer( # FBGEMM reuses the beta1 OptimizerArgs slot for RMSProp's alpha. optimizer_kwargs["beta1"] = optimizer_kwargs.pop("alpha") return optimizers.RMSProp, optimizer_kwargs + elif optimizer_type == "ftrl_optimizer": + register_ftrl_emb_opt_type() + return FTRL, optimizer_kwargs else: raise ValueError(f"Unknown optimizer: {optimizer_type}") diff --git a/tzrec/optim/optimizer_builder_test.py b/tzrec/optim/optimizer_builder_test.py index f3c70c2a..70212a82 100644 --- a/tzrec/optim/optimizer_builder_test.py +++ b/tzrec/optim/optimizer_builder_test.py @@ -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) + 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(.*)" diff --git a/tzrec/protos/optimizer.proto b/tzrec/protos/optimizer.proto index ce74d119..08161e83 100644 --- a/tzrec/protos/optimizer.proto +++ b/tzrec/protos/optimizer.proto @@ -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]; + // 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]; diff --git a/tzrec/tests/configs/multi_tower_din_dynamicemb_ftrl_mock.config b/tzrec/tests/configs/multi_tower_din_dynamicemb_ftrl_mock.config new file mode 100644 index 00000000..ac1d88f5 --- /dev/null +++ b/tzrec/tests/configs/multi_tower_din_dynamicemb_ftrl_mock.config @@ -0,0 +1,187 @@ +train_input_path: "" +eval_input_path: "" +model_dir: "experiments/multi_tower_din_dynamicemb_ftrl_mock" +train_config { + sparse_optimizer { + ftrl_optimizer { + lr: 0.01 + learning_rate_power: -0.5 + ftrl_beta: 1.0 + l1_reg: 0.001 + l2_reg: 0.001 + initial_accumulator_value: 0.1 + } + constant_learning_rate { + } + } + dense_optimizer { + adam_optimizer { + lr: 0.001 + } + constant_learning_rate { + } + } + num_epochs: 1 +} +eval_config { +} +data_config { + batch_size: 8192 + dataset_type: ParquetDataset + label_fields: "clk" + num_workers: 8 +} +feature_configs { + id_feature { + feature_name: "id_1" + embedding_dim: 16 + dynamicemb { + max_capacity: 200000 + score_strategy: "TIMESTAMP" + init_capacity_per_rank: 16384 + } + } +} +feature_configs { + id_feature { + feature_name: "id_2" + embedding_dim: 16 + dynamicemb { + max_capacity: 65536 + score_strategy: "STEP" + frequency_admission_strategy { + threshold: 2 + } + } + } +} +feature_configs { + id_feature { + feature_name: "id_3" + embedding_dim: 8 + dynamicemb { + max_capacity: 32768 + score_strategy: "LFU" + } + } +} +feature_configs { + id_feature { + feature_name: "id_4" + embedding_dim: 16 + embedding_name: "id_4_emb" + dynamicemb { + max_capacity: 20480 + score_strategy: "NO_EVICTION" + init_capacity_per_rank: 8192 + } + } +} +feature_configs { + id_feature { + feature_name: "id_5" + embedding_dim: 16 + embedding_name: "id_4_emb" + dynamicemb { + max_capacity: 20480 + score_strategy: "NO_EVICTION" + init_capacity_per_rank: 8192 + } + } +} +feature_configs { + raw_feature { + feature_name: "raw_2" + } +} +feature_configs { + raw_feature { + feature_name: "raw_3" + value_dim: 4 + } +} +feature_configs { + sequence_feature { + sequence_name: "click_50_seq" + sequence_length: 50 + sequence_delim: "|" + features { + id_feature { + feature_name: "id_2" + embedding_dim: 16 + dynamicemb { + max_capacity: 65536 + score_strategy: "STEP" + frequency_admission_strategy { + threshold: 2 + } + } + } + } + features { + id_feature { + feature_name: "id_3" + embedding_dim: 8 + dynamicemb { + max_capacity: 32768 + score_strategy: "LFU" + } + } + } + features { + raw_feature { + feature_name: "raw_2" + } + } + } +} +model_config { + feature_groups { + group_name: "deep" + feature_names: "id_1" + feature_names: "id_2" + feature_names: "id_3" + feature_names: "id_4" + feature_names: "id_5" + feature_names: "raw_2" + feature_names: "raw_3" + group_type: DEEP + } + feature_groups { + group_name: "seq" + feature_names: "id_2" + feature_names: "id_3" + feature_names: "raw_2" + feature_names: "click_50_seq__id_2" + feature_names: "click_50_seq__id_3" + feature_names: "click_50_seq__raw_2" + group_type: SEQUENCE + } + multi_tower_din { + towers { + input: 'deep' + mlp { + hidden_units: [512, 256, 128] + } + } + din_towers { + input: 'seq' + attn_mlp { + hidden_units: [256, 64] + } + } + final { + hidden_units: [64] + } + } + metrics { + auc {} + } + train_metrics { + auc {} + decay_step: 1 + } + losses { + binary_cross_entropy {} + } +} diff --git a/tzrec/tests/rank_integration_test.py b/tzrec/tests/rank_integration_test.py index 22f0e67e..d05c9079 100644 --- a/tzrec/tests/rank_integration_test.py +++ b/tzrec/tests/rank_integration_test.py @@ -1162,6 +1162,22 @@ def test_multi_tower_din_with_dynamicemb_train_eval(self): # ) self.assertTrue(self.success) + @unittest.skipIf( + gpu_unavailable[0] or not dynamicemb_util.has_dynamicemb, + "dynamicemb not available.", + ) + @mark_ci_scope("gpu") + def test_multi_tower_din_with_dynamicemb_ftrl_train_eval(self): + self.success = utils.test_train_eval( + "tzrec/tests/configs/multi_tower_din_dynamicemb_ftrl_mock.config", + self.test_dir, + ) + if self.success: + self.success = utils.test_eval( + os.path.join(self.test_dir, "pipeline.config"), self.test_dir + ) + self.assertTrue(self.success) + @unittest.skipIf( gpu_unavailable[0] or not trt_utils.has_tensorrt, "tensorrt not available." ) diff --git a/tzrec/utils/dynamicemb_util.py b/tzrec/utils/dynamicemb_util.py index 7c26fa4a..553031be 100644 --- a/tzrec/utils/dynamicemb_util.py +++ b/tzrec/utils/dynamicemb_util.py @@ -45,7 +45,7 @@ ) from torchrec.modules.embedding_configs import BaseEmbeddingConfig -from tzrec.optim.optimizer import sparse_init_accumulator_value +from tzrec.optim.optimizer import FTRL, sparse_init_accumulator_value from tzrec.protos import feature_pb2 from tzrec.utils.logging_util import logger @@ -207,6 +207,28 @@ def _log_dynamicemb_table_plan( ) +def _get_optimizer_multipler( + optimizer_class: Optional[Type[torch.optim.Optimizer]], shape: torch.Size +) -> float: + """Optimizer state size per embedding element, including dynamicemb's FTRL. + + torchrec's table maps any class it does not know to 1, and FTRL keeps a + linear and an accumulator term per element, so its rows are ``2 * dim`` wide + like Adam's. + + Args: + optimizer_class (type, optional): in-backward optimizer class of the + table, None when the table is not trained. + shape (torch.Size): unsharded table shape, used by row-wise optimizers. + + Returns: + the multiplier applied to the embedding width. + """ + if optimizer_class is FTRL: + return 2.0 + return shard_estimators._get_optimizer_multipler(optimizer_class, shape) + + has_dynamicemb = False try: import dynamicemb @@ -564,7 +586,7 @@ def _to_sharding_plan( # calc local_hbm_for_values tensor = sharding_option.tensor optimizer_class = getattr(tensor, "_optimizer_classes", [None])[0] - optimizer_multipler = shard_estimators._get_optimizer_multipler( + optimizer_multipler = _get_optimizer_multipler( optimizer_class, tensor.shape ) dynamicemb_options.training = optimizer_class is not None @@ -610,6 +632,22 @@ def _to_sharding_plan( ddr_bytes=int(shards[0].storage.ddr), ) else: + # A data_parallel table is replicated and updated by the dense + # optimizer, so the sparse optimizer never reaches it. + if ( + getattr(sharding_option.tensor, "_optimizer_classes", [None])[0] + is FTRL + and sharding_type != ShardingType.DATA_PARALLEL.value + ): + raise ValueError( + "sparse ftrl_optimizer only supports dynamicemb embedding " + "tables, but table[" + f"{sharding_option.path}.{sharding_option.name}] is planned " + f"with sharding_type[{sharding_type}] and " + f"compute_kernel[{sharding_option.compute_kernel}]. " + "Set `dynamicemb { }` on every sparse feature, or use " + "another sparse optimizer." + ) module_plan[sharding_option.name] = ParameterSharding( sharding_spec=sharding_spec, sharding_type=sharding_type, @@ -732,7 +770,7 @@ def _calculate_dynamicemb_storage_specific_sizes( optimizer_multipler = 0.0 optimizer_class = getattr(tensor, "_optimizer_classes", [None])[0] if not is_inference: - optimizer_multipler = shard_estimators._get_optimizer_multipler( + optimizer_multipler = _get_optimizer_multipler( optimizer_class, tensor.shape ) diff --git a/tzrec/utils/dynamicemb_util_test.py b/tzrec/utils/dynamicemb_util_test.py index c5bd540c..e101b0e7 100644 --- a/tzrec/utils/dynamicemb_util_test.py +++ b/tzrec/utils/dynamicemb_util_test.py @@ -20,7 +20,9 @@ from torchrec.distributed.embedding_types import EmbeddingComputeKernel from torchrec.distributed.types import ShardingType from torchrec.modules.embedding_configs import EmbeddingBagConfig +from torchrec.optim import optimizers, rowwise_adagrad +from tzrec.optim.optimizer import FTRL from tzrec.protos import feature_pb2 from tzrec.utils import dynamicemb_util from tzrec.utils.test_util import mark_ci_scope, parameterized_name_func @@ -154,6 +156,32 @@ def test_only_values_drops_metadata(self): self.assertGreater(with_meta, without_meta) +class OptimizerMultiplerTest(unittest.TestCase): + """Per-element optimizer state width used to size dynamicemb values.""" + + @parameterized.expand( + [ + param("untrained", optimizer_class=None, expected=0.0), + param("sgd", optimizer_class=optimizers.SGD, expected=0), + param("adam", optimizer_class=optimizers.Adam, expected=2), + param( + "rowwise_adagrad", + optimizer_class=rowwise_adagrad.RowWiseAdagrad, + expected=1 / 16, + ), + param("ftrl", optimizer_class=FTRL, expected=2.0), + ], + name_func=parameterized_name_func, + ) + def test_multipler(self, _name, optimizer_class, expected): + self.assertAlmostEqual( + dynamicemb_util._get_optimizer_multipler( + optimizer_class, torch.Size([1024, 16]) + ), + expected, + ) + + @unittest.skipUnless( dynamicemb_util.has_dynamicemb, "dynamicemb is not installed; skipping." ) diff --git a/tzrec/utils/plan_util_test.py b/tzrec/utils/plan_util_test.py index 24a5ea79..05d906e9 100644 --- a/tzrec/utils/plan_util_test.py +++ b/tzrec/utils/plan_util_test.py @@ -472,14 +472,25 @@ def _build_constraint(self, max_capacity=4096): dynamicemb_options=opts, ) - def _build_model(self): - table = EmbeddingBagConfig( - num_embeddings=4096, - embedding_dim=32, - name="table_de", - feature_names=["feat_de"], - ) - return TestSparseNN(tables=[table], sparse_device=torch.device("meta")) + def _build_model(self, with_plain_table=False): + tables = [ + EmbeddingBagConfig( + num_embeddings=4096, + embedding_dim=32, + name="table_de", + feature_names=["feat_de"], + ) + ] + if with_plain_table: + tables.append( + EmbeddingBagConfig( + num_embeddings=4096, + embedding_dim=32, + name="table_plain", + feature_names=["feat_plain"], + ) + ) + return TestSparseNN(tables=tables, sparse_device=torch.device("meta")) def test_enumerate_yields_both_modes_and_all_factors(self): from tzrec.utils.plan_util import ( @@ -609,6 +620,95 @@ def test_sharding_plan_carries_initial_accumulator_value(self): DynamicEmbParameterSharding.pop_additional_fused_params(fused_params) self.assertIn("initial_accumulator_value", fused_params) + def _enumerate(self, with_plain_table=False): + from tzrec.utils.plan_util import ( + EmbeddingEnumerator as _TzrecEmbeddingEnumerator, + ) + from tzrec.utils.plan_util import ( + get_default_sharders as _tzrec_get_default_sharders, + ) + + model = self._build_model(with_plain_table=with_plain_table) + topology = Topology(world_size=2, compute_device="cuda") + enumerator = _TzrecEmbeddingEnumerator( + topology=topology, + batch_size=128, + fqn_constraints={"sparse.ebc.table_de": self._build_constraint()}, + ) + search_space = enumerator.enumerate( + module=model, sharders=_tzrec_get_default_sharders() + ) + return search_space, topology + + def test_sharding_plan_sizes_ftrl_optimizer_state(self): + import inspect + + from dynamicemb.batched_dynamicemb_tables import BatchedDynamicEmbeddingTablesV2 + from torchrec.distributed.planner import planners + + from tzrec.optim.optimizer import FTRL + from tzrec.utils.dynamicemb_util import ( + _calculate_dynamicemb_table_storage_specific_size, + ) + + # FTRL's knobs travel as plain kwargs of the dynamicemb table module. + signature = inspect.signature(BatchedDynamicEmbeddingTablesV2.__init__) + for name in ("learning_rate_power", "ftrl_beta", "l1_reg", "l2_reg"): + self.assertIn(name, signature.parameters) + + search_space, topology = self._enumerate() + sharding_option = next( + so for so in search_space if getattr(so, "use_dynamicemb", False) + ) + for rank, shard in enumerate(sharding_option.shards): + shard.rank = rank + sharding_option.tensor._optimizer_classes = [FTRL] + + plan = planners.to_sharding_plan([sharding_option], topology) + + param_sharding = plan.plan[sharding_option.path][sharding_option.name] + options = param_sharding.dynamicemb_options + self.assertTrue(options.training) + # FTRL keeps `linear` and `accum` per element, so the values are sized + # as if the row were three times the embedding width. + self.assertEqual( + options.local_hbm_for_values, + _calculate_dynamicemb_table_storage_specific_size( + sharding_option.shards[0].size, + sharding_option.tensor.element_size(), + 2.0, + sharding_option.cache_load_factor, + is_hbm=True, + only_values=True, + bucket_capacity=options.bucket_capacity, + ), + ) + + def test_sharding_plan_rejects_ftrl_on_non_dynamicemb_table(self): + from torchrec.distributed.planner import planners + + from tzrec.optim.optimizer import FTRL + + search_space, topology = self._enumerate(with_plain_table=True) + for sharding_type, raises in ( + (ShardingType.ROW_WISE.value, True), + (ShardingType.DATA_PARALLEL.value, False), + ): + sharding_option = next( + so + for so in search_space + if not getattr(so, "use_dynamicemb", False) + and so.sharding_type == sharding_type + ) + for rank, shard in enumerate(sharding_option.shards): + shard.rank = rank + sharding_option.tensor._optimizer_classes = [FTRL] + if raises: + with self.assertRaisesRegex(ValueError, "ftrl_optimizer"): + planners.to_sharding_plan([sharding_option], topology) + else: + planners.to_sharding_plan([sharding_option], topology) + if __name__ == "__main__": unittest.main()