This document consists of two parts, one is the introduction to the API, which can be used as a reference, and the other part prompts specific functions from a functional perspective and which interfaces are related to them.
- DynamicEmbParameterConstraints
- DynamicEmbeddingEnumerator
- DynamicEmbeddingShardingPlanner
- Sharding planner
- DynamicEmbeddingCollectionSharder
- DynamicEmbCheckMode
- DynamicEmbInitializerMode
- DynamicEmbInitializerArgs
- DynamicEmbPoolingMode
- DynamicEmbTableOptions
- DynamicEmbDump
- DynamicEmbLoad
- incremental_dump
- replay_increment
- pop_evicted_keys
- pop_erased_keys
- get_score
- set_score
- Counter
- AdmissionStrategy and MultiTableAdmitter
The DynamicEmbParameterConstraints function inherits from TorchREC's ParameterConstraints function. It has the same basic parameters as ParameterConstraints and adds members for specifying and configuring dynamic embedding tables. This function serves as the only entry point for users to specify whether an embedding table is a dynamic embedding table..
```python
#How to import
from dynamicemb.planner import DynamicEmbParameterConstraints
#API arguments
@dataclass
class DynamicEmbParameterConstraints(ParameterConstraints):
"""
DynamicEmb-specific parameter constraints that extend ParameterConstraints.
Attributes
----------
use_dynamicemb : Optional[bool]
A flag indicating whether to use DynamicEmb storage. Defaults to False.
dynamicemb_options : Optional[DynamicEmbTableOptions]
Configuration for the dynamic embedding table, including initializer args.
Common choices include "uniform", "normal", etc. Defaults to "uniform".
"""
use_dynamicemb: Optional[bool] = False
dynamicemb_options: Optional[DynamicEmbTableOptions] = DynamicEmbTableOptions()
```
The DynamicEmbeddingEnumerator function inherits from TorchREC's EmbeddingEnumerator function and its usage is exactly the same as EmbeddingEnumerator. This class differentiates between TorchREC's embedding tables and dynamic embedding tables during enumeration in the sharding plan.
```python
#How to import
from dynamicemb.planner import DynamicEmbeddingEnumerator
#API arguments
class DynamicEmbeddingEnumerator(EmbeddingEnumerator):
def __init__(
self,
topology: Topology,
batch_size: Optional[int] = BATCH_SIZE,
constraints: Optional[Dict[str, DynamicEmbParameterConstraints]] = None,
estimator: Optional[Union[ShardEstimator, List[ShardEstimator]]] = None,
) -> None:
"""
DynamicEmbeddingEnumerator extends the EmbeddingEnumerator to handle dynamic embedding tables.
Parameters
----------
topology : Topology
The topology of the GPU and Host memory.
batch_size : Optional[int], optional
The batch size for training. Defaults to BATCH_SIZE.
The creation and usage are consistent with the same types in TorchREC.
constraints : Optional[Dict[str, DynamicEmbParameterConstraints]], optional
A dictionary of constraints for the parameters. Defaults to None.
estimator : Optional[Union[ShardEstimator, List[ShardEstimator]]], optional
An estimator or a list of estimators for estimating shard sizes. Defaults to None.
The creation and usage are consistent with the same types in TorchREC.
"""
```
Wrapped TorchREC's EmbeddingShardingPlanner to perform sharding for dynamic embedding tables. Unlike EmbeddingShardingPlanner, it requires an additional eb_configs argument so DynamicEmb can derive per-rank capacities and bucket widths from the global EmbeddingConfig and the process-group world size.
On construction it runs the internal preparation step described under Sharding planner, then builds the TorchREC sub-planner and DynamicEmb shard metadata as before.
For row-wise DynamicEmb sharding, the supported dist_type values are:
continuousroundrobinhash_roundrobin
hash_roundrobin hashes the raw key before rank assignment and is intended to reduce sensitivity to pathological raw-key patterns that can break plain modulo-based roundrobin. It is an opt-in routing mode; the default remains roundrobin for compatibility with existing checkpoints. It should be understood as a routing robustness improvement, not as a general solution to arbitrary hot-key or Zipf-skew load balancing.
```python
#How to import
from dynamicemb.planner import DynamicEmbeddingShardingPlanner
#API arguments
class DynamicEmbeddingShardingPlanner:
def __init__(self,
eb_configs: List[BaseEmbeddingConfig],
topology: Optional[Topology] = None,
batch_size: Optional[int] = None,
enumerator: Optional[Enumerator] = None,
storage_reservation: Optional[StorageReservation] = None,
proposer: Optional[Union[Proposer, List[Proposer]]] = None,
partitioner: Optional[Partitioner] = None,
performance_model: Optional[PerfModel] = None,
stats: Optional[Union[Stats, List[Stats]]] = None,
constraints: Optional[Dict[str, DynamicEmbParameterConstraints]] = None,
debug: bool = True):
"""
DynamicEmbeddingShardingPlanner wraps EmbeddingShardingPlanner and adds `eb_configs` (TorchREC
table configs) so per-rank DynamicEmb options can be filled before planning. See the
"Sharding planner" section in DynamicEmb_APIs.md for how `DynamicEmbTableOptions` are adjusted.
Parameters
----------
eb_configs : List[BaseEmbeddingConfig]
A list of TorchREC BaseEmbeddingConfig in the TorchREC model
topology : Optional[Topology], optional
The topology of GPU and Host memory. If None, a default topology will be created. Defaults to None.
The creation and usage are consistent with the same types in TorchREC.
Note: The memory budget does not include the consumption of dynamicemb.
batch_size : Optional[int], optional
The batch size for training. Defaults to None, will set 512 in Planner.
enumerator : Optional[Enumerator], optional
An enumerator for sharding. Defaults to None.
The creation and usage are consistent with the same types in TorchREC.
storage_reservation : Optional[StorageReservation], optional
Storage reservation details. Defaults to None.
The creation and usage are consistent with the same types in TorchREC.
proposer : Optional[Union[Proposer, List[Proposer]]], optional
A proposer or a list of proposers for proposing sharding plans. Defaults to None.
The creation and usage are consistent with the same types in TorchREC.
partitioner : Optional[Partitioner], optional
A partitioner for partitioning the embedding tables. Defaults to None.
The creation and usage are consistent with the same types in TorchREC.
performance_model : Optional[PerfModel], optional
A performance model for evaluating sharding plans. Defaults to None.
The creation and usage are consistent with the same types in TorchREC.
stats : Optional[Union[Stats, List[Stats]]], optional
Statistics or a list of statistics for the sharding process. Defaults to None.
The creation and usage are consistent with the same types in TorchREC.
constraints : Optional[Dict[str, DynamicEmbParameterConstraints]], optional
A dictionary of constraints for every TorchREC embedding table and Dynamic embedding table. Defaults to None.
debug : bool, optional
A flag indicating whether to enable debug mode. Defaults to True.
"""
```
When you construct DynamicEmbeddingShardingPlanner, the implementation first validates that constraints and eb_configs are consistent (every EmbeddingCollection / table name appears exactly once in eb_configs, matches the keys of constraints, and there are no extra keys). Then, for each table with DynamicEmbParameterConstraints.use_dynamicemb == True, it updates that table’s DynamicEmbTableOptions in dynamicemb_options via the internal routine _prepare_dynemb_table_options (order matters):
| Step | Field(s) | What happens |
|---|---|---|
| 1 | initializer_args |
complete_initializer_args returns a new DynamicEmbInitializerArgs when needed. For UNIFORM initialization only: if lower or upper is None, they are filled. With a TorchREC embedding_config, bounds are ±sqrt(1 / num_embeddings); without it, 0.0 and 1.0. Other modes are returned unchanged. |
| 2 | bucket_capacity, max_capacity |
_sharded_table_bucket_layout(embedding_config, world_size, bucket_capacity) (internal) returns (num_buckets, effective_bucket_width) per rank. The planner overwrites bucket_capacity with the effective width (after MAX_BUCKET_CAPACITY / alignment rules). max_capacity is set to num_buckets * effective_bucket_width, i.e. the same value as get_sharded_table_capacity(embedding_config, world_size, bucket_capacity). init_capacity: if unset, set to max_capacity; if set, align to bucket_capacity, then clamp to max_capacity if larger. User input: bucket_capacity on DynamicEmbTableOptions (multiple of BUCKET_ALIGNMENT (16) unless MAX_BUCKET_CAPACITY = 2**63 - 1). Sentinel: layout is (1, aligned_per_rank_rows) — one bucket spanning the shard. Otherwise: num_buckets = align_to_table_size(ceil(N/world), bucket_capacity) // bucket_capacity. |
| 3 | local_hbm_for_values |
Overwritten to ceil(global_hbm_for_values / world_size) so each rank gets an equal byte budget from the user-provided global_hbm_for_values (set on DynamicEmbTableOptions before planning). |
User-supplied values that should be set before planning (typical DMP path) include at least:
global_hbm_for_values: global HBM byte budget for the table’s values; the planner only splits it across ranks.bucket_capacity: requested hashtable bucket size in rows (see rules above), orMAX_BUCKET_CAPACITYfor “single bucket per rank table”.initializer_args: optional partial bounds forUNIFORM; missing bounds are completed as in step 1.
Downstream (after planning, when modules are built): the batched embedding path uses max_capacity from DynamicEmbTableOptions (e.g. allocation size and consistency checks against TorchREC shard row counts in _get_dynamicemb_options_per_table in batched_dynamicemb_compute_kernel.py).
Public helpers (same rules as the planner):
from dynamicemb.dynamicemb_config import complete_initializer_args(initializer completion; not re-exported fromdynamicembtop-level today)from dynamicemb import get_sharded_table_capacity— returns per-rank row capacity after sharding and bucket alignment (num_buckets * effective_bucket_width), matchingmax_capacityset by the planner. WithMAX_BUCKET_CAPACITY, the table is one bucket per rank whose width is the aligned shard row count.from dynamicemb import get_table_value_bytes— total bytes for all ranks’ value storage (embedding + optimizer state rows), using the same row layout asget_sharded_table_capacityfor the givenbucket_capacity(includingMAX_BUCKET_CAPACITY). Use this to sizeglobal_hbm_for_valuesbefore planning; apply your own caching fraction or HBM budget scale on top if needed (as in benchmarks / examples).from dynamicemb import BUCKET_ALIGNMENT, MAX_BUCKET_CAPACITY
Inherits from TorchREC's EmbeddingCollectionSharder and is used in exactly the same way. This API mainly overrides the deduplication process of indices through inheritance, making it compatible with dynamic embedding tables.
```python
#How to import
from dynamicemb.shard import DynamicEmbeddingCollectionSharder
```
When the dynamic embedding table capacity is small, a single feature with a large number of indices may lead to issues where the hashtable cannot insert the indices. Enabling safe check allows you to observe this behavior (the number of times indices cannot be inserted and the number of indices each time) to determine if the dynamic embedding table capacity is set too low.
```python
#How to import
from dynamicemb import DynamicEmbCheckMode
#API arguments
class DynamicEmbCheckMode(enum.IntEnum):
"""
Enumeration for different modes of checking dynamic embedding's insertion behaviors.
DynamicEmb uses a hashtable as the backend. If the embedding table capacity is small and the number of indices in a single feature is large,
it is easy for too many indices to be allocated to the same hash table bucket in one lookup, resulting in the inability to insert indices into the hashtable.
DynamicEmb resolves this issue by setting the lookup results of indices that cannot be inserted to 0.
Fortunately, in a hashtable with a large capacity, such insertion failures are very rare and almost never occur.
This issue is more frequent in hashtables with small capacities, which can affect training accuracy.
Therefore, we do not recommend using dynamic embedding tables for very small embedding tables.
To prevent this behavior from affecting training without user awareness, DynamicEmb provides a safe check mode.
Users can set whether to enable safe check when configuring DynamicEmbTableOptions.
Enabling safe check will add some overhead, but it can provide insights into whether the hash table frequently fails to insert indices.
If the number of insertion failures is high and the proportion of affected indices is large,
it is recommended to either increase the dynamic embedding capacity or avoid using dynamic embedding tables for small embedding tables.
Attributes
----------
ERROR : int
When there are indices that can't be inserted successfully:
This mode will throw a runtime error indicating how many indices failed to insert.
The program will crash.
WARNING : int
When there are indices that can't be inserted successfully:
This mode will give a warning about how many indices failed to insert.
The program will continue. For uninserted indices, their embeddings' values will be set to 0.0.
IGNORE : int
Don't check whether insertion is successful or not, therefore it doesn't bring additional checking overhead.
For uninserted indices, their embeddings' values will be set to 0.0 silently.
"""
ERROR = 0
WARNING = 1
IGNORE = 2
```
The initialization method for each embedding vector in the dynamic embedding table currently supports random UNIFORM distribution, random NORMAL distribution, and non-random constant initialization. The default distribution is UNIFORM.
```python
#How to import
from dynamicemb import DynamicEmbInitializerMode
#API arguments
class DynamicEmbInitializerMode(enum.Enum):
"""
Enumeration for different modes of initializing dynamic embedding vector values.
Attributes
----------
NORMAL : str
Normal Distribution.
UNIFORM : str
Uniform distribution of random values.
CONSTANT : str
All dynamic embedding vector values are a given constant.
DEBUG : str
Debug value generation mode for testing.
"""
NORMAL = "normal"
TRUNCATED_NORMAL = "truncated_normal"
UNIFORM = "uniform"
CONSTANT = "constant"
DEBUG = "debug"
```
Parameters for each random initialization method in DynamicEmbInitializerMode.
```python
#How to import
from dynamicemb import DynamicEmbInitializerArgs
#API arguments
@dataclass
class DynamicEmbInitializerArgs:
"""
Arguments for initializing dynamic embedding vector values.
Attributes
----------
mode : DynamicEmbInitializerMode
The mode of initialization, one of the DynamicEmbInitializerMode values.
mean : float, optional
The mean value for normal distributions. Defaults to 0.0.
std_dev : float, optional
The standard deviation for normal and distributions. Defaults to 1.0.
lower : float, optional
The lower bound for uniform distribution. Defaults to 0.0.
upper : float, optional
The upper bound for uniform distribution. Defaults to 1.0.
value : float, optional
The constant value for constant initialization. Defaults to 0.0.
"""
mode: DynamicEmbInitializerMode
mean: float = 0.0
std_dev: float = 1.0
lower: float = None
upper: float = None
value: float = 0.0
```
The storage space is limited, but the value range of sparse features is relatively large, so dynamicemb introduces the concept of score to perform customized eviction of sparse features within the limited storage space. dynamicemb provides the following strategies to set the score.
```python
#How to import
from dynamicemb import DynamicEmbScoreStrategy
#API arguments
class DynamicEmbScoreStrategy(enum.IntEnum):
"""
Enumeration for different modes to set index-embedding's score.
The index-embedding pair with smaller scores will be more likely to be evicted from the embedding table when the table is full.
dynamicemb allows configuring scores by table.
For a table, the scores in the subsequent forward passes are larger than those in the previous ones for modes TIMESTAMP and STEP.
Users can also provide customized score(mode CUSTOMIZED) for each table's forward pass.
Attributes
----------
TIMESTAMP:
In a forward pass, embedding table's scores will be set to global nanosecond timer of device, and due to the timing of GPU scheduling,
different scores may have slight differences.
Users must not set scores under TIMESTAMP mode.
STEP:
Each embedding table has a member `step` which will increment for every forward pass.
All scores in each forward pass are the same which is step's value.
Users must not set scores under STEP mode.
CUSTOMIZED:
Each embedding table's score are managed by users.
Users have to set the score before every forward pass using `set_score` interface.
LFU:
If there are not enough slots inside the bucket to store new keys, the least used key in the bucket will be evicted.
NO_EVICTION:
The table’s capacity doubles whenever there are not enough slots for new keys, and this continues until available memory is exhausted.
When the memory resources are insufficient, there will be a warning message, and training can continue but the accuracy of eviction cannot be guaranteed.
"""
TIMESTAMP = 0
STEP = 1
CUSTOMIZED = 2
LFU = 3
NO_EVICTION = 4
```
Users can specify the `DynamicEmbScoreStrategy` using `score_strategy` in `DynamicEmbTableOptions` per table.
A table normally keeps one score per key, but a compound score_strategy
given as a tuple keeps more than one. The supported compound
(DynamicEmbScoreStrategy.TIMESTAMP, DynamicEmbScoreStrategy.LFU) keeps two
scores per key — a last-access timestamp and an access frequency — so the table
can rank eviction by frequency while still supporting a time-based
incremental_dump.
The order you write the tuple in does not change eviction behavior:
(TIMESTAMP, LFU) and (LFU, TIMESTAMP) evict identically. It only controls the
order of the score columns as you see them, in two places that always stay
consistent with each other:
- Scores dumped to file (
DynamicEmbDump/load): each key's scores are written in the tuple order.(TIMESTAMP, LFU)writes[timestamp, frequency]per key;(LFU, TIMESTAMP)writes[frequency, timestamp].loadreads them back in the same order, so a checkpoint is interchangeable only between tables configured with the same tuple order. score_functionindexing: inside ascore_function,scores[0],scores[1], ... follow the same tuple order. For(TIMESTAMP, LFU),scores[0]is the timestamp andscores[1]the frequency; for(LFU, TIMESTAMP)the two are swapped.
Choose the order that matches how you want checkpoint columns laid out and how you
index scores in your score_function — the two always agree.
LFU alone evicts by raw access frequency. For finer control — e.g. a
time-decayed LFU that also favors recently-used keys — configure a table with the
compound score strategy (DynamicEmbScoreStrategy.TIMESTAMP, DynamicEmbScoreStrategy.LFU) (either order) and pass a Python score_function in
DynamicEmbTableOptions. When a bucket is full, the key(s) with the lowest
returned score are evicted.
Signature: score_function(scores, cur_timestamp) -> float64
scoresindexes the table's two scores in the same order as yourscore_strategytuple (see "Compound (multi-score) strategies and score order" above). For(TIMESTAMP, LFU):scores[0]is the last-access timestamp (device nanosecond timer) andscores[1]is the access frequency; for(LFU, TIMESTAMP)the two are swapped.cur_timestampis the current device nanosecond timer at eviction time.- The function is compiled to device code with numba-cuda, so it must be
numba-compilable: index
scoreswith integer constants only, and return a numeric value (the return is treated asfloat64, so an integer expression is cast). It should return a finite value; aNaNreturn is treated as evict-first. Any decay constant is written directly in the body. Ifscore_functionis omitted, eviction uses the default frequency-ranked evictor (ties broken by the older timestamp).
Constraints: score_function is supported only for the two-word compound
{TIMESTAMP, LFU} strategy — exactly one TIMESTAMP column plus one LFU column,
so scores always has length 2 (scores[0]/scores[1]). It is not
available for single strategies (TIMESTAMP, STEP, LFU, CUSTOMIZED,
NO_EVICTION) or any other score combination or count. Setting score_function
with a non-compound score_strategy raises ValueError. numba-cuda must be
installed.
import math
from dynamicemb import DynamicEmbScoreStrategy, DynamicEmbTableOptions
def lru_lfu_decay_score(scores, cur_timestamp):
# scores[0] = last-access timestamp (ns), scores[1] = access frequency.
# Time-decayed LFU: reward frequency, penalize staleness (~0.9 per second).
age_seconds = (cur_timestamp - scores[0]) * 1e-9
return math.log(max(scores[1], 1)) + age_seconds * math.log(0.9)
table_options = DynamicEmbTableOptions(
score_strategy=(
DynamicEmbScoreStrategy.TIMESTAMP,
DynamicEmbScoreStrategy.LFU,
),
score_function=lru_lfu_decay_score,
)DynamicEmb supports three pooling modes that determine how embedding lookups are aggregated. These modes correspond to how EmbeddingCollection (sequence) and EmbeddingBagCollection (pooled) work in TorchREC.
All pooling modes use fused CUDA kernels for both forward and backward passes. Tables with different embedding dimensions (mixed-D) are fully supported in SUM and MEAN modes.
```python
#How to import
from dynamicemb import DynamicEmbPoolingMode
#API arguments
@enum.unique
class DynamicEmbPoolingMode(enum.IntEnum):
"""
Enumeration for pooling modes in dynamic embedding lookup.
The values are taken from the bound C++ enum ``dyn_emb::PoolingMode``
(``src/utils.h``), which is the single source of truth, in the same way
``DynamicEmbEvictStrategy`` takes its values from ``EvictStrategy``. It
is an ``IntEnum`` because a pooling mode is handed straight to the
kernels and compared as an integer, so it stays usable wherever a plain
``int`` was before (``SUM == 0``, dict keys, JSON).
Attributes
----------
SUM : int
Sum pooling. For each sample, the embeddings of all indices in the bag
are summed. Output shape: (batch_size, total_D) where total_D is the
sum of embedding dimensions across all features.
MEAN : int
Mean pooling. For each sample, the embeddings of all indices in the bag
are averaged. Output shape: same as SUM.
NONE : int
No pooling (sequence mode). Each index produces its own embedding row.
Output shape: (total_indices, D).
"""
SUM = BagPoolingMode.KSum # 0
MEAN = BagPoolingMode.KMean # 1
NONE = BagPoolingMode.KNone # 2
```
Weighted SUM pooling — SUM pooling supports optional per-position float32 weights: out[b] = Σ wᵢ · embᵢ. Supported in both training and
eval mode, only for SUM (no weighted-MEAN variant), and works with mixed-D tables. Two ways to use it:
-
Through TorchRec's
EmbeddingBagCollection: build the collection withis_weighted=Trueand attach the weights to the features KJT (KeyedJaggedTensor(keys=..., values=indices, weights=weights, lengths=...)); callmodel(features)as usual. The weights ride the KJT through the all2all distribution and reach the DynamicEmb lookup as pooling weights.is_weighted=Trueis required, not merely conventional. DynamicEmb only reads the KJT weights as pooling weights when the collection declares itself weighted; on anis_weighted=Falsecollection the weights are ignored and the result is a plain unweighted sum, with no error raised. The weights channel of aKeyedJaggedTensoris shared — TorchRec also uses it to carry per-feature scores for virtual-table eviction, and the sequence (EmbeddingCollection) path reads it as LFU frequency counters — so the flag is what disambiguates them. -
Through
BatchedDynamicEmbeddingTablesV2.forwarddirectly:module(indices, offsets, pooling_weights=w)wherewis a float32 tensor aligned withindices(w.numel() == indices.numel()).
The weights are not differentiable. No gradient is computed for them, so learned per-position weights are not supported: passing a tensor with
requires_grad=True while grad mode is enabled raises ValueError rather than letting the producing module train against a gradient that is
silently zero. Weights that carry requires_grad are still accepted under torch.no_grad(), which is what keeps eval working.
Weighted pooling raises ValueError when: pooling_mode != SUM, the weights are not float32, weights.numel() != indices.numel(), the weights
require grad while grad mode is enabled, or frequency_counters is passed alongside them (both are carried by the single KJT weights channel, so
a caller has to pick one).
The optimizer entry in fused_params normally takes an FBGEMM EmbOptimType. DynamicEmb also implements
optimizers that EmbOptimType has no member for; those live in DynamicEmbOptimType and are accepted
anywhere an optimizer type is. The two enums never compare equal, so an EmbOptimType check elsewhere
cannot accidentally match one of them.
```python
#How to import
from dynamicemb import DynamicEmbOptimType
@enum.unique
class DynamicEmbOptimType(enum.Enum):
FTRL = "ftrl"
```
FTRL-Proximal, Algorithm 1 of Ad Click Prediction: a View from the Trenches (McMahan et al., KDD 2013), with the paper's fixed square root generalized to an arbitrary exponent. Per coordinate:
new_accum = accum + grad²
linear += grad − (new_accum^(−p) − accum^(−p)) / lr · weight p = learning_rate_power
weight = 0 if |linear| ≤ l1_reg
= (sign(linear)·l1_reg − linear)
/ ((ftrl_beta + new_accum^(−p)) / lr + l2_reg) otherwise
accum = new_accum
lr, ftrl_beta, l1_reg and l2_reg are the paper's α, β, λ1 and λ2; p = -0.5 makes n^(−p) a
square root and so recovers its α / (β + √n) learning rate exactly. Pass them through fused_params alongside optimizer:
| Parameter | Default | Meaning |
|---|---|---|
learning_rate |
0.01 |
α. Must be positive -- FTRL divides by it. |
learning_rate_power |
-0.5 |
Exponent on the accumulator in the learning rate, so it must be <= 0: negative decays the rate, 0 holds it fixed, and a positive value would make it grow without bound (rejected). -0.5 is the paper's own choice and takes a faster kernel path; other values follow TensorFlow's generalization. |
ftrl_beta |
0.0 |
β. Keeps the per-coordinate learning rate finite while accum is still small -- the paper's way of bounding the first steps. |
initial_accumulator_value |
0.0 |
Seeds accum; linear always starts at 0. Read the note below before setting it non-zero. |
l1_reg |
0.0 |
λ1. |
l2_reg |
0.0 |
λ2. |
Two behaviours worth knowing before choosing FTRL:
-
Per-row state is
2 × embedding_dim, laid out aslinearthenaccumright after the embedding -- the same shape of layout Adam uses for itsmandv. Checkpoints store it at full width. -
The weight is re-solved from
(linear, accum)each step rather than nudged from its previous value. That is what letsl1_regdrive a weight to exactly zero instead of merely shrinking it.It is also why seeding the state deserves care here. FTRL was written for linear regression, whose weights start at zero; an embedding's do not, and a state of
linear = 0is inconsistent with the weight already sitting in the row. The first update reconciles them, keeping a fraction(√n₁ − √n₀) / (ftrl_beta + √n₁) n₀ = initial_accumulator_value, n₁ = n₀ + g²of the initializer. Both knobs that bound the early steps cost retention, and for the small gradients typical of embeddings they cost nearly all of it:
n0ftrl_betag = 1.0g = 0.01g = 0.00010.00.0100% 100% 100% 0.10.069.8% 0.05% 0.000005% 0.01.050% 0.99% 0.01% 0.11.035.8% 0.012% 0.0000012% Only
initial_accumulator_value = 0together withftrl_beta = 0keeps the initializer whole. That combination has its own cost: withn₀ = 0the first step isw − α·sign(g), a full α-sized move however small the gradient. Which matters more is a modelling choice -- if the initializer is doing real work, keep both at 0 and pick α accordingly; if the early steps need damping, expect the initializer to be mostly overwritten and size it as if rows started near zero.
Unlike the EmbOptimType optimizers, FTRL has no FBGEMM counterpart, so construct_twin_module cannot
build a TorchRec twin for a model that uses it.
Per-table configuration for dynamic embedding, passed into DynamicEmbParameterConstraints as dynamicemb_options. The authoritative definition lives in dynamicemb.dynamicemb_config.DynamicEmbTableOptions (this section mirrors its docstring).
Fields declared first (through device_id) are planner/runtime-heavy: DynamicEmbeddingShardingPlanner fills them via _prepare_dynemb_table_options together with internal _sharded_table_bucket_layout (and thus the same per-rank row count as get_sharded_table_capacity). User-facing knobs such as training, bucket_capacity, global_hbm_for_values, and initializer_args follow. Hash-table scores are driven by score_strategy and kernels, not by a separate score-dtype field on this dataclass.
```python
#How to import
from dynamicemb import DynamicEmbTableOptions
#API arguments
@dataclass
class DynamicEmbTableOptions:
"""
Encapsulates the configuration options for dynamic embedding table.
This class includes parameters that control the behavior and performance of the embedding lookup module, specifically tailored for dynamic embeddings.
`get_grouped_key` will return fields used to group dynamic embedding tables.
Fields listed first (through ``device_id``) are often filled by the planner or runtime rather than
being the main user configuration knobs. Score handling for the hash table follows score
policies and kernels, not a separate score-dtype field on this dataclass.
Parameters
----------
embedding_dtype : Optional[torch.dtype], optional
Data (weight) type of dynamic embedding table. Also the precision a checkpoint stores:
``DynamicEmbDump`` writes value files at this dtype and records it in the table's meta JSON as
``embedding_dtype`` (alongside ``embedding_dim`` / ``optim_state_dtype``), which is how
``DynamicEmbLoad`` reads them back. Loading a checkpoint whose precision differs converts the
values and warns; a differing dim is an error.
dim : Optional[int], optional
Value vector dimension. With ``DynamicEmbeddingShardingPlanner``, ``_prepare_dynemb_table_options``
sets it from ``BaseEmbeddingConfig.embedding_dim``. The embedding kernel only warns if it
differs from the sharded ``local_cols`` (see ``_get_dynamicemb_options_per_table``).
max_capacity : Optional[int], optional
Per-shard maximum table rows on one GPU. With ``DynamicEmbeddingShardingPlanner``,
``_prepare_dynemb_table_options`` sets ``max_capacity`` to
per-rank row count from ``get_sharded_table_capacity``.
You may omit ``max_capacity`` on ``DynamicEmbTableOptions`` before planning and let the planner set it.
If ``init_capacity`` is unset it becomes ``max_capacity``; if set and aligned,
it is clamped to at most ``max_capacity``.
The embedding kernel checks consistency with TorchREC shard metadata (see
``_get_dynamicemb_options_per_table``).
evict_strategy : DynamicEmbEvictStrategy
Strategy used for evicting entries when the table exceeds its capacity.
Default is ``DynamicEmbEvictStrategy.LRU``.
local_hbm_for_values : int
High-bandwidth memory allocated for local values, in bytes. Default is 0.
With ``DynamicEmbeddingShardingPlanner``, this is set to
``ceil(global_hbm_for_values / world_size)`` per rank.
device_id : Optional[int], optional
CUDA device index.
training: bool
Flag to indicate dynamic embedding tables is working on training mode or evaluation mode, default to `True`.
If in training mode. **dynamicemb** stores embeddings and optimizer states together in the underlying key-value table. e.g.
key:torch.int64
value = torch.concat(embedding, opt_states, dim=1)
Therefore, if `training=True` the module allocates memory for optimizer states; the per-row state size follows the `optimizer` entry in `fused_params` (FBGEMM `EmbOptimType`), as used by `BatchedDynamicEmbeddingTablesV2`.
initializer_args : DynamicEmbInitializerArgs
Arguments for initializing dynamic embedding vector values when training, and default using uniform distribution.
For ``UNIFORM`` and ``TRUNCATED_NORMAL``, ``lower`` and ``upper`` default to
``±1/sqrt(N)`` where ``N`` is ``EmbeddingConfig.num_embeddings``.
eval_initializer_args: DynamicEmbInitializerArgs
The initializer args for evaluation mode, and will return torch.zeros(...) as embedding by default if index/sparse feature is missing.
caching: bool
Flag to indicate dynamic embedding tables is working on caching mode, default to `False`.
When the device memory on a single GPU is insufficient to accommodate a single shard of the dynamic embedding table,
dynamicemb supports the mixed use of device memory and host memory(pinned memory).
But by default, the values of the entire table are concatenated with device memory and host memory.
This means that the storage location of one embedding is determined by `hash_function(key)`, and mapping to device memory will bring better lookup performance.
However, sparse features in training are often with temporal locality.
In order to store hot keys in device memory, dynamicemb creates two table instances,
whose values are stored in device memory and host memory respectively, and store hot keys on the GPU table priorily.
If the GPU table is full, the evicted keys will be inserted into the host table.
If the host table is also full, the key will be evicted(all the eviction is based on the score per key).
The original intention of eviction is based on this insight: features that only appear once should not occupy memory(even host memory) for a long time.
In short:
set **`caching=True`** will create a GPU table and a host table, and make GPU table serves as a cache;
set **`caching=False`** will create a hybrid table which use GPU and host memory in a concatenated way to store value.
All keys and other meta data are always stored on GPU for both cases.
init_capacity : Optional[int], optional
The initial capacity of the table. If not set, it defaults to max_capacity after sharding.
``BatchedDynamicEmbeddingTablesV2`` also sets ``init_capacity`` to ``max_capacity`` when it is still ``None``.
If `init_capacity` is provided, it will serve as the initial table capacity on a single GPU.
With ``DynamicEmbeddingShardingPlanner``, it is rounded up to a multiple of the effective
``bucket_capacity`` in ``_prepare_dynemb_table_options``, then capped at ``max_capacity`` if the aligned value is larger.
As the `load_factor` of the table increases, its capacity will gradually double (rehash) until it reaches `max_capacity`.
Rehash will be done implicitly.
Note: This is the setting for a single table at each rank.
max_load_factor : float
The maximum load factor before rehashing occurs. Default is 0.5.
In NO_EVICTION mode, this option is ignored: the implementation uses a fixed effective max load factor of 0.5 for the key_index_map (initial sizing and expansion). See the "Table expansion" section for NO_EVICTION trigger conditions.
score_strategy(DynamicEmbScoreStrategy or Tuple[DynamicEmbScoreStrategy, ...]):
dynamicemb gives each key-value pair a score to represent its importance.
Once there is insufficient space, the key-value pair will be evicted based on the score.
The `score_strategy` is used to configure how to set the scores for keys in each batch.
Default to DynamicEmbScoreStrategy.TIMESTAMP.
For the multi-GPUs scenario of model parallelism, every rank's score_strategy should keep the same for one table,
as they are the same table, but stored on different ranks.
A *compound* score strategy may be given as a tuple to keep **more than
one** score per key. Only
`(DynamicEmbScoreStrategy.TIMESTAMP, DynamicEmbScoreStrategy.LFU)` (in
either order) is supported for now: it ranks eviction by the LFU access
count while ALSO keeping a per-key last-access timestamp (two scores per
key, +8 bytes/key) so items touched since the last dump can be selected by
a time-based `incremental_dump`. Both tuple orders evict identically; the
order only controls how the score columns are laid out when dumped to file
and how `score_function` indexes them (see "Compound (multi-score)
strategies and score order" under `DynamicEmbScoreStrategy`). A one-element
tuple `(X,)` is treated as the single strategy `X`.
score_function(Optional[Callable]):
Optional custom eviction score for a compound `{TIMESTAMP, LFU}` table.
Signature `score_function(scores, cur_timestamp) -> float64`; `scores` is
indexed in the configured (logical) tuple order, and the key(s) with the
lowest returned score are evicted first. It is compiled to device code
with numba-cuda, so subscripts must be integer constants. Supported
**only** for the two-word `{TIMESTAMP, LFU}` compound strategy (`scores`
has length 2); any other (single or non-compound) `score_strategy` with a
`score_function` raises `ValueError`. Defaults to None (the built-in
frequency → older-timestamp evictor). See the "Customized eviction via
`score_function`" section under `DynamicEmbScoreStrategy`.
bucket_capacity : int
Capacity of each bucket in the hash table, and default is 128 (using 1024 when the table serves as cache).
A key will only be mapped to one bucket.
When the bucket is full, the key with the smallest score in the bucket will be evicted, and its slot will be used to store a new key.
The larger the bucket capacity, the more accurate the score based eviction will be, but it will also result in performance loss.
safe_check_mode : DynamicEmbCheckMode
Used to check if all keys in the current batch have been successfully inserted into the table.
Should dynamic embedding table insert safe check be enabled? By default, it is disabled.
Please refer to the API documentation for DynamicEmbCheckMode for more information.
global_hbm_for_values : int
Total GPU memory allocated to store embedding + optimizer states, in bytes. Default is 0.
It has different meanings under `caching=True` and `caching=False`.
When `caching=False`, it decides how much GPU memory is in the total memory to store value in a single hybrid table.
When `caching=True`, it decides the table capacity of the GPU table.
To match planner row counts and optimizer state width, size the **global** budget with
``get_table_value_bytes(embedding_config, EmbOptimType, world_size, bucket_capacity)``
(same ``bucket_capacity`` you pass into ``get_sharded_table_capacity``, e.g. ``128`` or ``MAX_BUCKET_CAPACITY``),
then multiply by a cache ratio or scale if desired. The planner overwrites **per-rank**
``local_hbm_for_values`` to ``ceil(global_hbm_for_values / world_size)``.
external_storage: Storage
The external storage/ParamterServer which inherits the interface of Storage, and can be configured per table.
If not provided, will using DynamicEmbeddingTable as the Storage.
index_type : Optional[torch.dtype], optional
Index type of sparse features, will be set to DEFAULT_INDEX_TYPE(torch.int64) by default.
admit_strategy : Optional[AdmissionStrategy], optional
Admission strategy for controlling which keys are allowed to enter the embedding table.
If provided, only keys that meet the strategy's criteria will be inserted into the table.
Keys that don't meet the criteria will still be initialized and used in the forward pass,
but won't be stored in the table. Default is None (all keys are admitted).
Anything deciding takes -- a frequency counter, an initializer for
the rows it rejects -- is configured on the strategy itself, e.g.
``FrequencyAdmissionStrategy(threshold=..., counter=KVCounter(...))``.
The module turns the strategies of the tables it fuses into one
``MultiTableAdmitter``, which is what actually decides.
admission_counter : Optional[Counter], optional
Deprecated, and warns when set. Pass the counter to the strategy
that uses it instead.
evicted_item_mode : EvictedItemMode, optional
How the *last-tier* storage handles an item it evicts. ``DISCARD``
(default) drops evicted keys with zero overhead. ``RETAIN_KEY`` retains
the keys it evicts so they can be read back with ``pop_evicted_keys``.
Only the final tier that truly discards a key records it -- intermediate
cache / HBM tiers spill their evictions to the next tier and are not
recorded. Records the (key, table_id) only, no value/score. Default
``EvictedItemMode.DISCARD``.
Notes
-----
See ``DynamicEmb_APIs.md`` (this file) and ``dynamicemb/planner/planner.py`` for planner integration.
"""
embedding_dtype: Optional[torch.dtype] = None
dim: Optional[int] = None
max_capacity: Optional[int] = None
evict_strategy: DynamicEmbEvictStrategy = DynamicEmbEvictStrategy.LRU
local_hbm_for_values: int = 0 # in bytes
device_id: Optional[int] = None
training: bool = True
initializer_args: DynamicEmbInitializerArgs = field(
default_factory=DynamicEmbInitializerArgs
)
eval_initializer_args: DynamicEmbInitializerArgs = field(
default_factory=lambda: DynamicEmbInitializerArgs(
mode=DynamicEmbInitializerMode.CONSTANT,
value=0.0,
)
)
caching: bool = False
init_capacity: Optional[
int
] = None # if not set then set to max_capcacity after sharded
max_load_factor: float = 0.5 # max load factor before rehash(double capacity)
score_strategy: Union[
DynamicEmbScoreStrategy, Tuple[DynamicEmbScoreStrategy, ...]
] = DynamicEmbScoreStrategy.TIMESTAMP
bucket_capacity: int = 128
safe_check_mode: DynamicEmbCheckMode = DynamicEmbCheckMode.IGNORE
global_hbm_for_values: int = 0 # in bytes
external_storage: Storage = None
index_type: Optional[torch.dtype] = None
admit_strategy: Optional[AdmissionStrategy] = None
admission_counter: Optional[Counter] = None # deprecated
evicted_item_mode: EvictedItemMode = EvictedItemMode.DISCARD
```
Automatically find the dynamic embedding tables in the torch model and parallelly dump them into the file system, dumping into a single file.
```python
#How to import
from dynamicemb import DynamicEmbDump
#API arguments
def DynamicEmbDump(path: str, model: nn.Module, table_names: Optional[Dict[str, List[str]]] = None, optim: Optional[bool] = False , pg: Optional[dist.ProcessGroup] = None) -> None:
"""
Dump the distributed weights and corresponding optimizer states of dynamic embedding tables from the model to the filesystem.
The weights of the dynamic embedding table will be stored in each EmbeddingCollection or EmbeddingBagCollection folder.
The name of the collection is the path of the torch module within the model, with the input module defined as str of model.
Each dynamic embedding table will be stored as a key binary file and a value binary file, where the dtype of the key is int64_t,
and the dtype of the value is the table's own ``embedding_dtype`` (float32, float16 or bfloat16) -- values are stored at the
precision the table holds them at, not widened to float32. Each optimizer state is also treated as a dynamic embedding table
and shares that precision.
The value files carry no header, so the per-table meta JSON records ``embedding_dtype``, ``embedding_dim`` and
``optim_state_dtype``; :func:`DynamicEmbLoad` reads the files by them. A checkpoint written before those keys existed is read
as float32 with the row width recovered from the value file's size.
Parameters
----------
path : str
The main folder for weight files.
model : nn.Module
The model contains dynamic embedding tables.
table_names : Optional[Dict[str, List[str]]], optional
A dictionary specifying which embedding collection and which table to dump. The key is the name of the embedding collection,
and the value is a list of dynamic embedding table names within that collection. Defaults to None.
optim : Optional[bool], optional
Whether to dump the optimizer states. Defaults to False.
pg : Optional[dist.ProcessGroup], optional
The process group used to control the communication scope in the dump. Defaults to None.
Returns
-------
None
"""
```
Load embedding weights from the binary file generated by DynamicEmbDump into the dynamic embedding tables in the torch model.
```python
#How to import
from dynamicemb import DynamicEmbLoad
#API arguments
def DynamicEmbLoad(path: str, model: nn.Module, table_names: Optional[List[str]] = None, optim: bool = False , pg: Optional[dist.ProcessGroup] = None):
"""
Load the distributed weights and corresponding optimizer states of dynamic embedding tables from the filesystem into the model.
Each dynamic embedding table will be stored as a key binary file and a value binary file, where the dtype of the key is int64_t,
and the dtype of the value is the table's own ``embedding_dtype``. Each optimizer state is also treated as a dynamic embedding
table and shares that precision.
The value files are read at the precision and row width the per-table meta JSON records (``embedding_dtype`` /
``embedding_dim`` / ``optim_state_dtype``). A checkpoint written before those keys existed is read as float32 with the row
width recovered from the value file's size. A dim that disagrees with the runtime table is an error; a differing precision is
converted on load, with a warning.
Parameters
----------
path : str
The main folder for weight files.
model : nn.Module
The model containing dynamic embedding tables.
table_names : Optional[Dict[str, List[str]]], optional
A dictionary specifying which embedding collection and which table to load. The key is the name of the embedding collection,
and the value is a list of dynamic embedding table names within that collection. Defaults to None.
optim : bool, optional
Whether to load the optimizer states. Defaults to False.
pg : Optional[dist.ProcessGroup], optional
The process group used to control the communication scope in the load. Defaults to None.
Returns
-------
None
"""
```
Background In recommendation systems, an incremental dump refers to the process of exporting or updating only the new or changed data since the last data dump, rather than exporting the entire dataset each time. The background for using incremental dumps is that recommendation systems often operate on massive and continuously growing datasets, such as user interactions, item updates, and behavioral logs. Performing a full data dump frequently would be inefficient and resource-intensive, consuming significant storage and processing power.
Target The main purpose of incremental dump is to improve efficiency by reducing the amount of data that needs to be processed and transferred during each update cycle. This allows the recommendation system to stay up-to-date with the latest data changes while minimizing downtime, storage usage, and computational overhead. Incremental dumps enable faster data refresh and model retraining, ensuring that recommendations remain relevant and timely in dynamic, large-scale environments.
Behavior
Given a model contains one or more ShardedDynamicEmbeddingCollection, this API will dump the eligible indices and embeddings of all ranks, based on the input score_threshold.
The meaning of the threshold depends on the table's score_strategy:
- Strategies that carry a device-timestamp column — the single
TIMESTAMPstrategy, or a compound strategy that includesTIMESTAMPsuch as(TIMESTAMP, LFU)— produce a time-based incremental dump: only items whose last-access timestamp crosses the threshold (i.e. items that were touched/changed since the reference time) are dumped. Use the value returned byget_scoreas the reference threshold. - Strategies without a timestamp column (e.g.
STEP, singleLFU,NO_EVICTION) instead threshold on the absolute score value: items whose score is not less than the threshold are dumped. This is not a time-based increment.
Limitation —
dist_type:incremental_dumpsupports onlyroundrobinandhash_roundrobinsharding. A table sharded withdist_type="continuous"raisesNotImplementedError. The returnedslot_indexis meant for precisereplay_increment, which reconstructs each key's owning rank from the key via(key or hash(key)) % world_size;continuoususes a different, range-based key→rank mapping that this path does not implement. Useroundrobinorhash_roundrobinif you need incremental dump.
```python
#How to import
from dynamicemb.incremental_dump import incremental_dump
#API arguments
def incremental_dump(
model: torch.nn.Module,
score_threshold: Union[int, Dict[str, Dict[str, int]]],
pg: Optional[dist.ProcessGroup] = None,
) -> Dict[str, "DeltaDumpResult"]:
"""Dump the model's embedding tables incrementally based on the score threshold. The index-embedding pair whose score is not less than the threshold will be returned.
Args:
model(nn.Module):The model containing dynamic embedding tables.
score_threshold(Union[int, Dict[str, Dict[str, int]]]):
int: All embedding table's score threshold will be this integer. It will dump matched results for all tables in the model.
Dict[str, Dict[str, int]]: the first `str` is the name of embedding collection in the model. 'str' in Dict[str, int] is the name of dynamic embedding table, and `int` in Dict[str, int] is the table's score threshold. It will dump for only tables whose names present in this Dict.
pg(Optional[dist.ProcessGroup]): optional. The process group used to control the communication scope in the dump (the all_gather of keys/values/slot_index). Defaults to None.
Returns
-------
Dict[str, DeltaDumpResult]:
``{collection_path: DeltaDumpResult}`` -- one ``DeltaDumpResult`` per
embedding collection. Each ``DeltaDumpResult`` holds column-aligned
per-table lists (element ``i`` refers to ``table_names[i]``):
- ``table_names: List[str]`` -- the dumped table names.
- ``keys: List[torch.Tensor]`` -- per-table matched keys on host.
- ``values: List[torch.Tensor]`` -- per-table `[N, dim]` embeddings on
host. **Embeddings only**: the rest of the stored row is in
`optimizer_states`, and concatenating the two along dim 1
reproduces the row as the table holds it.
- ``optimizer_states: List[Optional[torch.Tensor]]`` -- per-table
optimizer state on host, at the **same width the file checkpoint
uses** — narrower than the runtime row for rowwise Adagrad, whose
fused layout reserves 16 bytes per row but fills one scalar.
``None`` for a table whose optimizer keeps no per-row state, e.g.
plain SGD.
- ``scores: List[torch.Tensor]`` -- per-table `[N, num_scores]` score
words on host, in the table's configured (logical) column order.
Timestamp columns hold an **age** (`current_score - score`), not a
raw timestamp, since `%globaltimer` is per device and resets across
boots; a consumer rebases them onto its own clock. Every other
column (LFU frequency, STEP, CUSTOMIZED, NO_EVICTION row) is
carried verbatim.
- ``erased_keys: List[Optional[torch.Tensor]]`` -- per-table keys an
explicit erase removed, for tables whose ``evicted_item_mode``
erase asked to have recorded (each `erase` call passes its own
`EvictedItemMode`), so always present and possibly empty. These
are the removals `replay_increment` applies.
- ``evicted_keys: List[Optional[torch.Tensor]]`` -- per-table retained
evicted keys on host for tables with ``evicted_item_mode=RETAIN_KEY``,
else ``None``. Returning them drains that table's retained-evicted
buffer (each evicted key reported once across successive calls).
- ``meta: List[Dict[str, Any]]`` -- per-table metadata, a flat dict:
- ``"current_score": int`` -- the table's current score after this
dump; usable as the next forward's score and as the next
``incremental_dump`` threshold.
- ``"slot_index": torch.Tensor`` -- int64 host tensor aligned with
``keys``; the storage slot each dumped key occupies (for
``replay_increment``).
- ``"current_capacity": int`` -- key-map slots: the modulus that
decides a key's home bucket.
- ``"row_capacity": Tuple[int, ...]`` -- value-buffer rows per
storage tier, which is what bounds a row write. Equal to
`current_capacity` except under NO_EVICTION, whose key map is
deliberately larger than its value buffer.
- ``"bucket_capacity": int`` -- slots per hash bucket.
- ``"num_scores": int`` -- score words per key; part of the
slot layout ``replay_increment`` compares against.
- ``"world_size": int`` -- ranks the table was sharded across.
- ``"table_options": DynamicEmbTableOptions`` -- the table config.
"""
```
More usage please see test
Background
incremental_dump produces a delta; replay_increment is the other half of that pipeline — it writes the delta back into a model. The typical deployment is delta replication: a training job dumps periodically, the delta is shipped to a serving replica, and the replica replays it to catch up without reloading a full checkpoint.
Behavior
For every table in the delta, replay_increment erases the table's erased_keys (so the target converges to the source) and then upserts keys / values at the slots in meta["slot_index"]. evicted_keys is never applied: the key that took an evicted key's slot is in the same delta and overwrites it.
Write-back is by slot. Every key is written at the slot and value row it held in the source table, leaving the target layout-identical to it. A key can only be found inside its own home bucket, and that bucket is hash(key) % capacity / bucket_capacity, so this requires the target's layout to match the source's. replay_increment compares the delta's meta (current_capacity, row_capacity, bucket_capacity, num_scores, world_size, and the table_options fields score_strategy / dim / dist_type) against the target table, and checks that every column has the shape the write path needs. Both happen for all of the delta's tables before any of them is written, so a mismatch anywhere raises ValueError naming the first offending field with the collection untouched — a delta spans a whole collection, and a partly applied one is worse to recover from than a rejected one. Configure the target to match the source, or rebuild it from a full checkpoint (DynamicEmbLoad) instead.
The same rule is enforced per key inside the kernel — a slot that does not land in its key's home bucket raises rather than dropping the key.
Precondition: writing a key at its source slot overwrites whatever occupies that slot in the target. That is what makes a replica converge — if the source evicted key
Bto make room forA, a same-capacity replica must do the same. It also means a replayed table must be built only by loading/replaying from its source: a table that also takes independent writes can lose a key whose slot a delta key claims.
Scores and optimizer state travel under ReplayContent, which asks for both by default. Drop SCORE and a restored key is scored as if it had just been inserted into the target, so the replica ranks its own future evictions by when it received each key rather than by how the source ranked it — the two can evict in different orders, but never disagree on the value of a key they both hold. Drop OPTIMIZER_STATE and a key keeps its state only if it already occupies the target row. NO_EVICTION is unaffected by SCORE either way and is exact: its score word is a value row, not a score.
Sharding. Replay keeps only the keys this rank owns, recomputing ownership from the key with this model's global world size — the fan-out the tables were sharded over. There is deliberately no process-group argument: replay is local (filter, then write), so a group could only narrow the modulus and mis-route every key. A delta gathered over a process group (incremental_dump(..., pg)) holds the whole group's keys and so fans out correctly when the same delta is handed to every rank. A per-rank delta (pg=None) holds only the producing rank's keys and should be replayed there — replaying it on another rank is not an error, every key simply belongs to someone else and is skipped, which ReplayStats.skipped reports. Like incremental_dump, only roundrobin and hash_roundrobin are supported; continuous raises NotImplementedError.
Optimizer state is not part of a delta. A key that already occupies its target row keeps its optimizer state; a row taken over from another key (or a brand-new one) is reset to the table's initial optimizer state.
```python
#How to import
from dynamicemb import replay_increment
#API arguments
def replay_increment(
model: torch.nn.Module,
deltas: Dict[str, "DeltaDumpResult"],
content: ReplayContent = ReplayContent.ALL,
) -> Dict[str, Dict[str, "ReplayStats"]]:
"""Write incremental_dump results back into a model's dynamic embedding tables.
Args:
model(nn.Module): The model containing dynamic embedding tables.
deltas(Dict[str, DeltaDumpResult]): `incremental_dump`'s return value, keyed by embedding-collection path. Collections or tables the model does not have are skipped with a warning.
content(ReplayContent): what travels with each key besides its embedding, which always does — optimizer state, score words, or both. Defaults to both; `ReplayContent.EMBEDDING_ONLY` for neither.
Returns
-------
Dict[str, Dict[str, ReplayStats]]:
`{collection_path: {table_name: ReplayStats}}`, where `ReplayStats` has:
- `upserted: int` -- keys written back at their source slot.
- `erased: int` -- keys actually removed before the upsert, counted as really removed (a delta can name keys this replica never held).
- `skipped: int` -- delta keys this rank does not own.
Raises
------
ValueError: a target table's layout does not match the source's, or a delta is missing the per-key data replay needs.
NotImplementedError: a table is sharded with `dist_type="continuous"`.
"""
```
Example — replicate a training model's deltas into a serving model:
```python
from dynamicemb import replay_increment
from dynamicemb.incremental_dump import get_score, incremental_dump
threshold = get_score(train_model) # reference point for the next window
... # train for a while
deltas = incremental_dump(train_model, threshold, pg)
threshold = {c: {r.table_names[i]: r.meta[i]["current_score"]
for i in range(len(r.table_names))}
for c, r in deltas.items()} # threshold for the NEXT dump
stats = replay_increment(serve_model, deltas) # raises if layouts differ
for collection, per_table in stats.items():
for name, s in per_table.items():
print(collection, name, s.upserted, "keys replayed at their source slot")
```
Background
When a table's last-tier storage is full, evicting a key drops it from the system entirely. With evicted_item_mode=RETAIN_KEY (see DynamicEmbTableOptions), the last tier instead retains the keys it evicts so they can be read back -- e.g. to feed a downstream key-value store, a cold-tier archive, or an offline pipeline.
Behavior
Returns, per table, the keys evicted since the previous call, deduplicated within a table. This is a read-and-clear (incremental) operation: each evicted key is reported exactly once across successive calls, and returning a table's keys drains its retained-evicted buffer on this rank. Tables without evicted_item_mode=RETAIN_KEY are omitted from the result.
Keys removed by an explicit erase are not here — they live in a separate buffer read by pop_erased_keys. The two are kept apart because a consumer usually wants different things from them: an eviction says the table ran out of room, an erase says someone asked for the key to go.
```python
#How to import
from dynamicemb import pop_evicted_keys
#API arguments
def pop_evicted_keys(
model: torch.nn.Module,
table_names: Optional[Dict[str, List[str]]] = None,
pg: Optional[dist.ProcessGroup] = None,
) -> Dict[str, Dict[str, torch.Tensor]]:
"""Return (and clear) the keys evicted and retained by last-tier storage, per table.
Only tables created with evicted_item_mode=RETAIN_KEY are included; all other tables are omitted.
Args:
model(nn.Module): the model containing dynamic embedding tables.
table_names(Optional[Dict[str, List[str]]]): optional filter, keyed by
embedding-collection path -> [table_name, ...]. None pops every
retain-enabled table in the model.
pg(Optional[dist.ProcessGroup]): optional. None returns each rank's LOCAL
evicted keys (row-wise sharded, hence disjoint across ranks; zero
communication). When given, keys are all_gathered within pg so every
rank in the group receives the group-wide union. Clearing always
affects only this rank's buffer, regardless of pg.
Returns
-------
Dict[str, Dict[str, torch.Tensor]]:
{collection_path: {table_name: keys}} where keys is a 1-D host
tensor of table-unique evicted keys. Empty dict if the model has no
tables retaining evictions.
"""
```
Background The counterpart of pop_evicted_keys for the other way a key leaves a table: an explicit removal, rather than the table running out of room.
Unlike evictions, this is not a table setting. Retaining evictions has to be decided up front — it swaps in a collecting insert kernel and costs an extra kernel output plus a device sync on every evicting insert. An erase already holds its keys, so recording them costs a copy and nothing has to be prepared; each erase call therefore passes its own EvictedItemMode and decides for itself.
Behavior
Same arguments and the same read-and-clear semantics as pop_evicted_keys, draining the erased-key buffer instead. Tables whose mode retains no erases are omitted.
Only these keys are removals a replica has to perform for itself, which is why replay_increment applies erased_keys and never evicted_keys.
```python
#How to import
from dynamicemb import pop_erased_keys
#API arguments
def pop_erased_keys(
model: torch.nn.Module,
table_names: Optional[Dict[str, List[str]]] = None,
pg: Optional[dist.ProcessGroup] = None,
) -> Dict[str, Dict[str, torch.Tensor]]:
"""Return (and clear) the keys an explicit erase removed, per table.
Every table is included -- whether an erase was recorded is that
erase call's decision, not the table's, so there is nothing to
filter on. Arguments are identical to pop_evicted_keys.
"""
```
dynamicemb also provides a get_score interface whose returns are the current scores which will be used in the next forward pass.
Users can use get_score’s returns at an earlier time and use it as a threshold for the later incremental dump.
It's recommended for TIMESTAMP and STEP mode, and can also be used under CUSTOMIZED mode, but please note that get_score will only return the scores used in the next forward.
Under CUSTOMIZED mode, users need to understand the meaning of get_score's returns, and dynamicemb is not responsible for managing the score any more. For example, if users call get_score firstly to get the threshold, then decrease the score later to train the model, and finally call incremental_dump using the previous larger threshold, it will not dump anything.
```python
#How to import
from dynamicemb.incremental_dump import get_score
#API arguments
def get_score(model: torch.nn.Module) -> Union[Dict[str, Dict[str, int]], None]:
"""Get score for each dynamic embediing table.
Args:
model(torch.nn.Module): The model containing dynamic embedding tables.
Returns:
Dict[str, Dict[str,int]]:
- The first `str` is the name of embedding collection in the model.
- The second `str` is the name of dynamic embedding table.
- `int` represents:
* TIMESTAMP mode: global timer of device
* STEP mode: table's step after last forward pass
* CUSTOMIZED mode: score set in last forward pass
- Returns None if no dynamic embedding tables exist or scores unavailable
"""
```
Under CUSTOMIZED mode, users have to set the scores for each table. Users can set the scores one time and use it repeatedly in the several forward passes later. Generally speaking, the score increases as training progresses, and dynamicemb will throw a warning when the new score is less than the old one. Setting the environment variable DYNAMICEMB_CSTM_SCORE_CHECK to 0 will not throw the warnings.
```python
#How to import
from dynamicemb.incremental_dump import set_score
#API arguments
def set_score(
model: torch.nn.Module, table_score: Union[int, Dict[str, Dict[str, int]]]
) -> None:
"""Set the score for each dynamic embedding table. It will not reset the scores of each embedding table, but register a score for the
Args:
model(torch.nn.Module): The model containing dynamic embedding tables.
table_score(Union[int, Dict[str, Dict[str, int]]):
int: all embedding table's scores will be set to this integer.
Dict[str, Dict[str, int]]: the first `str` is the name of embedding collection in the model. 'str' in Dict[str, int] is the name of dynamic embedding table, and `int` in Dict[str, int] is the table's score which will broadcast to all scores in the same batch for the table.
Returns:
None.
"""
```
A counter maps a key to an accumulated count. A strategy that admits by
frequency owns one; the framework only carries it, reporting its memory and
writing it into checkpoints, which is what AdmissionStrategy.state() hands
over. Custom counters inherit Counter.
class Counter(abc.ABC):
"""Interface of a counter table which maps a key to a counter."""
@abc.abstractmethod
def add(
self,
keys: torch.Tensor,
table_ids: torch.Tensor,
frequencies: torch.Tensor,
) -> torch.Tensor:
"""Add frequencies to these keys and return their accumulated counts."""
@abc.abstractmethod
def erase(self, keys: torch.Tensor, table_ids: torch.Tensor) -> None:
"""Erase these keys."""
@abc.abstractmethod
def memory_usage(self, mem_type=MemoryType.DEVICE) -> int:
"""Consumption of one kind of memory."""
@abc.abstractmethod
def load(self, key_file, counter_file, table_id: int) -> None:
"""Load one table's keys and counts from these files."""
@abc.abstractmethod
def dump(self, key_file, counter_file, table_id: int) -> None:
"""Dump one table's keys and counts to these files."""dynamicemb provides KVCounter, which sizes one table's share of a counter.
It is configuration, not a Counter: the strategy holding it turns the shares
of the tables it serves into a single fused MultiTableKVCounter when a module
materializes it. The table is bucketized, and a full bucket evicts its
smallest-frequency key to make room, so size the capacity for the keys still
waiting to be admitted rather than for the embedding table.
class KVCounter:
def __init__(
self,
capacity: int,
bucket_capacity: int = 1024,
key_type: torch.dtype = torch.int64,
)Admission is two classes: what a table is configured with, and what a module runs.
AdmissionStrategy is the configuration. It is inert -- it allocates nothing,
touches no device, and may be given to as many tables as you like -- and it
decides nothing. A fused module turns the configurations of the tables it fuses
into the one thing that decides by calling create_admitter, which is where
anything device-side comes into being.
class AdmissionStrategy(abc.ABC):
@classmethod
@abc.abstractmethod
def create_admitter(
cls,
table_strategies: List["AdmissionStrategy"],
device: torch.device,
) -> "MultiTableAdmitter":
"""The admitter these tables share, with its device state allocated.
Called once, by the module, with one configuration per table it fuses.
"""
@classmethod
def one_configuration(cls, table_strategies) -> "AdmissionStrategy":
"""The single configuration these tables agree on.
They were grouped on get_grouped_key, so they are already
interchangeable and the first stands for all.
"""
def get_grouped_key(self):
"""What has to match for two tables to share one fused module.
The default is the instance itself, so only the same object groups.
Override it to let equal but separately constructed configurations
share a module, returning everything that changes the decision and, for
what is resolved per table, only what the tables must agree on.
"""MultiTableAdmitter is what the module runs, and holds everything deciding
takes. Only admit has to be implemented; the other two have defaults, so an
admitter that merely decides is one method.
class MultiTableAdmitter(abc.ABC):
@abc.abstractmethod
def admit(
self,
keys: torch.Tensor,
table_ids: torch.Tensor,
frequencies: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""Which of these missing keys may enter the table.
``frequencies`` is how often each key occurred in *this batch*, where
the module counts occurrences at all; None means treat each key as one.
It is an increment, not a running total -- any total is the admitter's
own to keep.
"""
def state(self) -> Optional[Counter]:
"""Persistent state the framework has to carry, or None.
Whatever an admitter accumulates across steps has to be reported in
memory accounting and survive a checkpoint. Return it here and the
framework does both.
"""
@property
def non_admitted_initializer(self):
"""What writes the rows this admitter rejects, or None for the table's.
A rejected key still takes part in the forward, so its row has to be
written by something. Return a ``MultiTableInitializer`` to take that
on; the module does the calling.
"""dynamicemb provides two pairs.
FrequencyAdmissionStrategy admits a key once it has been seen often enough,
and carries the counter it accumulates into, so admission is configured in one
place. Its admitter erases a key from that counter the moment it is admitted.
class FrequencyAdmissionStrategy(AdmissionStrategy):
def __init__(
self,
threshold: int,
counter: Optional[KVCounter] = None,
initializer_args: Optional[DynamicEmbInitializerArgs] = None,
)
# create_admitter -> MultiTableFrequencyAdmitterProbabilisticAdmissionStrategy admits a key by a coin toss instead, and its
admitter keeps no state at all -- state() is None and no counter is built for
it. Admission is consulted only for a key that is missing, so a key gets a
fresh toss every time it turns up and is not yet in the table: it takes
1 / probability appearances on average to get in. A batch holding a key k
times counts as k tosses.
class ProbabilisticAdmissionStrategy(AdmissionStrategy):
def __init__(
self,
probability: float,
initializer_args: Optional[DynamicEmbInitializerArgs] = None,
)
# create_admitter -> MultiTableProbabilisticAdmitterFor both, initializer_args says how to initialize the rows the admitter
rejects; None leaves them to the table's own initializer, which is also the only
way to get UNIFORM bounds derived from a table's row count.
Once the model containing EmbeddingCollection is built and initialized through DistributedModelParallel, it can be trained and evaluated on each GPU like a single GPU, with torchrec completing communication between different GPUs.
The switching between training and evaluation modes should be consistent with nn.Module, while training in DynamicEmbTableOptions is used to guide whether to allocate memory to optimizer states when builds the table.
Due to limited resources, the dynamic embedding table does not pre allocate memory for all keys. If a key appears for the first time during training, it will be initialized immediately during the training process. Please see initializer_args and eval_initializer_args in DynamicEmbTableOptions for more information.
Weighted-sum pooling for EmbeddingBagCollection tables backed by DynamicEmb is supported in training and eval (see
DynamicEmbPoolingMode). Gradients are scaled per-position: grad_embᵢ = wᵢ · grad_pooled, verified against an exact SGD
reference in test/unit_tests/test_weighted_pooled_embedding_v2.py and end-to-end, sharded, forward and backward, against a twin TorchRec model in test/unit_tests/test_twin_module.py (is_weighted=True).
BatchedDynamicEmbeddingTablesV2.forward previously took per_sample_weights, which was consumed as LFU frequency counters. That parameter is now split:
frequency_counters: Optional[Tensor]— per-position LFU frequency counters (the oldper_sample_weightssemantics).pooling_weights: Optional[Tensor]— per-position float32 weights for weighted-SUM pooling (new; SUM only). The two are mutually exclusive.
Callers that passed per_sample_weights=... for LFU counting must pass frequency_counters=... instead.
InferenceEmbeddingTable's per_sample_weights (inference/export path) is a separate API and is unchanged.
The size of the table is finite, but the set of keys during training may be infinite. dynamicemb provides the function of automatic eviction, which constrains the size of tables reasonably when there is no available space. See score_strategy and bucket_capacity for more information.
dynamicemb supports caching hot embeddings on GPU memory, and you can prefetch keys from host to device like torchrec. Caching and prefetch work for both sequence mode (NONE) and pooling modes (SUM/MEAN). See test_prefetch_flush_in_cache in test prefetch for usage examples.
dynamicemb supports external storage once external_storage in DynamicEmbTableOptions inherits the Storage interface under types.py.
Refer to demo PyDictStorage in unit test for detailed usage.
Users can specify the initial capacity of a table on a single GPU. When the table needs more space, the implementation may double the key_index_map and embedding table capacity (per table, only for tables that need it) before insert. Expansion is triggered in these paths so that insert does not fail for lack of capacity:
- Prefetch HBM direct: before inserting admitted keys into
DynamicEmbStorage. - Cache write-back: before writing evicted keys back to storage (only
DynamicEmbStorageuses cache mode;HybridStoragedoes not). - Generic forward (DEFAULT mode): before
storage.insert()when usingDynamicEmbStorageorHybridStorage(for the latter, only the host tier is expanded). - HybridStorage.load: before inserting keys evicted from HBM into the host tier.
Trigger conditions:
- Non–NO_EVICTION: When
max_load_factorwould be exceeded or (if set) atmax_capacity, the table(s) that need it are doubled. Seeinit_capacity,max_load_factor,max_capacityinDynamicEmbTableOptions. - NO_EVICTION: The option
max_load_factoris not used. The key_index_map is sized and expanded with a fixed effective max load factor of 0.5 (key_index_map capacity = ceil(init_capacity/0.5) at creation; expansion whenneeded > table_rowsorneeded > key_index_map.capacity(table_id)). Both the key_index_map and the embedding buffer are doubled for the affected table(s).
Dump/Load and incremental dump is different from general module in PyTorch, because dynamicemb's underlying implementation is a hash table instead of a dense torch.Tensor.
So dynamicemb provides dedicated interface to load/save models' states, and provide conditional dump to support online training.
Please see DynamicEmbDump, DynamicEmbLoad, incremental_dump in APIs Doc for more information.
When handling cache eviction and final table eviction, dynamicemb encounters randomness in the eviction of keys with the same score. To eliminate this uncertainty, dynamicemb provides a deterministic mode. Enabling this mode ensures that, under the same training script, the evicted keys will be determined.
This mode is enabled by setting the environment variable DEMB_DETERMINISM_MODE.
The following environment variables control runtime behavior of dynamicemb:
| Variable | Default | Effect |
|---|---|---|
DYNAMICEMB_CSTM_SCORE_CHECK |
1 |
When set to 0, suppresses warnings in set_score when a new score is less than the previous one (CUSTOMIZED eviction mode). |
DEMB_DETERMINISM_MODE |
unset | When set, enables deterministic eviction of keys with equal scores during cache eviction and final table eviction. |
DYNAMICEMB_DEBUG |
unset | When set to any non-empty value, enables extra runtime validation in segmented_unique_cuda. Copies segmented_range to CPU and verifies: (1) segmented_range[0] == 0, (2) segmented_range[num_tables] == num_keys, (3) monotonically non-decreasing. Raises an error with a descriptive message on violation. Intended for development and testing only; incurs one device-to-host copy per call. |