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

Keras中如何实现基于样本级条件的动态拼接操作?

Keras实现基于样本条件的动态拼接方案

核心思路

要实现按样本条件调整拼接顺序,关键是用Keras的Lambda层结合TensorFlow的tf.where函数,针对批次内每个样本的条件值,动态选择拼接顺序。需要额外传入一个条件输入张量(每个样本对应一个条件标记,比如0或1),用来控制拼接逻辑。

具体实现代码

import tensorflow as tf
from tensorflow.keras.layers import Input, Dense, Lambda, Concatenate
from tensorflow.keras.models import Model

# 1. 定义输入:主输入(生成器输出的128维向量,拆分两个64维分支)+ 条件输入
main_input = Input(shape=(128,), name='main_input')
# 拆分主输入为两个64维分支
x1 = Lambda(lambda x: x[:, :64])(main_input)
x2 = Lambda(lambda x: x[:, 64:])(main_input)
# 条件输入:每个样本对应一个标记(0或1,0表示h1在前,1表示h2在前)
condition_input = Input(shape=(1,), name='condition_input', dtype='int32')

# 2. 双分支线性层处理
h1 = Dense(64, activation=None, name='branch1_dense')(x1)
h2 = Dense(64, activation=None, name='branch2_dense')(x2)

# 3. 实现条件拼接的Lambda层
def conditional_concat(h1, h2, condition):
    # 先准备两种拼接结果:h1在前 和 h2在前
    concat_order1 = Concatenate(axis=-1)([h1, h2])
    concat_order2 = Concatenate(axis=-1)([h2, h1])
    # 把条件转为布尔型,tf.where会根据每个样本的条件选择对应拼接结果
    condition_bool = tf.cast(condition, tf.bool)
    # 扩展条件张量维度,和拼接结果维度对齐
    condition_bool = tf.tile(condition_bool, [1, tf.shape(concat_order1)[1]])
    return tf.where(condition_bool, concat_order2, concat_order1)

# 包装成Lambda层
conditional_concat_layer = Lambda(lambda inputs: conditional_concat(*inputs))([h1, h2, condition_input])

# 4. 后续全连接层处理
fc1 = Dense(128, activation=None)(conditional_concat_layer)
fc2 = Dense(64, activation=None)(fc1)
output = Dense(你的输出维度, activation=None)(fc2)

# 构建模型
model = Model(inputs=[main_input, condition_input], outputs=output)
model.summary()

关键细节说明

  • 条件输入的形状:条件输入是形状为(batch_size, 1)的张量,每个元素对应一个样本的拼接规则(0或1)。
  • tf.where的作用:逐元素检查条件张量,为每个样本选择对应的拼接结果,完美适配批次处理场景。
  • 生成器适配:生成器需要同时输出主输入向量和条件标记,比如每个批次返回([main_batch, condition_batch], label_batch),确保模型能接收对应输入。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.18 04:30:30