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

如何在TensorFlow的sampled_softmax_loss中指定特定样本为负样本?

自定义tf.nn.sampled_softmax_loss的负样本(指定Hard Negative)

要主动指定特定负样本(比如推荐系统中的hard negative),无需依赖内置候选采样函数,直接手动构造sampled_values参数传入即可。

关键原理

sampled_values是一个三元组(sampled_candidates, true_expected_count, sampled_expected_count),完全可以手动生成,而非必须由tf.random.fixed_unigram_candidate_sampler等函数返回:

  • sampled_candidates:你指定的负样本ID集合
  • true_expected_count:真实标签的期望采样次数(无特殊需求时设为全1即可)
  • sampled_expected_count:自定义负样本的期望采样次数(无特殊需求时设为全1即可)

具体实现步骤

  1. 准备自定义负样本:根据业务需求筛选出hard negative样本,整理成张量,数据类型需与labels一致(通常为tf.int64)。

    • 若每个样本对应不同的负样本,张量形状为[batch_size, num_sampled]
    • 若用全局统一的负样本,张量形状为[num_sampled]
  2. 构造sampled_values三元组:

    • sampled_candidates:直接传入自定义的负样本张量
    • true_expected_count:创建与labels形状相同的全1张量
    • sampled_expected_count:创建与sampled_candidates形状相同的全1张量
  3. 传入sampled_softmax_loss计算损失

代码示例

import tensorflow as tf

# 假设以下参数已提前定义
batch_size = 32
num_classes = 10000
embedding_dim = 128
num_sampled = 10
num_true = 1

# 模型参数与输入示例
weights = tf.Variable(tf.random.normal([num_classes, embedding_dim]))
biases = tf.Variable(tf.zeros([num_classes]))
inputs = tf.random.normal([batch_size, embedding_dim])
labels = tf.random.uniform([batch_size, num_true], maxval=num_classes, dtype=tf.int64)

# 自定义hard negative样本:每个样本对应10个hard负样本
hard_negatives = tf.constant(
    [[100, 201, 305, 420, 512, 630, 777, 810, 901, 999] for _ in range(batch_size)],
    dtype=tf.int64
)

# 构造sampled_values
sampled_candidates = hard_negatives
true_expected_count = tf.ones_like(labels, dtype=tf.float32)
sampled_expected_count = tf.ones_like(sampled_candidates, dtype=tf.float32)
sampled_values = (sampled_candidates, true_expected_count, sampled_expected_count)

# 计算采样softmax损失
loss = tf.nn.sampled_softmax_loss(
    weights=weights,
    biases=biases,
    labels=labels,
    inputs=inputs,
    num_sampled=num_sampled,
    num_classes=num_classes,
    num_true=num_true,
    sampled_values=sampled_values
)

注意事项

  • 确保sampled_candidates的维度与num_sampled匹配,避免维度不兼容报错
  • 若不需要调整损失权重,true_expected_count和sampled_expected_count设为全1即可,不影响自定义负样本的生效

内容的提问来源于stack exchange,提问作者IMAPOTATO

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.09 04:10:30