模型返回2D张量时的1D概率处理及强化学习代码错误修复
问题解决:Kerath Actor-Critic代码的维度错误与优化
错误原因分析
核心问题是状态输入未做扁平处理:
- 输入模型的
state是(1,10,10)的3D张量,Dense层默认对最后一维做变换,导致输出的x_probs为(1,10,10)的二维概率分布,但np.random.choice要求传入一维概率数组,因此触发ValueError。 - 额外错误:选择y和tile_type时误用
max_x,应分别替换为max_y和max_tile_type。
修改后的完整代码
import tensorflow as tf from keras import layers import numpy as np class CustomModel(tf.keras.Model): def __init__(self, num_hidden, max_x, max_y, n_tile_type): super(CustomModel, self).__init__() self.flatten = layers.Flatten() self.common = layers.Dense(num_hidden, activation="relu") self.x_probs = layers.Dense(max_x, activation="softmax") self.y_probs = layers.Dense(max_y, activation="softmax") self.tile_prob = layers.Dense(n_tile_type, activation="softmax") def call(self, inputs): # 将二维board扁平为一维向量 x = self.flatten(inputs) common = self.common(x) x_probs = self.x_probs(common) y_probs = self.y_probs(common) tile_prob = self.tile_prob(common) return [x_probs, y_probs, tile_prob] max_x = 10 max_y = 10 max_tile_type = 5 model = CustomModel(512, max_x, max_y, max_tile_type) board = np.full((max_y, max_x), -1) # 状态归一化:将board的-1转换为0,统一到0~1范围 def normalize_state(state): return np.where(state == -1, 0, state) / max_tile_type while True: state = board state = normalize_state(state) # 加入状态归一化 state = tf.convert_to_tensor(state, dtype=tf.float32) state = tf.expand_dims(state, 0) # 添加batch维度 x_probs, y_probs, tile_probs = model(state) # 提取一维概率数组,满足np.random.choice要求 x = np.random.choice(max_x, p=x_probs.numpy()[0]) y = np.random.choice(max_y, p=y_probs.numpy()[0]) tile_type = np.random.choice(max_tile_type, p=tile_probs.numpy()[0]) board[x][y] = tile_type # 填满board后终止循环 if np.all(board != -1): break
关键修改点说明
- 输入扁平处理:在模型中加入
layers.Flatten(),将二维board转为一维向量,确保后续Dense层输出一维概率分布(形状(1, max_x))。 - 参数修正:选择y时用
max_y,选择tile_type时用max_tile_type,避免维度不匹配。 - 状态归一化:新增
normalize_state函数,将board中的-1替换为0并归一化到0~1区间,符合神经网络输入要求。 - 张量转numpy:通过
.numpy()[0]提取batch内的一维概率数组,满足np.random.choice的参数格式。
反向传播补充
原代码缺少训练逻辑,若要实现Actor-Critic的反向传播,需:
- 定义损失函数:Actor损失用策略梯度(对数概率乘优势函数),Critic损失用均方误差(预测值与真实回报的差值)。
- 优化器:选用
tf.keras.optimizers.Adam等,每步计算损失后执行梯度更新。 - 轨迹记录:保存每一步的状态、动作、回报数据,用于计算优势函数和更新模型。
内容的提问来源于stack exchange,提问作者donghyunlee
相关产品推荐
相关产品推荐

