@@ -326,8 +326,6 @@ def _populate_zero_collision_tbe_params(
326326 meta_header_lens [i ] = table .virtual_table_eviction_policy .get_meta_header_len ()
327327 if not isinstance (table .virtual_table_eviction_policy , NoEvictionPolicy ):
328328 enabled = True
329-
330- fs_eviction_enabled : bool = False
331329 if enabled :
332330 counter_thresholds = [0 ] * len (config .embedding_tables )
333331 ttls_in_mins = [0 ] * len (config .embedding_tables )
@@ -386,7 +384,6 @@ def _populate_zero_collision_tbe_params(
386384 raise ValueError (
387385 f"Do not support multiple eviction strategy in one tbe { eviction_strategy } and 5 for tables { table_names } "
388386 )
389- fs_eviction_enabled = True
390387 elif isinstance (policy_t , TimestampBasedEvictionPolicy ):
391388 training_id_eviction_trigger_count [i ] = (
392389 policy_t .training_id_eviction_trigger_count
@@ -462,7 +459,6 @@ def _populate_zero_collision_tbe_params(
462459 backend_return_whole_row = (backend_type == BackendType .DRAM ),
463460 eviction_policy = eviction_policy ,
464461 embedding_cache_mode = embedding_cache_mode_ ,
465- feature_score_collection_enabled = fs_eviction_enabled ,
466462 )
467463
468464
@@ -2905,7 +2901,6 @@ def __init__(
29052901 _populate_zero_collision_tbe_params (
29062902 ssd_tbe_params , self ._bucket_spec , config , backend_type
29072903 )
2908- self ._kv_zch_params : KVZCHParams = ssd_tbe_params ["kv_zch_params" ]
29092904 compute_kernel = config .embedding_tables [0 ].compute_kernel
29102905 embedding_location = compute_kernel_to_embedding_location (compute_kernel )
29112906
@@ -3190,40 +3185,7 @@ def forward(self, features: KeyedJaggedTensor) -> torch.Tensor:
31903185 self ._split_weights_res = None
31913186 self ._optim .set_sharded_embedding_weight_ids (sharded_embedding_weight_ids = None )
31923187
3193- weights = features .weights_or_none ()
3194- per_sample_weights = None
3195- score_weights = None
3196- if weights is not None and weights .dtype == torch .float64 :
3197- fp32_weights = weights .view (torch .float32 )
3198- per_sample_weights = fp32_weights [:, 0 ]
3199- score_weights = fp32_weights [:, 1 ]
3200- elif weights is not None and weights .dtype == torch .float32 :
3201- if self ._kv_zch_params .feature_score_collection_enabled :
3202- score_weights = weights .view (- 1 )
3203- else :
3204- per_sample_weights = weights .view (- 1 )
3205- if features .variable_stride_per_key () and isinstance (
3206- self .emb_module ,
3207- (
3208- SplitTableBatchedEmbeddingBagsCodegen ,
3209- DenseTableBatchedEmbeddingBagsCodegen ,
3210- SSDTableBatchedEmbeddingBags ,
3211- ),
3212- ):
3213- return self .emb_module (
3214- indices = features .values ().long (),
3215- offsets = features .offsets ().long (),
3216- weights = score_weights ,
3217- per_sample_weights = per_sample_weights ,
3218- batch_size_per_feature_per_rank = features .stride_per_key_per_rank (),
3219- )
3220- else :
3221- return self .emb_module (
3222- indices = features .values ().long (),
3223- offsets = features .offsets ().long (),
3224- weights = score_weights ,
3225- per_sample_weights = per_sample_weights ,
3226- )
3188+ return super ().forward (features )
32273189
32283190
32293191class BatchedFusedEmbeddingBag (
0 commit comments