from typing import Dict from algorithm.base.algorithm_base import AlgorithmBase from pympler import asizeof import os import dotenv SCRIPT_DIR = os.path.dirname(__file__) # 脚本所在路径 dotenv.load_dotenv() # 加载脚本所在目录的.env环境变量 from algorithm.uf_rl.env.env_config_loader import EnvConfigLoader, create_env_params_from_yaml from algorithm.uf_rl.rl_model.DQN.uf_decide.run_dqn_decide import build_physics, replace, check_state_bounds # ========== 决策器 ========== from algorithm.uf_rl.rl_model.DQN.uf_decide.dqn_decider import UFDQNDecider import logging logger = logging.getLogger(__name__) class DQNDecide(AlgorithmBase): def __init__(self, params: Dict = None): super().__init__(params=None) # 静态配置参数 # ========== 模型及配置路径指定 ========== self.IS_TIMES = int(os.getenv("IS_TIMES", '0')) # 外部指定变量,表示CEB间隔为时间控制/次数控制,T表示48次bw一次CEB,F表示48h一次CEB self.PLANT = os.getenv("PLANT", None) self.MODEL_PATH = os.path.join(SCRIPT_DIR, "config_and_model", self.PLANT, "48times_dqn_model.zip" if self.IS_TIMES else "48h_dqn_model.zip") # 需根据IS_TIMES变量值指定模型为48h_dqn_model.zip/48times_dqn_model.zip self.ENV_CONFIG_PATH = os.path.join(SCRIPT_DIR, "config_and_model", self.PLANT, "env_config.yaml") # 环境配置路径 # ========== 模型及配置加载 ========== config_loader = EnvConfigLoader(self.ENV_CONFIG_PATH) config_loader.validate_config() config_loader.print_config_summary() ( self.uf_state_default, # UFState默认值 phys_params, # UFPhysicsParams action_spec, # UFActionSpec reward_params, # UFRewardParams self.state_bounds # UFStateBounds ) = create_env_params_from_yaml(self.ENV_CONFIG_PATH) # 构造决策器, 环境实例化,模型加载等功能放在UFDQNDecider类中 self.decider = UFDQNDecider( physics=build_physics(self.IS_TIMES, phys_params, self.state_bounds), action_spec=action_spec, reward_params=reward_params, state_bounds=self.state_bounds, model_path=self.MODEL_PATH, seed=0, ) def __call__(self, x: Dict, *args, **kwargs)->Dict: if not isinstance(x, dict): logger.warning("输入不为字典", x) return {} # ========== 外部调用输入 ========== # 轻量版,仅输入当前周期起始状态变量 units_to_run = x.get('units_to_run', None) # 新增输入:本次调用的机组对象名 TMP0 = x.get('TMP0', None) # 原始 TMP0 q_UF = x.get('q_UF', None) # 进水流量 temp = x.get('temp', None) # 进水温度 if None in [TMP0, q_UF, temp, units_to_run]: logger.warning("输入不合法", x) return {} # ========== 调用模型生成模型指令 ========== # 基于外部输入构建当前状态 current_state = replace( self.uf_state_default, TMP=TMP0, q_UF=q_UF, temp=temp ) # 状态异常检查(仅检查,不中断,出现异常时后续归一化中将异常状态强制归一化至上下限) for unit_name in units_to_run: error_result = check_state_bounds(current_state, self.state_bounds, unit_name) if error_result: print(f"错误发生时间: {error_result['error_time']};错误特征量:{error_result['error_feature']}") # 模型输出指令 decision = self.decider.decide(current_state) # 生成plc指令放到业务层 return { 'action_id': decision["action_id"], 'model_L_s': decision["action_id"], 'model_t_bw_s': decision["t_bw_s"], } def set(self): pass def get(self): pass def __del__(self): """显式执行清理""" pass if __name__ == "__main__": decider = DQNDecide() res = decider({ 'units_to_run': ["UF1"], # 新增输入:本次调用的机组对象名 'TMP0': 0.07, # 原始 TMP0 'q_UF': 300, # 进水流量 'temp': 20.0 # 进水温度 }) print(res)