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

推理阶段Attention层批量尺寸异常排查与高效推理优化求助

批量推理时Encoder-Decoder模型Attention层批量尺寸不匹配问题排查与优化方案

问题描述

基于Encoder-Decoder架构训练seq-to-seq模型,批量推理时需处理输入上下文向量以提升生产效率,但Decoder的Attention层出现批量尺寸不匹配错误:传入批量大小为64的输入,Attention层却收到批量大小为32的张量,导致维度兼容报错。单样本或批量大小为1时推理正常,附实现代码与错误信息,请求排查问题并提供高效的输出序列生成方案。

实现代码

### Define Inference Encoder
def define_inference_encoder(input_shape):
  encoder_input = Input(shape=input_shape, name='en_input_layer')

  ### First Bidirectional GRU Layer
  bidi_gru1 = Bidirectional(GRU(160, return_sequences=True), name='en_bidirect_gru1')
  gru1_out = bidi_gru1(encoder_input)

  gru1_out = Dropout(0.46598303573163413, name='bidirect_gru1_dropout')(gru1_out)

  ### Second GRU Layer
  # hp_units_2 = hp.Int('enc_lstm2', min_value=32, max_value=800, step=32)
  gru2 = GRU(hsize, return_sequences=True, return_state=True, name='en_gru2_layer')
  gru2_out, gru2_states = gru2(gru1_out)

  encoder_model = Model(inputs=encoder_input, outputs=[gru2_out, gru2_states])
  return encoder_model

### Define Inference Decoder
def define_inf_decoder(context_vec, input_shape):
  decoder_input = Input(shape=input_shape)
  decoder_state_input = Input(shape=(hsize,))

  de_gru1 = GRU(hsize, return_sequences=True, return_state=True, name='de_gru1_layer')
  de_gru1_out, de_state_out = de_gru1(decoder_input, initial_state=decoder_state_input)

  attn_layer = Attention(use_scale=True, name='attn_layer')
  attn_out = attn_layer([de_gru1_out, context_vec])

  attn_added = Concatenate(name='attn_source_concat_layer')([de_gru1_out, attn_out])

  attn_dense_layer = Dense(736, name='tanh_dense_layer', activation='tanh')

  h_hat = attn_dense_layer(attn_added)

  ### Output Layer
  preds = Dense(1, name='output_layer')(h_hat)

  decoder_model = Model(inputs=[decoder_input, decoder_state_input], outputs=[preds, de_state_out])
  return decoder_model

def set_weights(untrained_model, trained_model):
  trained_layers = [l.name for l in trained_model.layers]
  print(f"No. of trained layers: {len(trained_layers)}")

  for l in untrained_model.layers:
    if l.name in trained_layers:
      trained_wts = trained_model.get_layer(l.name).get_weights()
      if len(trained_wts)>0:
        untrained_model.get_layer(l.name).set_weights(trained_wts)
        print(f"Layer {l.name} weight set")
        
  return untrained_model

生成输出序列代码

inference_encoder = define_inference_encoder((12, 1))
inference_encoder = set_weights(inference_encoder, tuned_model)

for (ex_context, ex_target_in), ex_target_out in test_ds.take(1):
  print(ex_context.shape, ex_target_in.shape) ### (64, 12, 1) (64, 3, 1)

test_context, test_states = inference_encoder.predict(tf.reshape(ex_context, shape=(-1,seq_len, 1)))
print(test_context.shape, test_states.shape) ### (64, 12, 256) (64, 256)

inf_decoder = define_inf_decoder(test_context, (1,1))
inf_decoder = set_weights(inf_decoder, tuned_model)

dec_inp = tf.reshape(ex_context[:,-1], shape=(-1,1,1))
dec_inp.shape ### (64,1,1)

test_inf_decoder_out = inf_decoder.predict([dec_inp, test_states])

错误信息

ValueError: Exception encountered when calling layer 'attn_layer' (type Attention).

Dimensions must be equal, but are 32 and 64 for '{{node model_7/attn_layer/MatMul}} = BatchMatMulV2[T=DT_FLOAT, adj_x=false,

adj_y=true](model_7/de_gru1_layer/PartitionedCall:1, model_7/15181)'
with input shapes: [32,1,256], [64,12,256].

Call arguments received by layer 'attn_layer' (type Attention):
  • inputs=['tf.Tensor(shape=(32, 1, 256), dtype=float32)', 'tf.Tensor(shape=(64, 12, 256), dtype=float32)']
  • mask=None
  • training=False
  • return_attention_scores=False
  • use_causal_mask=False

问题分析与修复方案

核心问题

推理阶段定义Decoder时,直接传入了具体的批量上下文张量test_context,导致Decoder模型被固定绑定了该张量的批量维度(64)。但TensorFlow在运行时会自动进行设备并行(如多GPU拆分批量),将输入批量拆分为32,此时Decoder中硬编码的64批量上下文张量与拆分后的32批量输入维度不匹配,引发报错。

修复代码

修改define_inf_decoder函数,通过Input层接收上下文向量,而非传入具体张量,让模型支持动态批量:

### Define Inference Decoder
def define_inf_decoder(input_shape, context_shape):
  decoder_input = Input(shape=input_shape)
  decoder_state_input = Input(shape=(hsize,))
  # 新增上下文输入层,支持动态批量
  context_input = Input(shape=context_shape)

  de_gru1 = GRU(hsize, return_sequences=True, return_state=True, name='de_gru1_layer')
  de_gru1_out, de_state_out = de_gru1(decoder_input, initial_state=decoder_state_input)

  attn_layer = Attention(use_scale=True, name='attn_layer')
  # 使用上下文输入层而非固定张量
  attn_out = attn_layer([de_gru1_out, context_input])

  attn_added = Concatenate(name='attn_source_concat_layer')([de_gru1_out, attn_out])

  attn_dense_layer = Dense(736, name='tanh_dense_layer', activation='tanh')

  h_hat = attn_dense_layer(attn_added)

  ### Output Layer
  preds = Dense(1, name='output_layer')(h_hat)

  # 上下文输入纳入模型输入列表
  decoder_model = Model(inputs=[decoder_input, decoder_state_input, context_input], outputs=[preds, de_state_out])
  return decoder_model

修改生成序列代码,传入上下文的shape而非具体张量,推理时将上下文作为输入传入:

inference_encoder = define_inference_encoder((12, 1))
inference_encoder = set_weights(inference_encoder, tuned_model)

for (ex_context, ex_target_in), ex_target_out in test_ds.take(1):
  print(ex_context.shape, ex_target_in.shape) ### (64, 12, 1) (64, 3, 1)

test_context, test_states = inference_encoder.predict(tf.reshape(ex_context, shape=(-1,seq_len, 1)))
print(test_context.shape, test_states.shape) ### (64, 12, 256) (64, 256)

# 传入上下文的shape而非具体张量
inf_decoder = define_inf_decoder((1,1), test_context.shape[1:])
inf_decoder = set_weights(inf_decoder, tuned_model)

dec_inp = tf.reshape(ex_context[:,-1], shape=(-1,1,1))
dec_inp.shape ### (64,1,1)

# 推理时传入上下文作为第三个输入
test_inf_decoder_out = inf_decoder.predict([dec_inp, test_states, test_context])

优化的输出序列生成方案

  1. 动态批量适配:使用tf.data.Dataset的batch方法动态调整批量大小,结合修改后的Decoder模型,适配不同硬件的并行处理能力。
  2. 预计算编码器输出:对整个测试集预计算所有上下文向量和状态,避免重复调用编码器,大幅提升批量推理效率。
  3. 向量化生成策略:采用贪心搜索(每步选择概率最高的输出)或束搜索(保留Top-K候选序列),利用TensorFlow的向量化操作加速批量生成,避免逐样本循环。
  4. 混合精度推理:开启TensorFlow混合精度模式(tf.keras.mixed_precision.set_global_policy('mixed_float16')),减少内存占用,提升批量处理速度。

内容的提问来源于stack exchange,提问作者Krishnang K Dalal

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.04 18:34:55