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

TensorFlow Keras中Transformer不同批次可变长度输入的掩码构建问题

解决Keras MultiHeadAttention中attention_mask的形状匹配问题

核心问题梳理

你当前的掩码设计混淆了序列长度维度和特征维度,导致形状不匹配。你的输入数据形状是[B, T, C](B=批量大小,T=序列长度=4,C=特征维度=3),而掩码应该针对无效的序列行(比如含nan的行),而非特征维度。

正确的掩码构建步骤

1. 生成基础序列掩码

首先从输入数据中识别无效序列行,生成形状为[B, T]的基础掩码(1表示有效行,0表示无效行):

import numpy as np
import tensorflow as tf

# 假设你的输入数据为data,形状[10,4,3]
mask = ~np.isnan(data).any(axis=-1)  # 一行中只要有nan就标记为无效(False)
mask = mask.astype(np.int32)  # 转换为int类型,最终形状[10,4]

2. 转换为MultiHeadAttention兼容的形状

自注意力场景下(Query=Key=Value),需要将基础掩码扩展为[B, T, T]或[B, num_heads, T, T],确保每个Query位置只能关注到有效的Key位置:

  • 方法一:生成三维掩码[B, T, T]
# 确保Query和Key位置都有效时才允许注意力计算
attention_mask = tf.expand_dims(mask, 1) & tf.expand_dims(mask, 2)  # 形状[B,T,T]
  • 方法二:生成四维掩码适配多头注意力(自动广播到所有head)
# 先得到[B,T,T],再扩展head维度
attn_mask_3d = tf.matmul(tf.expand_dims(mask, -1), tf.expand_dims(mask, 1))
attention_mask = tf.expand_dims(attn_mask_3d, 1)  # 形状[B,1,T,T],适配num_heads维度

模型适配修正

不能将掩码作为固定参数传入模型,需将其作为输入之一,与每个batch的样本一一对应:

def transformer_encoder(inputs, head_size, num_heads, ff_dim, dropout=0.0, mask=None):
    x = layers.LayerNormalization(epsilon=1e-6)(inputs)
    x = layers.MultiHeadAttention(
        key_dim=head_size, num_heads=num_heads, dropout=dropout
    )(x, x, attention_mask=mask)
    x = layers.Dropout(dropout)(x)
    res = x + inputs

    x = layers.LayerNormalization(epsilon=1e-6)(res)
    x = layers.Conv1D(filters=ff_dim, kernel_size=1, activation="relu")(x)
    x = layers.Dropout(dropout)(x)
    x = layers.Conv1D(filters=inputs.shape[-1], kernel_size=1)(x)
    return x + res


def build_model(
    n_classes,
    input_shape,
    head_size,
    num_heads,
    ff_dim,
    num_transformer_blocks,
    mlp_units,
    dropout=0.0,
    mlp_dropout=0.0,
) -> keras.Model:
    inputs = keras.Input(shape=input_shape)
    mask_input = keras.Input(shape=(input_shape[0],))  # 掩码输入,形状[T]
    
    x = inputs
    for _ in range(num_transformer_blocks):
        # 实时转换掩码为兼容形状
        attn_mask_3d = tf.matmul(tf.expand_dims(mask_input, -1), tf.expand_dims(mask_input, 1))
        attn_mask = tf.expand_dims(attn_mask_3d, 1)
        x = transformer_encoder(x, head_size, num_heads, ff_dim, dropout, mask=attn_mask)

    x = layers.GlobalAveragePooling2D(data_format="channels_first")(x)
    for dim in mlp_units:
        x = layers.Dense(dim, activation="relu")(x)
        x = layers.Dropout(mlp_dropout)(x)
    outputs = layers.Dense(n_classes, activation="softmax")(x)
    return keras.Model(inputs=[inputs, mask_input], outputs=outputs)

训练时的使用方式

将输入数据和对应掩码一起传入模型:

# 假设labels是你的标签数据
model.fit([data, mask], labels, batch_size=5, epochs=10)

关键注意点

  • 掩码始终针对序列长度维度,而非特征维度,不要生成[B,T,C]形状的掩码
  • 自注意力场景下,掩码需要覆盖Query和Key的所有位置组合,因此形状需扩展为[B,T,T]或适配多头的四维形状
  • 掩码必须与输入数据一一对应,每个batch的掩码随样本动态传入,不能预先固定

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.14 03:15:39