TensorFlow中忽略数组元素计算tf.reduce_sum及自定义MSE损失
Hey there! Let's tackle your two TensorFlow questions one by one:
tf.reduce_sum in TensorFlow The core idea here is to use a mask to exclude the elements you don't want included in the sum. There are two straightforward approaches:
Approach 1: Multiply by a binary mask
Create a mask tensor where positions you want to keep are 1.0 and positions to ignore are 0.0. Multiply this mask with your original tensor before running tf.reduce_sum—this zeros out the unwanted elements so they don't contribute to the sum:
import tensorflow as tf # Example input tensor (shape [batch_size, 10]) tensor = tf.random.normal([32, 10]) # Create mask: keep all elements except indices 4, 5, 6 mask = tf.ones([10], dtype=tf.float32) # Set the target indices to 0 mask = tf.tensor_scatter_nd_update(mask, indices=[[4], [5], [6]], updates=[0.0, 0.0, 0.0]) # Apply mask and compute sum sum_result = tf.reduce_sum(tensor * mask, axis=1)
Approach 2: Use tf.boolean_mask to filter elements
If you prefer to directly remove unwanted elements from the tensor before summing, tf.boolean_mask extracts only elements where the mask is True:
# Create boolean mask: True for positions to keep bool_mask = tf.ones([10], dtype=tf.bool) bool_mask = tf.tensor_scatter_nd_update(bool_mask, indices=[[4], [5], [6]], updates=[False, False, False]) # Filter the tensor and compute sum filtered_tensor = tf.boolean_mask(tensor, bool_mask, axis=1) sum_result = tf.reduce_sum(filtered_tensor, axis=1)
Mask multiplication is usually more efficient (no tensor reshaping), while tf.boolean_mask is more readable if you need to work directly with the filtered elements.
We can adapt your existing code to ignore indices 4, 5, 6 using the same masking logic. Here are two solid options:
Option 1: Masked Squared Difference (Most Efficient)
Zero out the squared differences at ignored indices before summing—this keeps your tensor shape intact and avoids extra operations:
import tensorflow as tf def custom_mse_loss(target, output): # Create mask: 1.0 for positions to consider, 0.0 for indices 4,5,6 mask = tf.ones([10], dtype=tf.float32) mask = tf.tensor_scatter_nd_update(mask, indices=[[4], [5], [6]], updates=[0.0, 0.0, 0.0]) # Calculate squared difference, apply mask, then sum per sample squared_difference = tf.reduce_sum(tf.square(target - output) * mask, axis=1) # Optional: If you want MSE to be the average of the 7 valid elements per sample, # uncomment the line below to divide each sum by the number of valid elements: # squared_difference = squared_difference / tf.reduce_sum(mask) # Take the mean over the batch mse_loss = tf.reduce_mean(squared_difference) return mse_loss # Test with your example target target = tf.constant([[0,1,2,3,0.5,0.5,0.5,7,8,9]]) output = tf.random.normal([1, 10]) # Random output for testing print(custom_mse_loss(target, output))
Option 2: Filter Elements First
Explicitly extract only the non-ignored elements from target and output before computing the loss:
def custom_mse_loss(target, output): # Boolean mask for valid positions bool_mask = tf.ones([10], dtype=tf.bool) bool_mask = tf.tensor_scatter_nd_update(bool_mask, indices=[[4], [5], [6]], updates=[False, False, False]) # Filter target and output to keep only valid indices target_filtered = tf.boolean_mask(target, bool_mask, axis=1) output_filtered = tf.boolean_mask(output, bool_mask, axis=1) # Calculate squared difference sum per sample, then batch mean squared_difference = tf.reduce_sum(tf.square(target_filtered - output_filtered), axis=1) mse_loss = tf.reduce_mean(squared_difference) return mse_loss
Quick Note:
Your original code sums squared differences per sample then averages across the batch. The first option preserves this logic—each sample's sum excludes the 3 ignored indices, then we average those sums. If you want the loss to represent the average of valid elements per sample (not the sum), uncomment the division line to divide each sample's sum by 7.
内容的提问来源于stack exchange,提问作者Peter111

