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

如何在Keras模型中结合使用TensorFlow feature_column?

在Keras模型中结合TensorFlow Feature Column(含TF Hub嵌入列)

当然可以实现!TensorFlow提供了专门的tf.keras.layers.DenseFeatures层,能完美把Feature Column(包括TF Hub的文本嵌入列)接入Keras模型。下面我结合你的需求给出完整示例:

步骤1:导入依赖并定义Feature Column

首先我们还是先定义TF Hub的文本嵌入列,和你在Estimator里的用法一致:

import tensorflow as tf
import tensorflow_hub as hub

# 定义TF Hub文本嵌入列
embedded_text_feature_column = hub.text_embedding_column(
    key="sentence",
    module_spec="https://tfhub.dev/google/nnlm-en-dim128/1"
)

步骤2:构建Keras模型(接入Feature Column)

这里的关键是用DenseFeatures层把Feature Column转换成Keras可处理的张量,因为Feature Column通常对应字典格式的输入,所以我们的输入层要以字典形式定义:

# 定义字典形式的输入,key要和feature column的key对应
inputs = {
    "sentence": tf.keras.layers.Input(shape=(), dtype=tf.string, name="sentence")
}

# 用DenseFeatures解析feature column,输出Keras张量
x = tf.keras.layers.DenseFeatures([embedded_text_feature_column])(inputs)

# 后面就可以接你原本计划的Keras层了
x = tf.keras.layers.Dense(100, activation='relu')(x)
# 这里如果是分类任务,建议加softmax激活(和Estimator的DNNClassifier行为对齐)
outputs = tf.keras.layers.Dense(2, activation='softmax')(x)

# 组装完整模型
model = tf.keras.Model(inputs=inputs, outputs=outputs)

步骤3:编译并训练模型

接下来的编译和训练流程和普通Keras模型完全一致:

model.compile(
    optimizer=tf.keras.optimizers.Adam(learning_rate=0.001),
    loss=tf.keras.losses.SparseCategoricalCrossentropy(),  # 如果标签是整数形式用这个
    metrics=['accuracy']
)

# 示例训练数据(字典格式输入,和输入层对应)
train_data = {
    "sentence": ["this is positive", "that's negative", "great day", "bad experience"]
}
train_labels = [0, 1, 0, 1]

# 启动训练
model.fit(train_data, train_labels, epochs=5)

关键说明

  • DenseFeatures层是连接Feature Column和Keras的桥梁,它能自动处理Feature Column的输出(比如这里的文本嵌入),把它们转换成标准的Keras张量。
  • 如果有多个Feature Column(比如同时有数值特征、类别特征),只需要把所有列放进DenseFeatures的列表里,它会自动拼接所有特征。
  • 输入数据必须是字典格式,字典的key要和Feature Column的key一一对应,这样DenseFeatures才能正确解析。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.28 09:03:03