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

如何创建支持双输入tf.matmul的Keras自定义层(修复batch维度报错)

问题描述

自定义Keras层实现矩阵计算逻辑时,因未适配批量输入维度规则触发维度不匹配报错,原实现代码如下:

class KeyQuery(keras.layers.Layer):
    def __init__(self, v):
        super(KeyQuery, self).__init__()
        self.v = tf.convert_to_tensor(v)
        
    def build(self, input_shape): 
        self.v = tf.Variable(self.v, trainable = True)
        print(self.v.shape)
    def call(self, inputs1, inputs2):
        y1 = tf.matmul(self.v, tf.transpose(inputs1))
        y2 = tf.matmul(y2, inputs2)
        return y2
    
keyquery = KeyQuery(v)

inputs1 = keras.Input(shape=(50,768))
inputs2 = keras.Input(shape=(50,3))
outputs = keyquery(inputs1,inputs2)
model = keras.Model([inputs1,inputs2], outputs)
model.summary()
  • 初始化层传入的参数v是尺寸为(1,768)的二维数组
  • 预期计算逻辑:
    • 单样本场景下:self.v形状为(1,768),inputs1形状为(50,768),矩阵乘后y1预期形状为(1,50);再与形状为(50,3)的inputs2做矩阵乘,最终输出形状为(1,3)
    • 带batch维度场景下:inputs1形状为(None,50,768),inputs2形状为(None,50,3),预期返回结果形状为(None,1,3)
  • 实际运行报错:

ValueError: Dimensions must be equal, but are 768 and 50 for '{{node key_query_4/MatMul}} = BatchMatMulV2[T=DT_FLOAT, adj_x=false, adj_y=false](key_query_4/MatMul/ReadVariableOp, key_query_4/transpose)' with input shapes: [1,768], [768,50,?].

报错根因
  1. 转置逻辑错误:直接调用tf.transpose(inputs1)会逆序所有维度,带batch的输入原本形状为(batch,50,768),全维度转置后会变成(768,50,batch),完全不符合批量矩阵乘的维度对齐要求。批量矩阵乘仅需转置最后两个特征维度,保留最前置的batch维度。
  2. 参数维度未适配批量广播:可训练参数self.v形状为(1,768),缺少batch对应的维度位,无法和带batch维度的输入做批量矩阵乘广播对齐。
  3. 代码笔误:计算y2时直接传入未定义的y2作为乘法输入,正确输入应为前一步计算得到的y1。
修复代码

调整转置规则、补充参数的批量广播维度、修正变量引用错误,修复后可正常运行,输出形状符合预期:

import tensorflow as tf
from tensorflow import keras

class KeyQuery(keras.layers.Layer):
    def __init__(self, v):
        super(KeyQuery, self).__init__()
        self.v_init = tf.convert_to_tensor(v, dtype=tf.float32)
        
    def build(self, input_shape): 
        self.v = tf.Variable(self.v_init, trainable=True)
        super().build(input_shape)
        
    def call(self, inputs1, inputs2):
        # 仅转置最后两个维度,保留batch维度:(batch,50,768) -> (batch,768,50)
        inputs1_t = tf.transpose(inputs1, perm=[0, 2, 1])
        # 给参数扩展batch维度,从(1,768)变为(1,1,768),可自动适配任意batch大小
        v_broadcast = self.v[None, ...]
        # 第一步批量矩阵乘:(batch,1,768) @ (batch,768,50) = (batch,1,50)
        y1 = tf.matmul(v_broadcast, inputs1_t)
        # 第二步批量矩阵乘:(batch,1,50) @ (batch,50,3) = (batch,1,3)
        y2 = tf.matmul(y1, inputs2)
        return y2

# 测试初始化
v = tf.random.normal(shape=(1,768))
keyquery = KeyQuery(v)

inputs1 = keras.Input(shape=(50,768))
inputs2 = keras.Input(shape=(50,3))
outputs = keyquery(inputs1,inputs2)
model = keras.Model([inputs1,inputs2], outputs)
model.summary()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.29 08:39:09