dqn_decide.py 4.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113
  1. from typing import Dict
  2. from algorithm.base.algorithm_base import AlgorithmBase
  3. from pympler import asizeof
  4. import os
  5. import dotenv
  6. SCRIPT_DIR = os.path.dirname(__file__) # 脚本所在路径
  7. dotenv.load_dotenv() # 加载脚本所在目录的.env环境变量
  8. from algorithm.uf_rl.env.env_config_loader import EnvConfigLoader, create_env_params_from_yaml
  9. from algorithm.uf_rl.rl_model.DQN.uf_decide.run_dqn_decide import build_physics, replace, check_state_bounds
  10. # ========== 决策器 ==========
  11. from algorithm.uf_rl.rl_model.DQN.uf_decide.dqn_decider import UFDQNDecider
  12. import logging
  13. logger = logging.getLogger(__name__)
  14. class DQNDecide(AlgorithmBase):
  15. def __init__(self, params: Dict = None):
  16. super().__init__(params=None)
  17. # 静态配置参数
  18. # ========== 模型及配置路径指定 ==========
  19. self.IS_TIMES = int(os.getenv("IS_TIMES", '0')) # 外部指定变量,表示CEB间隔为时间控制/次数控制,T表示48次bw一次CEB,F表示48h一次CEB
  20. self.PLANT = os.getenv("PLANT", None)
  21. 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
  22. self.ENV_CONFIG_PATH = os.path.join(SCRIPT_DIR, "config_and_model", self.PLANT, "env_config.yaml") # 环境配置路径
  23. # ========== 模型及配置加载 ==========
  24. config_loader = EnvConfigLoader(self.ENV_CONFIG_PATH)
  25. config_loader.validate_config()
  26. config_loader.print_config_summary()
  27. (
  28. self.uf_state_default, # UFState默认值
  29. phys_params, # UFPhysicsParams
  30. action_spec, # UFActionSpec
  31. reward_params, # UFRewardParams
  32. self.state_bounds # UFStateBounds
  33. ) = create_env_params_from_yaml(self.ENV_CONFIG_PATH)
  34. # 构造决策器, 环境实例化,模型加载等功能放在UFDQNDecider类中
  35. self.decider = UFDQNDecider(
  36. physics=build_physics(self.IS_TIMES, phys_params, self.state_bounds),
  37. action_spec=action_spec,
  38. reward_params=reward_params,
  39. state_bounds=self.state_bounds,
  40. model_path=self.MODEL_PATH,
  41. seed=0,
  42. )
  43. def __call__(self, x: Dict, *args, **kwargs)->Dict:
  44. if not isinstance(x, dict):
  45. logger.warning("输入不为字典", x)
  46. return {}
  47. # ========== 外部调用输入 ==========
  48. # 轻量版,仅输入当前周期起始状态变量
  49. units_to_run = x.get('units_to_run', None) # 新增输入:本次调用的机组对象名
  50. TMP0 = x.get('TMP0', None) # 原始 TMP0
  51. q_UF = x.get('q_UF', None) # 进水流量
  52. temp = x.get('temp', None) # 进水温度
  53. if None in [TMP0, q_UF, temp, units_to_run]:
  54. logger.warning("输入不合法", x)
  55. return {}
  56. # ========== 调用模型生成模型指令 ==========
  57. # 基于外部输入构建当前状态
  58. current_state = replace(
  59. self.uf_state_default,
  60. TMP=TMP0,
  61. q_UF=q_UF,
  62. temp=temp
  63. )
  64. # 状态异常检查(仅检查,不中断,出现异常时后续归一化中将异常状态强制归一化至上下限)
  65. for unit_name in units_to_run:
  66. error_result = check_state_bounds(current_state, self.state_bounds, unit_name)
  67. if error_result:
  68. print(f"错误发生时间: {error_result['error_time']};错误特征量:{error_result['error_feature']}")
  69. # 模型输出指令
  70. decision = self.decider.decide(current_state)
  71. # 生成plc指令放到业务层
  72. return {
  73. 'action_id': decision["action_id"],
  74. 'model_L_s': decision["action_id"],
  75. 'model_t_bw_s': decision["t_bw_s"],
  76. }
  77. def set(self):
  78. pass
  79. def get(self):
  80. pass
  81. def __del__(self):
  82. """显式执行清理"""
  83. pass
  84. if __name__ == "__main__":
  85. decider = DQNDecide()
  86. res = decider({
  87. 'units_to_run': ["UF1"], # 新增输入:本次调用的机组对象名
  88. 'TMP0': 0.07, # 原始 TMP0
  89. 'q_UF': 300, # 进水流量
  90. 'temp': 20.0 # 进水温度
  91. })
  92. print(res)