自定义EV强化学习环境运行时Keras DQN输入维度不匹配报错求解
错误根因分析
1. 输入维度不匹配问题
- 你在DQN智能体中设置了
SequentialMemory(window_length=2),该参数会将连续2步的观测堆叠作为模型输入,原本单步观测维度是(2,)(包含电价、SOC两个值),堆叠后输入维度变为(2, 2),加上批处理维度后最终输入shape是(1, 2, 2) - 但你的模型输入层定义为
input_shape=states,也就是(2,),期望接收二维输入(None, 2),维度不一致直接触发报错
2. EVEnv环境代码逻辑错误
你的环境代码还存在多个语法和逻辑问题,会导致后续运行继续报错:
step方法中提前执行了return reward,后面的done判断、观测返回代码永远不会执行cost函数定义在__init__内部,且没有正确调用,step里直接使用cost变量会触发未定义报错- step里的
for self.soc in range (30,80)逻辑错误,会直接覆盖原有SOC值 - reset和step返回的观测是元组,不是numpy数组,和observation_space定义不匹配
解决方案
方案1:修改模型输入层适配窗口长度
如果你需要保留window_length=2的设置,直接修改模型输入层的形状即可:
from tensorflow.keras.layers import Flatten def build_model(states, actions): model = Sequential() # 输入形状适配堆叠后的观测,加Flatten层压平给全连接层 model.add(Flatten(input_shape=(2, states[0]))) model.add(Dense(24, activation='relu')) model.add(Dense(24, activation='relu')) model.add(Dense(actions, activation='linear')) return model
方案2:移除窗口长度设置(适合简单场景)
如果不需要堆叠多帧观测,直接把window_length设置为1即可,不需要修改模型:
memory = SequentialMemory(limit=50000, window_length=1)
修复EVEnv环境的错误代码
替换你的EVEnv代码为以下修复后的版本:
import numpy as np import random from gym import Env from gym.spaces import Discrete, Box class EVEnv(Env): def __init__(self): self.min_soc = 30 self.max_soc = 80 self.max_price = 5.0 self.min_price = 0.5 self.soc = 45 + random.randint(-10,10) self.price = 3.5 + random.randint(-3,3) self.low = np.array([self.min_price, self.min_soc], dtype=np.float32) self.high = np.array([self.max_price, self.max_soc], dtype=np.float32) self.action_space = Discrete(2) self.observation_space = Box(self.low, self.high, dtype=np.float32) def _calc_cost(self): return self.price * self.soc def step(self, action): # 动作对应充放电调整,可根据实际需求修改步长 if action == 1: self.soc = min(self.soc + 1, self.max_soc) else: self.soc = max(self.soc - 1, self.min_soc) cost = self._calc_cost() reward = 1 if 15 <= cost <=20 else -1 done = self.soc == self.max_soc info = {} return np.array([self.price, self.soc], dtype=np.float32), reward, done, info def reset(self): self.soc = 45 + random.randint(-10,10) self.price = 3.5 + random.randint(-3,3) return np.array([self.price, self.soc], dtype=np.float32)
内容的提问来源于stack exchange,提问作者Zonuna Fanai
相关产品推荐
相关产品推荐

