TensorFlow中如何按指定行索引将张量部分行缩放2倍
Solution: Scale Specific Rows in Embedding Tensor by 2x
Alright, let's break down how to scale the rows specified in indices_of_scaling_ids by a factor of 2.0 in your emb_inputs tensor. The core idea is to create a scaling weight tensor that applies 2.0 to the target rows and 1.0 to all others, then use element-wise multiplication (leveraging TensorFlow's broadcasting) to apply the scaling.
Step-by-Step Implementation
Here's the code that will work for your setup:
# Get the length of the rows dimension from emb_inputs all_row_len = tf.shape(emb_inputs)[1] # Initialize a weight tensor with all 1.0s (no scaling by default) scaling_weights = tf.ones(shape=(all_row_len,), dtype=emb_inputs.dtype) # Prepare indices and update values for the rows we want to scale # tf.tensor_scatter_nd_update expects indices to be 2D, so we expand the dimension target_indices = tf.expand_dims(indices_of_scaling_ids, axis=1) scale_updates = tf.ones_like(indices_of_scaling_ids, dtype=emb_inputs.dtype) * 2.0 # Update the weights to set target rows to 2.0 scaling_weights = tf.tensor_scatter_nd_update(scaling_weights, target_indices, scale_updates) # Reshape the weights to match the broadcast requirements of emb_inputs # We need (1, all_row_len, 1) to broadcast across batch and embedding dimensions scaling_weights = tf.expand_dims(tf.expand_dims(scaling_weights, axis=0), axis=-1) # Apply the scaling to the embedding inputs scaled_emb_inputs = emb_inputs * scaling_weights
How This Works
- Weight Initialization: We start with a tensor of 1.0s, so by default, no rows are scaled.
- Scatter Update: We use
tf.tensor_scatter_nd_updateto set the weights for your target rows to 2.0. This works even ifindices_of_scaling_idsis a dynamic tensor (not a fixed list). - Broadcasting: By expanding the weights tensor to shape
(1, all_row_len, 1), TensorFlow will automatically broadcast it across the batch dimension (batch_size=4) and the embedding dimension (emb_size=128), ensuring every element in the target rows gets multiplied by 2.0.
This approach is efficient and compatible with GPU execution, which aligns with your setup since your tensor is on /device:GPU:0.
内容的提问来源于stack exchange,提问作者Vishal Anand
相关产品推荐
相关产品推荐

