使用magenta中tf.nn.dynamic_rnn时出现维度不匹配错误求助
解决Dynamic RNN输入与初始状态维度不匹配的问题
这个错误其实是批量大小(batch size)不匹配导致的,咱们一步步拆解来看:
错误根源解析
你看到的报错ConcatOp : Dimensions of inputs should match: shape[0] = [1,38] vs. shape[1] = [128,512],这里的第一个维度是批量大小,第二个维度分别是输入特征数(38)和LSTM隐藏层单元数(512):
- 输入
m_seq的形状[1,38]:1是当前输入的批量大小,38是输入的独热向量维度 - 初始状态中
c/h的形状[128,512]:128是状态对应的批量大小,512是隐藏层单元数
Dynamic RNN要求输入的批量大小必须和初始状态的批量大小完全一致——因为每个样本的状态要对应到输入里的每个样本,这样RNN才能逐个处理batch里的序列。而输入特征数(38)和隐藏层单元数(512)根本不需要匹配,RNN内部的权重矩阵会自动完成维度转换,这不是报错的原因。
解决方法
你只需要把初始状态的批量大小调整为和输入一致的1即可,分两种情况处理:
情况1:初始状态是预加载的(比如从训练好的模型中读取)
可以对状态的批量维度做切片,只取其中一个样本的状态:
# 假设initial_state是包含两个LSTMStateTuple的元组,每个的batch size是128 new_initial_state = tuple( tf.nn.rnn_cell.LSTMStateTuple(c=state.c[0:1], h=state.h[0:1]) for state in initial_state ) # 然后用new_initial_state作为dynamic_rnn的initial_state参数 outputs, final_state = tf.nn.dynamic_rnn( self.cell, m_seq, sequence_length=lens, initial_state=new_initial_state, swap_memory=swap_memory, parallel_iterations=parallel_iterations)
情况2:初始化全新的状态
直接调用cell的zero_state方法时指定正确的batch size:
initial_state = self.cell.zero_state(batch_size=1, dtype=tf.float32)
总结
记住这个关键点:Dynamic RNN只要求输入的批量大小和初始状态的批量大小一致,输入特征维度和隐藏层维度不需要相同——RNN本身就是用来处理不同维度的输入到隐藏层的转换的,所以不用纠结38和512的差异。
内容的提问来源于stack exchange,提问作者Rodrigo Mira
相关产品推荐
相关产品推荐

