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

如何在Keras函数式API中实现广播加法?Add层失效原因及解决

问题描述

我在Keras中定义了如下TextDec模型类:

class TextDec(tf.keras.Model):
    def __init__(
        self,
        n_vocab: int,
        n_ctxs: int,
        n_states: int,
        n_heads: int,
        n_layers: int,
    ):
        super(TextDec, self).__init__(name="dec")

        self.token_emb = tf.keras.layers.Embedding(
            n_vocab, n_states, name="dec-token-emb"
        )

        self.position_emb = tf.Variable(
            np.zeros((n_ctxs, n_states)),
            name="dec-position-emb",
            dtype=tf.float32,
        )

        self.add = tf.keras.layers.Add(name="dec-add")

        # ...

    def call(self, inputs: List[tf.Tensor]):
        tokens, audio_features = inputs
        offset = 0

        x = self.add(
            [
                self.token_emb(tokens),
                tf.slice(self.position_emb, [offset, 0], [tokens.shape[-1], -1]),
            ]
        )

        # ...

        return x

当设置n_ctxs=448、n_states=512调用该模型时,出现报错:

ValueError: Cannot merge tensors with different batch sizes. Got tensors with shapes [(5, 1, 512), (1, 512)]

我尝试使用tf.expand_dims后仍报错:

x = self.add([ self.token_emb(tokens),  tf.expand_dims(tf.slice(self.position_emb, [offset, 0], [tokens.shape[-1], -1]), 0) ])

ValueError: Cannot merge tensors with different batch sizes. Got tensors with shapes [(5, 1, 512), (1, 1, 512)]

还尝试了已弃用的tf.compat.v1.placeholder_with_default,同样失败:

z = tf.slice(self.position_emb, [offset, 0], [tokens.shape[-1], -1])
z = tf.compat.v1.placeholder_with_default(z, [None, 1, 512])
x = self.add([ self.token_emb(tokens),  z ])

ValueError: Shapes must be equal rank, but are 2 and 3 for '{{node PlaceholderWithDefault}} = PlaceholderWithDefaultdtype=DT_FLOAT, shape=[?,1,512]' with input shapes: [1,512].

tf.broadcast_to因形状含None也无法使用。

但直接使用+运算符时却能正常运行:

x = self.token_emb(tokens) + tf.slice(self.position_emb, [offset, 0], [tokens.shape[-1], -1])

请问为何使用Add层无法实现广播加法,该如何解决?我希望使用Add层是因为模型摘要中会显示该层,而使用+运算符会出现tf.__operators__.add (TFOpLambda),担心影响模型序列化。


问题原因与解决方案

原因

Keras原生Add层要求输入张量的静态形状完全匹配,不会自动触发广播逻辑;而+是TensorFlow原生操作,只要张量动态形状符合广播规则就会自动执行广播计算,这是两者核心差异。

解决方案

以下两种方式可同时满足保留Add层显式结构、实现广播加法的需求:

方法1:自定义支持广播的加法层

继承tf.keras.layers.Layer实现自定义层,内部用+完成广播,同时保持层的可识别性:

class BroadcastAdd(tf.keras.layers.Layer):
    def __init__(self, name=None):
        super().__init__(name=name)
    
    def call(self, inputs):
        return inputs[0] + inputs[1]

在TextDec类的__init__方法中替换原有Add层:

self.add = BroadcastAdd(name="dec-add")

该方式下模型摘要会显示BroadcastAdd层,且不影响序列化。

方法2:手动对齐输入张量形状

在传入Add层前,将位置嵌入扩展为与词嵌入完全匹配的形状,通过动态获取批量大小并重复张量实现:

def call(self, inputs: List[tf.Tensor]):
    tokens, audio_features = inputs
    offset = 0
    
    token_emb_out = self.token_emb(tokens)
    # 动态获取批量大小
    batch_size = tf.shape(token_emb_out)[0]
    # 切片位置嵌入并扩展维度,再按批量大小重复
    pos_emb_slice = tf.slice(self.position_emb, [offset, 0], [tokens.shape[-1], -1])
    pos_emb_expanded = tf.expand_dims(pos_emb_slice, 0)
    pos_emb_broadcasted = tf.repeat(pos_emb_expanded, repeats=batch_size, axis=0)
    
    # 使用原生Add层计算
    x = self.add([token_emb_out, pos_emb_broadcasted])
    # ...
    return x

此方法让两个输入张量静态形状完全匹配,原生Add层可正常工作,同时保留原层结构。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.07 08:35:20