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

如何在TensorFlow中结合Flatten层实现特征拼接?

问题解决方法

你的错误根源有两个:

  1. 拼接轴设置错误:你用了axis=0(样本维度)拼接,但输入特征维度不匹配;应该用axis=1(特征维度)来拼接不同特征。
  2. 特征维度不一致:经纬度输入是(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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.18 06:45:33