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

在回归问题中使用tf.keras.layers.Embedding处理分类变量的实现疑问

实现指引:用Embedding替代独热编码构建回归ANN

1. 数据准备(基于你的现有代码)

先完成数据预处理并划分训练/测试集:

import pandas as pd
from sklearn import datasets
from sklearn.model_selection import train_test_split
import tensorflow as tf

# 你的原有数据处理逻辑
iris = datasets.load_iris()
df = pd.DataFrame(iris['data'], columns = iris['feature_names'])
df['iris_class'] = pd.Series(iris['target'], name = 'target_values')
df['iris_class_name'] = df['iris_class'].replace([0,1,2], ['iris-' + species for species in iris['target_names'].tolist()])
df.columns = df.columns.str.replace("[() ]", "")

# 分离特征与目标变量
X = df[['iris_class_name', 'sepalwidthcm', 'petallengthcm']]
y = df['sepallengthcm']

# 划分训练集、测试集
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=42)

2. 构建分类特征的Embedding层

参考你提供的代码,适配iris_class_name特征:

# 获取分类特征的词汇表
iris_vocabulary = X['iris_class_name'].unique().tolist()

# 字符串转整数索引的Lookup层
iris_lookup = tf.keras.layers.StringLookup(
    vocabulary=iris_vocabulary, 
    mask_token=None,
    output_mode='int'
)

# 构建Embedding序列层
embedding_dim = 2  # 针对3类特征,2-5维足够,可按需调整
iris_embedding = tf.keras.Sequential([
    iris_lookup,
    tf.keras.layers.Embedding(
        input_dim=len(iris_vocabulary) + 1,  # +1预留未见过词汇的处理空间
        output_dim=embedding_dim,
        name='iris_class_embedding'
    )
], name='iris_embedding')

3. 搭建完整的回归ANN模型

将Embedding输出与数值特征拼接后接入全连接层:

# 定义各输入层
iris_class_input = tf.keras.Input(shape=(1,), dtype=tf.string, name='iris_class_name_input')
sepal_width_input = tf.keras.Input(shape=(1,), dtype=tf.float32, name='sepalwidthcm_input')
petal_length_input = tf.keras.Input(shape=(1,), dtype=tf.float32, name='petallengthcm_input')

# 处理分类特征:生成Embedding并展平
iris_embedding_output = iris_embedding(iris_class_input)
iris_embedding_flat = tf.keras.layers.Flatten()(iris_embedding_output)

# 拼接所有特征
concatenated_features = tf.keras.layers.concatenate([
    iris_embedding_flat,
    sepal_width_input,
    petal_length_input
])

# 全连接层构建
dense1 = tf.keras.layers.Dense(32, activation='relu')(concatenated_features)
dense2 = tf.keras.layers.Dense(16, activation='relu')(dense1)
# 回归任务输出层:无激活函数,输出单个数值
output = tf.keras.layers.Dense(1, name='sepallengthcm_output')(dense2)

# 组装模型
model = tf.keras.Model(
    inputs=[iris_class_input, sepal_width_input, petal_length_input],
    outputs=output
)

# 编译模型:回归任务用MSE损失,Adam优化器
model.compile(optimizer='adam', loss='mse', metrics=['mae'])

# 查看模型结构
model.summary()

4. 训练与评估模型

整理输入格式并完成训练:

# 准备训练/测试输入(字典格式对应输入层名称)
train_inputs = {
    'iris_class_name_input': X_train['iris_class_name'].values,
    'sepalwidthcm_input': X_train['sepalwidthcm'].values,
    'petallengthcm_input': X_train['petallengthcm'].values
}

test_inputs = {
    'iris_class_name_input': X_test['iris_class_name'].values,
    'sepalwidthcm_input': X_test['sepalwidthcm'].values,
    'petallengthcm_input': X_test['petallengthcm'].values
}

# 训练模型
history = model.fit(
    train_inputs,
    y_train.values,
    epochs=50,
    batch_size=8,
    validation_split=0.1
)

# 评估模型
loss, mae = model.evaluate(test_inputs, y_test.values)
print(f"测试集MSE损失: {loss:.4f}, MAE: {mae:.4f}")

# 示例预测
sample_input = {
    'iris_class_name_input': ['iris-setosa'],
    'sepalwidthcm_input': [3.5],
    'petallengthcm_input': [1.4]
}
predicted_sepal_length = model.predict(sample_input)
print(f"预测的花萼长度: {predicted_sepal_length[0][0]:.4f}")

关键注意点

  • StringLookup层:负责将字符串分类标签转为整数索引,是Embedding层的必要前置步骤。
  • Embedding维度:针对仅3类的特征,无需设置过高维度,避免过拟合。
  • 特征拼接:Embedding输出需先展平,再与数值特征拼接,保证输入全连接层的是一维向量。
  • 回归任务配置:输出层无激活函数,损失选用均方误差(MSE),评估常用平均绝对误差(MAE)。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.19 01:46:30