TensorFlow中不依赖sampled_softmax_loss的负采样实现问询
sampled_softmax_loss) Absolutely! You don’t have to rely on sampled_softmax_loss to implement negative sampling in TensorFlow. Building a custom solution lets you fully control how negative samples are weighted by label frequency—exactly what you’re looking for. Let’s walk through the steps with concrete code examples.
Step 1: Calculate Label Frequency & Weighted Sampling Probabilities
First, you need to compute how often each label appears in your training data, then adjust those frequencies to create a sampling distribution. A common trick (used in Word2Vec) is to raise frequencies to a power (like 0.75) to soften the bias toward super-high-frequency labels, but you can tweak this to fit your needs.
# Assume `train_labels` is a 1D tensor of all your training labels (shape: [num_total_samples]) unique_labels, _, label_counts = tf.unique_with_counts(train_labels) # Convert counts to float and apply frequency weighting # Use exponent 0.75 to reduce over-sampling of extremely frequent labels (adjust as needed) weighted_counts = tf.pow(tf.cast(label_counts, tf.float32), 0.75) sampling_probs = weighted_counts / tf.reduce_sum(weighted_counts)
Step 2: Build a Custom Negative Sampling Function
With your sampling probabilities ready, you can write a function to draw negative samples on the fly during training. We’ll use tf.random.categorical to sample based on our weighted distribution:
def sample_negative_labels(num_neg_samples_per_pos, batch_pos_labels, unique_labels, sampling_probs): batch_size = tf.shape(batch_pos_labels)[0] total_neg_samples = batch_size * num_neg_samples_per_pos # Sample indices using our custom probability distribution # We use log(probs) because tf.random.categorical expects logits neg_indices = tf.random.categorical(tf.math.log(tf.expand_dims(sampling_probs, 0)), total_neg_samples) neg_labels = tf.gather(unique_labels, tf.squeeze(neg_indices)) # Reshape to match batch structure: [batch_size, num_neg_samples_per_pos] return tf.reshape(neg_labels, [batch_size, num_neg_samples_per_pos])
Step 3: Integrate into Your Training Loop
Now, you can use this sampling function alongside your model to compute a custom loss (like Noise-Contrastive Estimation, NCE, which works well for this setup):
# Example training batch setup batch_pos_labels = ... # 1D tensor of positive labels for the batch (shape: [batch_size]) model_inputs = ... # Your model's input tensor num_neg_samples = 5 # Number of negative samples per positive example # Get model outputs (logits for all classes) model_logits = your_model(model_inputs) # Sample negative labels neg_labels = sample_negative_labels(num_neg_samples, batch_pos_labels, unique_labels, sampling_probs) # Calculate loss for positive and negative samples # Positive samples: we want their logits to be classified as 1 pos_logits = tf.gather_nd(model_logits, tf.stack([tf.range(batch_size), batch_pos_labels], axis=1)) pos_loss = tf.nn.sigmoid_cross_entropy_with_logits( labels=tf.ones_like(pos_logits), logits=pos_logits ) # Negative samples: we want their logits to be classified as 0 neg_logits = tf.gather(model_logits, neg_labels, batch_dims=1) neg_loss = tf.nn.sigmoid_cross_entropy_with_logits( labels=tf.zeros_like(neg_logits), logits=neg_logits ) # Combine and average the loss total_loss = tf.reduce_mean(pos_loss) + tf.reduce_mean(neg_loss)
Key Notes
- Flexibility: You can adjust the exponent in the weighted counts (e.g., use 1.0 for pure frequency-based sampling, or lower values to reduce bias toward frequent labels).
- Efficiency: If you have an extremely large number of unique labels, precomputing the sampling distribution and using
tf.random.choice(withsampling_probs) might be slightly more efficient, buttf.random.categoricalworks great for most cases. - Loss Choice: While we used sigmoid cross-entropy here, you can also adapt this setup to use other losses like triplet loss if that fits your task better.
内容的提问来源于stack exchange,提问作者Long Le Minh

