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

TF2.4使用sampled_softmax_loss出现类型错误,求解决方案

解决TensorFlow中sampled_softmax_loss的符号张量转换错误

你遇到的这个TypeError是因为自定义损失函数的实现方式不符合Keras的符号张量处理规范,具体来说有两个核心问题:

  1. 你的SampledSoftmaxLoss类没有继承tf.keras.losses.Loss基类,Keras无法正确识别和处理这个自定义损失的符号运算逻辑
  2. 你直接绑定了模型的符号输入张量(self.input = output_layer.input),在训练时实际批量数据张量和这个符号张量混用,导致NumPy无法处理符号类型,触发转换错误

下面给你两种可行的解决方案,第一种是修复自定义损失的实现,第二种是提供替代的候选采样方案:


方案一:修复自定义SampledSoftmaxLoss,符合Keras规范

这个方案严格遵循Keras的设计逻辑,通过多输出模型传递计算采样损失所需的特征张量,避免符号冲突:

import tensorflow as tf
import numpy as np
TextVectorization = tf.keras.layers.experimental.preprocessing.TextVectorization

# 正确实现SampledSoftmaxLoss,继承Keras Loss基类
class SampledSoftmaxLoss(tf.keras.losses.Loss):
    def __init__(self, dense_layer, num_classes, num_sampled=3):
        super().__init__()
        self.dense_layer = dense_layer
        self.num_classes = num_classes
        self.num_sampled = num_sampled
    
    def call(self, y_true, y_pred):
        # 将one-hot标签转换为采样损失需要的索引格式
        labels = tf.argmax(y_true, axis=1)
        labels = tf.expand_dims(labels, axis=-1)
        # 获取Dense层的权重和偏置
        weights = self.dense_layer.weights[0]
        biases = self.dense_layer.weights[1]
        # y_pred是模型输出的列表:[Dense层结果, 特征层结果]
        features = y_pred[1]
        # 计算采样损失并对batch取平均
        loss = tf.nn.sampled_softmax_loss(
            weights=weights,
            biases=biases,
            labels=labels,
            inputs=features,
            num_sampled=self.num_sampled,
            num_classes=self.num_classes
        )
        return tf.reduce_mean(loss)

# 配置参数
max_features = 50  # 最大词汇量
max_len = 10       # 序列填充长度
embedding_dims = 5 # 词嵌入维度
n_classes = 50     # 分类类别数(和标签维度一致)

# 准备训练数据
input_data = np.array([
    "Python Machine Learning",
    "Data Science from Scratch: First Principles with Python",
    "Hands-On Machine Learning with Scikit-Learn and TensorFlow: Concepts, Tools, and Techniques for Building Intelligent Systems",
    "Introduction to Machine Learning with Python: A Guide for Data Scientists",
    "Vital Introduction to Machine Learning with Python: Best Practices to Improve and Optimize Machine Learning Systems and Algorithms",
    "Machine Learning in Python: Essential Techniques for Predictive Analysis",
    "Python Data Science Handbook: Essential Tools for Working with Data",
    "Introducing Data Science: Big Data, Machine Learning, and more, using Python tools",
    "Real-World Machine Learning"])

# 生成单分类标签(sampled_softmax更适合单分类场景,修正原代码的标签逻辑)
labels = np.random.randint(0, n_classes, size=len(input_data))
labels_one_hot = tf.one_hot(labels, depth=n_classes).numpy()

# 文本向量化层
vectorize_layer = TextVectorization(
    max_tokens=max_features,
    output_mode='int',
    output_sequence_length=max_len)
vectorize_layer.adapt(input_data)  # 修正原代码中text_dataset的错误

# 构建多输出模型:同时输出Dense层结果和特征层结果
inp = tf.keras.Input(shape=(1,), dtype=tf.string)
idxs = vectorize_layer(inp)
embed = tf.keras.layers.Embedding(max_features + 1, embedding_dims, input_length=max_len)(idxs)
flat = tf.keras.layers.Flatten()(embed)
dense_layer = tf.keras.layers.Dense(n_classes)
out = dense_layer(flat)
# 把特征层和Dense输出都作为模型输出,供损失函数使用
model = tf.keras.models.Model(inp, [out, flat])

# 初始化损失函数并编译模型
loss_fn = SampledSoftmaxLoss(dense_layer, n_classes)
# 第二个输出不需要单独损失,设置为None
model.compile(optimizer='adam', loss=[loss_fn, None])

# 训练模型
model.fit(input_data, labels_one_hot, epochs=5)

# 预测时只取Dense层的输出
predictions = model.predict(input_data)[0]

关键修复点:

  • 让损失类继承tf.keras.losses.Loss,Keras会自动处理符号张量的运算逻辑
  • 模型输出包含特征层(Flatten结果),让损失函数能拿到计算采样损失所需的输入张量
  • 修正了原代码中vectorize_layer.adapt(text_dataset)的错误(应该用input_data)
  • 对批量损失取平均,确保返回Keras需要的标量损失值

方案二:使用内置候选采样损失替代

如果你不想自定义损失,可以使用TensorFlow提供的tf.nn.nce_loss(负采样损失,和sampled_softmax逻辑类似),通过add_loss直接集成到模型中:

import tensorflow as tf
import numpy as np
TextVectorization = tf.keras.layers.experimental.preprocessing.TextVectorization

# 配置参数
max_features = 50
max_len = 10
embedding_dims = 5
n_classes = 50

# 准备数据(同方案一)
input_data = np.array([
    "Python Machine Learning",
    "Data Science from Scratch: First Principles with Python",
    "Hands-On Machine Learning with Scikit-Learn and TensorFlow: Concepts, Tools, and Techniques for Building Intelligent Systems",
    "Introduction to Machine Learning with Python: A Guide for Data Scientists",
    "Vital Introduction to Machine Learning with Python: Best Practices to Improve and Optimize Machine Learning Systems and Algorithms",
    "Machine Learning in Python: Essential Techniques for Predictive Analysis",
    "Python Data Science Handbook: Essential Tools for Working with Data",
    "Introducing Data Science: Big Data, Machine Learning, and more, using Python tools",
    "Real-World Machine Learning"])

labels = np.random.randint(0, n_classes, size=len(input_data))
labels_one_hot = tf.one_hot(labels, depth=n_classes).numpy()

# 文本向量化层
vectorize_layer = TextVectorization(
    max_tokens=max_features,
    output_mode='int',
    output_sequence_length=max_len)
vectorize_layer.adapt(input_data)

# 构建模型并添加NCE损失
inp = tf.keras.Input(shape=(1,), dtype=tf.string)
labels_inp = tf.keras.Input(shape=(n_classes,))  # 把标签作为模型输入之一
idxs = vectorize_layer(inp)
embed = tf.keras.layers.Embedding(max_features + 1, embedding_dims, input_length=max_len)(idxs)
flat = tf.keras.layers.Flatten()(embed)

# 定义NCE损失的权重和偏置
nce_weights = tf.Variable(tf.random.normal([n_classes, embedding_dims*max_len]))
nce_biases = tf.Variable(tf.zeros([n_classes]))

# 定义NCE损失函数
def nce_loss(y_true):
    labels = tf.expand_dims(tf.argmax(y_true, axis=1), -1)
    loss = tf.nn.nce_loss(
        weights=nce_weights,
        biases=nce_biases,
        labels=labels,
        inputs=flat,
        num_sampled=3,
        num_classes=n_classes
    )
    return tf.reduce_mean(loss)

# 添加损失到模型
model = tf.keras.Model([inp, labels_inp], flat)
model.add_loss(nce_loss(labels_inp))

# 编译并训练
model.compile(optimizer='adam')
model.fit([input_data, labels_one_hot], epochs=5)

这个方案不需要自定义损失类,直接通过add_loss将采样损失绑定到模型上,适合快速迭代开发。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.11 08:57:27