如何在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

