如何在TensorFlow中结合Flatten层实现特征拼接?
问题解决方法
你的错误根源有两个:
- 拼接轴设置错误:你用了
axis=0(样本维度)拼接,但输入特征维度不匹配;应该用axis=1(特征维度)来拼接不同特征。 - 特征维度不一致:经纬度输入是
(None,)的一维张量,而展平后的embedding是(None,10)的二维张量,拼接时需要所有输入在非拼接轴上形状一致,所以要把一维的经纬度特征扩展为二维的(None,1)。
修改后的代码如下:
import numpy as np import pandas as pd import tensorflow as tf from matplotlib import pyplot as plt from tensorflow import keras from tensorflow.keras import Model from tensorflow.keras.callbacks import TensorBoard from tensorflow.keras.layers import ( CategoryEncoding, Concatenate, Dense, Discretization, Embedding, Flatten, Input, Reshape ) from tensorflow.keras.layers.experimental.preprocessing import HashedCrossing dnn_hidden_units = [32, 8] NBUCKETS = 16 latbuckets = np.linspace(start=38.0, stop=42.0, num=NBUCKETS).tolist() lonbuckets = np.linspace(start=-76.0, stop=-72.0, num=NBUCKETS).tolist() # 假设你已定义inputs字典(示例): # inputs = { # "pickup_longitude": Input(shape=(), dtype=tf.float32), # "pickup_latitude": Input(shape=(), dtype=tf.float32), # "dropoff_longitude": Input(shape=(), dtype=tf.float32), # "dropoff_latitude": Input(shape=(), dtype=tf.float32) # } # Bucketization with Discretization layer plon = Discretization(lonbuckets, name="plon_bkt")(inputs["pickup_longitude"]) plat = Discretization(latbuckets, name="plat_bkt")(inputs["pickup_latitude"]) dlon = Discretization(lonbuckets, name="dlon_bkt")(inputs["dropoff_longitude"]) dlat = Discretization(latbuckets, name="dlat_bkt")(inputs["dropoff_latitude"]) # Feature Cross with HashedCrossing layer p_fc = HashedCrossing(num_bins=NBUCKETS * NBUCKETS, name="p_fc")((plon, plat)) d_fc = HashedCrossing(num_bins=NBUCKETS * NBUCKETS, name="d_fc")((dlon, dlat)) pd_fc = HashedCrossing(num_bins=NBUCKETS**4, name="pd_fc")((p_fc, d_fc)) # Embedding with Embedding layer pd_embed = Embedding(input_dim=NBUCKETS**4, output_dim=10, name="pd_embed")( pd_fc ) # 移除无意义的单个张量拼接 # unk = Concatenate(axis=1)([pd_embed]) # 扩展一维经纬度特征的维度,从(None,)变为(None,1) plon_input = Reshape((1,))(inputs["pickup_longitude"]) plat_input = Reshape((1,))(inputs["pickup_latitude"]) dlon_input = Reshape((1,))(inputs["dropoff_longitude"]) dlat_input = Reshape((1,))(inputs["dropoff_latitude"]) # 按特征维度(axis=1)拼接所有特征 deep = Concatenate(name="deep_input", axis=1)( [ plon_input, plat_input, dlon_input, dlat_input, Flatten(name="flatten_embedding")(pd_embed), ] )
关键修改点说明
- 用
Reshape((1,))将每个经纬度的一维输入转换为二维的(None,1),保证和embedding展平后的(None,10)在样本维度(第一维)一致。 - 将
Concatenate的axis参数从0改为1,这样是在特征维度上拼接,最终得到形状为(None, 14)的特征张量(4个1维特征+10维embedding),可以直接输入到后续的DNN层中。
内容的提问来源于stack exchange,提问作者Jose_Peeterson
相关产品推荐
相关产品推荐

