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

TensorFlow中TabTransformer表格数据多分类的文档/示例咨询及报错问题

TabTransformer 表格数据多分类实现(TensorFlow/Keras环境)

目前Keras官方并未提供专门针对TabTransformer的多分类官方示例,现有的官方示例仅支持二分类场景。针对你修改输出层后出现错误的问题,核心原因通常是仅调整了输出层,但未同步修改损失函数、标签编码方式及评估指标,以下是具体的修正方案:

关键修改点

  • 输出层调整:将原二分类的Dense(1, activation='sigmoid')替换为Dense(C, activation='softmax'),其中C为你的多分类任务总类别数,这一步你已经完成,但需确保C的数值与实际类别数一致。
  • 损失函数替换:
    • 若你的标签是one-hot编码格式(例如类别数为3时,标签为[1,0,0]、[0,1,0]等),将损失函数从BinaryCrossentropy()改为CategoricalCrossentropy()。
    • 若你的标签是整数索引格式(例如类别数为3时,标签为0、1、2),则使用SparseCategoricalCrossentropy(),这种方式无需对标签做one-hot编码,更节省内存。
  • 评估指标同步修改:
    • 对应one-hot标签,使用CategoricalAccuracy()作为评估指标。
    • 对应整数索引标签,使用SparseCategoricalAccuracy()。

代码示例片段

以下是修改后的核心代码部分(基于官方二分类示例调整):

# 替换输出层
outputs = layers.Dense(C, activation="softmax")(x)
model = keras.Model(inputs=inputs, outputs=outputs)

# 选择对应损失函数和指标(以整数索引标签为例)
model.compile(
    optimizer=keras.optimizers.Adam(learning_rate=0.001),
    loss=keras.losses.SparseCategoricalCrossentropy(),
    metrics=[keras.metrics.SparseCategoricalAccuracy()],
)

额外注意事项

  • 确保训练数据的标签格式与选择的损失函数匹配,若标签是整数但误用了CategoricalCrossentropy,会直接引发维度不匹配的错误。
  • 若你在数据预处理阶段对标签做了one-hot编码(例如使用tf.keras.utils.to_categorical),则必须使用CategoricalCrossentropy和CategoricalAccuracy。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.12 19:11:31