You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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_update to set the weights for your target rows to 2.0. This works even if indices_of_scaling_ids is 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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.26 09:40:58