| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113 |
- 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)
|