使用TF-Agents训练井字棋DQN时遭遇维度不兼容错误
错误原因解析
1. 观测空间(Observation Spec)定义不符合DQN输入要求
你自定义的TicTacToeEnvironment里,观测空间的维度设置错误。DQN的QNetwork默认要求输入是批量维度+特征维度的2维张量(例如(batch_size, feature_size)),但你的观测输出是(1,)形状的标量张量,缺少特征维度。当批量处理50个样本时,输入就变成了(50,)的1维张量,而全连接层(dense_2)需要至少2维输入,因此触发维度不兼容错误。
比如井字棋的观测应该对应9个格子的状态,正确的观测空间应定义为tf.TensorSpec(shape=(9,), dtype=tf.int32),如果错误定义成标量(shape=())或者错误编码成单个整数输出,就会出现这个问题。
2. 观测输出未正确保留特征维度
即使观测空间定义正确,环境输出观测时可能错误压缩了维度。比如把9格棋盘的状态编码成一个整数输出,导致观测形状变成(1,),而非(9,)的特征张量。传入QNetwork后,批量数据的形状就变成(batch_size,),不满足全连接层对输入维度的要求。
3. QNetwork初始化参数不匹配观测空间
初始化QNetwork时,如果没有正确传入观测空间的spec,或者手动设置了错误的输入层形状,会导致网络内部层结构与实际输入维度不匹配。比如观测空间是(9,),但QNetwork的输入层被设置为期望(1,),就会引发维度错误。
4. 数据预处理阶段丢失维度
在训练数据的生成、回放缓冲区(ReplayBuffer)或迭代器处理过程中,可能不小心压缩了观测的特征维度。例如没有正确设置batch_size相关的参数,导致输出的观测张量丢失了特征维度,变成1维的批量数据。
内容的提问来源于stack exchange,提问作者SXKA
相关产品推荐
相关产品推荐

