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

