Compare commits
28
Commits
7fcc809852
...
master
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
24a8999b18 | ||
|
|
2bb4207a98 | ||
|
|
646399dcc7 | ||
|
|
88f4896086 | ||
|
|
55d6e8708e | ||
|
|
f582e72151 | ||
|
|
1e99bf1b8c | ||
|
|
f93d4bb119 | ||
|
|
0932bb353a | ||
|
|
3dfe1aa673 | ||
|
|
845ca708a7 | ||
|
|
2c1bbc1a31 | ||
|
|
041e0ec1bd | ||
|
|
36a33e74e5 | ||
|
|
f4d45d3cfd | ||
|
|
1b93699501 | ||
|
|
65190dffea | ||
|
|
6cb93ad56d | ||
|
|
e2e8db1f04 | ||
|
|
7ee8272034 | ||
|
|
f0cc7ba9c4 | ||
|
|
3eb0cc7b60 | ||
|
|
a4f898c3ad | ||
|
|
c3111ad5be | ||
|
|
088b7d4733 | ||
|
|
ce2019e060 | ||
|
|
1f7ecc301f | ||
|
|
0dab7a6cec |
@@ -1,15 +1,12 @@
|
|||||||
<div align="center">
|
<div align="center">
|
||||||
<img src='./logo.png' width="250px">
|
<img src='./logo.svg' width="250px">
|
||||||
<h2>NuCon</h2>
|
<h2>NuCon</h2>
|
||||||
<br>
|
<br>
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
NuCon (Nucleares Controller) is a Python library designed to interface with and control parameters in [Nucleares](https://store.steampowered.com/app/1428420/Nucleares/), a nuclear reactor simulation game. It provides a robust, type-safe foundation for reading and writing game parameters, allowing users to easily create their own automations and control systems.
|
NuCon (Nucleares Controller) is a Python library designed to interface with and control parameters in [Nucleares](https://store.steampowered.com/app/1428420/Nucleares/), a nuclear reactor simulation game. It provides a robust, type-safe foundation for reading and writing game parameters, allowing users to easily create their own automations and control systems.
|
||||||
|
|
||||||
NuCon further provides a work in progress implementation of a reinforcement learning environment for training control policies and a simulator based on model learning.
|
NuCon further provides a reinforcement learning environment for training control policies and a simulator based on model learning.
|
||||||
|
|
||||||
> [!NOTE]
|
|
||||||
> Nucleares only exposes RODS_POS_ORDERED as writable parameter, and no parameters about core chemistry e.g. Xenon concentration. While NuCon is already usable, it's capabilities are still very limited based on these restrictions. The capabilites are supposed to be extended in future updates to Nucleares, development on the advanced features (Reinforcement / Model Learning) are paused till then.
|
|
||||||
|
|
||||||
## Features
|
## Features
|
||||||
|
|
||||||
@@ -109,11 +106,11 @@ Custom Enum Types:
|
|||||||
|
|
||||||
\*: Truthy value (will be treated as true in e.g. if statements).
|
\*: Truthy value (will be treated as true in e.g. if statements).
|
||||||
|
|
||||||
So if you're not in the mood to play the game manually, this API can be used to easily create your own automations and control systems. Maybe a little PID controller for the rods? Or, if you wanna go crazy, why not try some
|
So if you're not in the mood to play the game manually, this API can be used to easily create your own automations and control systems. Maybe a little PID controller for the rods — or a full classical reactor operator with grid-demand following, pressurizer control, and a live TUI, like the one in `scripts/reactor_control.py`? Or, if you wanna go crazy, why not try some
|
||||||
|
|
||||||
## Reinforcement Learning (Work in Progress)
|
## Reinforcement Learning
|
||||||
|
|
||||||
NuCon includes a preliminary Reinforcement Learning (RL) environment based on the OpenAI Gym interface. This allows you to train control policies for the Nucleares game instead of writing them yourself. This feature is currently a work in progress and requires additional dependencies.
|
NuCon includes a Reinforcement Learning (RL) environment based on the OpenAI Gym interface. This allows you to train control policies for the Nucleares game instead of writing them yourself. Requires additional dependencies.
|
||||||
|
|
||||||
### Additional Dependencies
|
### Additional Dependencies
|
||||||
|
|
||||||
@@ -123,18 +120,24 @@ To use you'll need to install `gymnasium` and `numpy`. You can do so via
|
|||||||
pip install -e '.[rl]'
|
pip install -e '.[rl]'
|
||||||
```
|
```
|
||||||
|
|
||||||
### RL Environment
|
### Environments
|
||||||
|
|
||||||
The `NuconEnv` class in `nucon/rl.py` provides a Gym-compatible environment for reinforcement learning tasks in the Nucleares simulation. Key features include:
|
Two environment classes are provided in `nucon/rl.py`:
|
||||||
|
|
||||||
- Observation space: Includes all readable parameters from the NuCon system.
|
**`NuconEnv`**: classic fixed-objective environment. You define one or more objectives at construction time (e.g. maximise power output, keep temperature in range). The agent always trains toward the same goal.
|
||||||
- Action space: Encompasses all writable parameters in the NuCon system.
|
|
||||||
- Step function: Applies actions to the NuCon system and returns new observations.
|
|
||||||
- Objective function: Allows for predefined or custom objective functions to be defined for training.
|
|
||||||
|
|
||||||
### Usage
|
- Observation space: all readable numeric parameters (~290 dims).
|
||||||
|
- Action space: all readable-back writable parameters (~30 dims): 9 individual rod bank positions, 3 MSCVs, 3 turbine bypass valves, 6 coolant pump speeds, condenser pump, freight/vent switches, resistor banks, and more.
|
||||||
|
- Objectives: predefined strings (`'max_power'`, `'episode_time'`) or arbitrary callables `(obs) -> float`. Multiple objectives are weighted-summed.
|
||||||
|
|
||||||
|
**`NuconGoalEnv`**: goal-conditioned environment. The desired goal (e.g. target generator output) is sampled at the start of each episode and provided as part of the observation. A single policy learns to reach *any* goal in the specified range, making it far more useful than a fixed-objective agent. Designed for training with [Hindsight Experience Replay (HER)](https://arxiv.org/abs/1707.01495), which makes sparse-reward goal-conditioned training tractable.
|
||||||
|
|
||||||
|
- Observation space: `Dict` with keys `observation` (non-goal params), `achieved_goal` (current goal param values, normalised to [0,1]), `desired_goal` (target, normalised to [0,1]).
|
||||||
|
- Goals are sampled uniformly from the specified `goal_range` each episode.
|
||||||
|
- Reward defaults to negative L2 distance in normalised goal space (dense). Pass `tolerance` for a sparse `{0, -1}` reward; this works particularly well with HER.
|
||||||
|
|
||||||
|
### NuconEnv Usage
|
||||||
|
|
||||||
Here's a basic example of how to use the RL environment:
|
|
||||||
```python
|
```python
|
||||||
from nucon.rl import NuconEnv, Parameterized_Objectives
|
from nucon.rl import NuconEnv, Parameterized_Objectives
|
||||||
|
|
||||||
@@ -154,47 +157,97 @@ env.close()
|
|||||||
|
|
||||||
Objectives takes either strings of the name of predefined objectives, or lambda functions which take an observation and return a scalar reward. Final rewards are (weighted) summed across all objectives. `info['objectives']` contains all objectives and their values.
|
Objectives takes either strings of the name of predefined objectives, or lambda functions which take an observation and return a scalar reward. Final rewards are (weighted) summed across all objectives. `info['objectives']` contains all objectives and their values.
|
||||||
|
|
||||||
You can e.g. train an PPO agent using the [sb3](https://github.com/DLR-RM/stable-baselines3) implementation:
|
You can e.g. train a PPO agent using the [sb3](https://github.com/DLR-RM/stable-baselines3) implementation:
|
||||||
```python
|
```python
|
||||||
from nucon.rl import NuconEnv
|
from nucon.rl import NuconEnv
|
||||||
from stable_baselines3 import PPO
|
from stable_baselines3 import PPO
|
||||||
|
|
||||||
env = NuconEnv(objectives=['max_power'], seconds_per_step=5)
|
env = NuconEnv(objectives=['max_power'], seconds_per_step=5)
|
||||||
|
|
||||||
# Create the PPO (Proximal Policy Optimization) model
|
|
||||||
model = PPO(
|
model = PPO(
|
||||||
"MlpPolicy",
|
"MlpPolicy",
|
||||||
env,
|
env,
|
||||||
verbose=1,
|
verbose=1,
|
||||||
learning_rate=3e-4, # You can adjust hyperparameters as needed
|
learning_rate=3e-4,
|
||||||
n_steps=2048,
|
n_steps=2048,
|
||||||
batch_size=64,
|
batch_size=64,
|
||||||
n_epochs=10,
|
n_epochs=10,
|
||||||
gamma=0.99,
|
gamma=0.99,
|
||||||
gae_lambda=0.95,
|
gae_lambda=0.95,
|
||||||
clip_range=0.2,
|
clip_range=0.2,
|
||||||
ent_coef=0.01
|
ent_coef=0.01,
|
||||||
)
|
)
|
||||||
|
model.learn(total_timesteps=100_000)
|
||||||
|
|
||||||
# Train the model
|
|
||||||
model.learn(total_timesteps=100000) # Adjust total_timesteps as needed
|
|
||||||
|
|
||||||
# Test the trained model
|
|
||||||
obs, info = env.reset()
|
obs, info = env.reset()
|
||||||
for _ in range(1000):
|
for _ in range(1000):
|
||||||
action, _states = model.predict(obs, deterministic=True)
|
action, _states = model.predict(obs, deterministic=True)
|
||||||
obs, reward, terminated, truncated, info = env.step(action)
|
obs, reward, terminated, truncated, info = env.step(action)
|
||||||
|
|
||||||
if terminated or truncated:
|
if terminated or truncated:
|
||||||
obs, info = env.reset()
|
obs, info = env.reset()
|
||||||
|
|
||||||
# Close the environment
|
|
||||||
env.close()
|
env.close()
|
||||||
```
|
```
|
||||||
|
|
||||||
But theres a problem: RL algorithms require a huge amount of training steps to get passable policies, and Nucleares is a very slow simulation and can not be trivially parallelized. That's why NuCon also provides a
|
### NuconGoalEnv + HER Usage
|
||||||
|
|
||||||
## Simulator (Work in Progress)
|
HER works by relabelling past trajectories with the goal that was *actually achieved*, turning every episode into useful training signal even when the agent never reaches the intended target. This makes it much more sample-efficient than standard RL for goal-reaching tasks. This matters a lot given how slow the real game is.
|
||||||
|
|
||||||
|
```python
|
||||||
|
from nucon.rl import NuconGoalEnv, Parameterized_Objectives, Parameterized_Terminators
|
||||||
|
from stable_baselines3 import SAC
|
||||||
|
from stable_baselines3.her.her_replay_buffer import HerReplayBuffer
|
||||||
|
|
||||||
|
|
||||||
|
env = NuconGoalEnv(
|
||||||
|
goal_params=['GENERATOR_0_KW', 'GENERATOR_1_KW', 'GENERATOR_2_KW'],
|
||||||
|
goal_range={
|
||||||
|
'GENERATOR_0_KW': (0.0, 1200.0),
|
||||||
|
'GENERATOR_1_KW': (0.0, 1200.0),
|
||||||
|
'GENERATOR_2_KW': (0.0, 1200.0),
|
||||||
|
},
|
||||||
|
tolerance=0.05, # sparse: within 5% of range counts as success (recommended with HER)
|
||||||
|
seconds_per_step=5,
|
||||||
|
simulator=simulator, # use a pre-trained simulator for fast pre-training
|
||||||
|
# Keep policy within the simulator's known data distribution.
|
||||||
|
# SIM_UNCERTAINTY (kNN-GP posterior std) is injected into obs when a simulator is active.
|
||||||
|
# Tune start/scale/threshold to taste.
|
||||||
|
additional_objectives=[Parameterized_Objectives['uncertainty_penalty'](start=0.3, scale=1.0)],
|
||||||
|
terminators=[Parameterized_Terminators['uncertainty_abort'](threshold=0.7)],
|
||||||
|
)
|
||||||
|
# Or use a preset: env = gym.make('Nucon-goal_power-v0', simulator=simulator)
|
||||||
|
|
||||||
|
model = SAC(
|
||||||
|
'MultiInputPolicy',
|
||||||
|
env,
|
||||||
|
replay_buffer_class=HerReplayBuffer,
|
||||||
|
replay_buffer_kwargs={'n_sampled_goal': 4, 'goal_selection_strategy': 'future'},
|
||||||
|
verbose=1,
|
||||||
|
learning_rate=1e-3,
|
||||||
|
batch_size=256,
|
||||||
|
tau=0.005,
|
||||||
|
gamma=0.98,
|
||||||
|
train_freq=1,
|
||||||
|
gradient_steps=1,
|
||||||
|
)
|
||||||
|
model.learn(total_timesteps=500_000)
|
||||||
|
```
|
||||||
|
|
||||||
|
At inference time, inject any target by constructing the observation manually:
|
||||||
|
```python
|
||||||
|
import numpy as np
|
||||||
|
obs, _ = env.reset()
|
||||||
|
# Override the desired goal (values are normalised to [0,1] within goal_range)
|
||||||
|
obs['desired_goal'] = np.array([0.8, 0.8, 0.8], dtype=np.float32) # ~960 kW per generator
|
||||||
|
action, _ = model.predict(obs, deterministic=True)
|
||||||
|
```
|
||||||
|
|
||||||
|
Predefined goal environments:
|
||||||
|
- `Nucon-goal_power-v0`: target total generator output (3 × 0–1200 kW)
|
||||||
|
- `Nucon-goal_temp-v0`: target core temperature (280–380 °C)
|
||||||
|
|
||||||
|
RL algorithms require a huge number of training steps, and Nucleares is slow and cannot be trivially parallelised. That's why NuCon provides a built-in simulator.
|
||||||
|
|
||||||
|
## Simulator
|
||||||
|
|
||||||
NuCon provides a built-in simulator to address the challenge of slow training times in the actual Nucleares game. This simulator allows for rapid prototyping and testing of control policies without the need for the full game environment. Key features include:
|
NuCon provides a built-in simulator to address the challenge of slow training times in the actual Nucleares game. This simulator allows for rapid prototyping and testing of control policies without the need for the full game environment. Key features include:
|
||||||
|
|
||||||
@@ -228,10 +281,7 @@ simulator.load_model('path/to/model.pth')
|
|||||||
# Set initial state (optional)
|
# Set initial state (optional)
|
||||||
simulator.set_state(OperatingState.NOMINAL)
|
simulator.set_state(OperatingState.NOMINAL)
|
||||||
|
|
||||||
# Run the simulator, will start the web server
|
# The web server starts automatically in __init__; access via nucon using the simulator's port
|
||||||
simulator.run()
|
|
||||||
|
|
||||||
# Access via nucon by using the simulator's port
|
|
||||||
nucon = Nucon(port=simulator.port)
|
nucon = Nucon(port=simulator.port)
|
||||||
|
|
||||||
# Or use the simulator with NuconEnv
|
# Or use the simulator with NuconEnv
|
||||||
@@ -242,16 +292,16 @@ env = NuconEnv(simulator=simulator) # When given a similator, instead of waiting
|
|||||||
# ...
|
# ...
|
||||||
```
|
```
|
||||||
|
|
||||||
But theres yet another problem: We do not know the exact simulation dynamics of the game and can therefore not implement an accurate simulator. Thats why NuCon also provides
|
The simulator needs an accurate dynamics model of the game. NuCon provides tools to learn one from real gameplay data.
|
||||||
|
|
||||||
## Model Learning (Work in Progress)
|
## Model Learning
|
||||||
|
|
||||||
To address the challenge of unknown game dynamics, NuCon provides tools for collecting data, creating datasets, and training models to learn the reactor dynamics. Key features include:
|
To address the challenge of unknown game dynamics, NuCon provides tools for collecting data, creating datasets, and training models to learn the reactor dynamics. Key features include:
|
||||||
|
|
||||||
- **Data Collection**: Gathers state transitions from human play or automated agents. `time_delta` is specified in game-time seconds; wall-clock sleep is automatically adjusted for `GAME_SIM_SPEED` so collected deltas are uniform regardless of simulation speed.
|
- **Data Collection**: Gathers state transitions from human play or automated agents. `time_delta` is specified in game-time seconds; wall-clock sleep is automatically adjusted for `GAME_SIM_SPEED` so collected deltas are uniform regardless of simulation speed.
|
||||||
- **Automatic param filtering**: Junk params (GAME_VERSION, TIME, ALARMS_ACTIVE, …) and params from uninstalled subsystems (returns `None`) are automatically excluded from model inputs/outputs.
|
- **Automatic param filtering**: Junk params (GAME_VERSION, TIME, ALARMS_ACTIVE, …) and params from uninstalled subsystems (returns `None`) are automatically excluded from model inputs/outputs.
|
||||||
- **Two model backends**: Neural network (NN) or k-Nearest Neighbours with GP interpolation (kNN).
|
- **Two model backends**: Neural network (NN) or a local Gaussian Process approximated via k-Nearest Neighbours (kNN-GP).
|
||||||
- **Uncertainty estimation**: The kNN backend returns a GP posterior standard deviation alongside each prediction — 0 means the query lies on known data, ~1 means it is out of distribution.
|
- **Uncertainty estimation**: The kNN-GP backend returns a GP posterior standard deviation alongside each prediction; 0 means the query lies on known data, ~1 means it is out of distribution.
|
||||||
- **Dataset management**: Tools for saving, loading, merging, and pruning datasets.
|
- **Dataset management**: Tools for saving, loading, merging, and pruning datasets.
|
||||||
|
|
||||||
### Additional Dependencies
|
### Additional Dependencies
|
||||||
@@ -260,6 +310,10 @@ To address the challenge of unknown game dynamics, NuCon provides tools for coll
|
|||||||
pip install -e '.[model]'
|
pip install -e '.[model]'
|
||||||
```
|
```
|
||||||
|
|
||||||
|
### Model selection
|
||||||
|
|
||||||
|
**kNN-GP** (the `ReactorKNNModel` backend) is a local Gaussian Process: it finds the `k` nearest neighbours in the training set, fits an RBF kernel on them, and returns a prediction plus a GP posterior std as uncertainty. It works well from a few hundred samples and requires no training. **NN** needs input normalisation and several thousand samples to generalise; use it once you have a large dataset. For initial experiments, start with kNN-GP (`k=10`).
|
||||||
|
|
||||||
### Usage
|
### Usage
|
||||||
|
|
||||||
```python
|
```python
|
||||||
@@ -277,19 +331,19 @@ learner.save_dataset('reactor_dataset.pkl')
|
|||||||
learner.merge_datasets('other_session.pkl')
|
learner.merge_datasets('other_session.pkl')
|
||||||
|
|
||||||
# --- Neural network backend ---
|
# --- Neural network backend ---
|
||||||
nn_learner = NuconModelLearner(model_type='nn', dataset_path='reactor_dataset.pkl')
|
nn_learner = NuconModelLearner(dataset_path='reactor_dataset.pkl')
|
||||||
nn_learner.train_model(batch_size=32, num_epochs=50)
|
nn_learner.train_model(batch_size=32, num_epochs=50) # creates NN model on first call
|
||||||
# Drop samples the NN already predicts well (keep hard cases for further training)
|
# Drop samples the NN already predicts well (keep hard cases for further training)
|
||||||
nn_learner.drop_well_fitted(error_threshold=1.0)
|
nn_learner.drop_well_fitted(error_threshold=1.0)
|
||||||
nn_learner.save_model('reactor_nn.pth')
|
nn_learner.save_model('reactor_nn.pth')
|
||||||
|
|
||||||
# --- kNN + GP backend ---
|
# --- kNN-GP backend ---
|
||||||
knn_learner = NuconModelLearner(model_type='knn', knn_k=10, dataset_path='reactor_dataset.pkl')
|
knn_learner = NuconModelLearner(dataset_path='reactor_dataset.pkl')
|
||||||
# Drop near-duplicate samples before fitting (keeps diverse coverage).
|
# Drop near-duplicate samples before fitting (keeps diverse coverage).
|
||||||
# A sample is dropped only if BOTH its input state AND output transition
|
# A sample is dropped only if BOTH its input state AND output transition
|
||||||
# are within the given distances of an already-kept sample.
|
# are within the given distances of an already-kept sample.
|
||||||
knn_learner.drop_redundant(min_state_distance=0.1, min_output_distance=0.05)
|
knn_learner.drop_redundant(min_state_distance=0.1, min_output_distance=0.05)
|
||||||
knn_learner.fit_knn()
|
knn_learner.fit_knn(k=10) # creates kNN-GP model on first call
|
||||||
|
|
||||||
# Point prediction
|
# Point prediction
|
||||||
state = knn_learner._get_state()
|
state = knn_learner._get_state()
|
||||||
@@ -306,6 +360,78 @@ knn_learner.save_model('reactor_knn.pkl')
|
|||||||
|
|
||||||
The trained models can be integrated into the NuconSimulator to provide accurate dynamics based on real game data.
|
The trained models can be integrated into the NuconSimulator to provide accurate dynamics based on real game data.
|
||||||
|
|
||||||
|
## Full Training Loop
|
||||||
|
|
||||||
|
The recommended end-to-end workflow for training an RL operator is an iterative cycle of real-game data collection, model fitting, and simulated training. The real game is slow and cannot be parallelised, so the bulk of RL training happens in the simulator. The game is used only as an oracle for data and evaluation.
|
||||||
|
|
||||||
|
```
|
||||||
|
┌─────────────────────────────────────────────────────────────┐
|
||||||
|
│ 1. Human dataset collection │
|
||||||
|
│ Play the game: start up the reactor, operate it across │
|
||||||
|
│ a range of states. NuCon records state transitions. │
|
||||||
|
└───────────────────────┬─────────────────────────────────────┘
|
||||||
|
│
|
||||||
|
▼
|
||||||
|
┌─────────────────────────────────────────────────────────────┐
|
||||||
|
│ 2. Initial model fitting │
|
||||||
|
│ Fit NN or kNN dynamics model to the collected dataset. │
|
||||||
|
│ kNN is instant; NN needs gradient steps but generalises │
|
||||||
|
│ better with more data. │
|
||||||
|
└───────────────────────┬─────────────────────────────────────┘
|
||||||
|
│
|
||||||
|
┌─────────▼──────────┐
|
||||||
|
│ 3. Train RL │◄───────────────────────┐
|
||||||
|
│ in simulator │ │
|
||||||
|
│ (fast, many │ │
|
||||||
|
│ trajectories) │ │
|
||||||
|
└─────────┬──────────┘ │
|
||||||
|
│ │
|
||||||
|
▼ │
|
||||||
|
┌─────────────────────┐ │
|
||||||
|
│ 4. Eval in game │ │
|
||||||
|
│ + collect new data │ │
|
||||||
|
│ (merge & prune │ │
|
||||||
|
│ dataset) │ │
|
||||||
|
└─────────┬───────────┘ │
|
||||||
|
│ │
|
||||||
|
▼ │
|
||||||
|
┌─────────────────────┐ model improved? │
|
||||||
|
│ 5. Refit model ├──────── yes ──────────┘
|
||||||
|
│ on expanded data │
|
||||||
|
└─────────────────────┘
|
||||||
|
```
|
||||||
|
|
||||||
|
**Step 1 — Human dataset collection**: Run `scripts/collect_dataset.py` during your play session (see [Scripts](#scripts)). Cover a wide range of states: startup from cold, ramping power, individual rod bank adjustments. Diversity in the dataset directly determines simulator accuracy. See [Model Learning](#model-learning) for collection details.
|
||||||
|
|
||||||
|
**Step 2 — Initial model fitting**: Fit a kNN-GP model (instant) or NN (better extrapolation with larger datasets) using `fit_knn()` or `train_model()`. Prune near-duplicate samples with `drop_redundant()` before fitting. See [Model Learning](#model-learning).
|
||||||
|
|
||||||
|
**Step 3 — Train RL in simulator**: Load the fitted model into `NuconSimulator`, then train a `NuconGoalEnv` policy with SAC + HER. The simulator runs far faster than the real game, allowing many trajectories in reasonable time. Pass `Parameterized_Objectives['uncertainty_penalty']` and `Parameterized_Terminators['uncertainty_abort']` as additional objectives/terminators to discourage the policy from wandering into regions the model hasn't seen; `SIM_UNCERTAINTY` is automatically injected into the obs dict when a simulator is active. See [NuconGoalEnv + HER Usage](#nucongoalenv--her-usage) and `scripts/train_sac.py` for a complete example.
|
||||||
|
|
||||||
|
**Step 4 — Eval in game + collect new data**: Run the trained policy against the real game. This validates simulator accuracy and simultaneously collects new data from states the policy visits, which may be regions the original dataset missed. Run a second `NuconModelLearner` in a background thread to collect concurrently.
|
||||||
|
|
||||||
|
**Step 5 — Refit model on expanded data**: Merge new data into the original dataset with `merge_datasets()`, prune with `drop_redundant()`, and refit. Then return to Step 3 with the improved model. Each iteration the simulator gets more accurate and the policy improves.
|
||||||
|
|
||||||
|
Stop when the policy performs well in the real game and kNN-GP uncertainty stays low throughout an episode, indicating the policy stays within the known data distribution.
|
||||||
|
|
||||||
|
## Scripts
|
||||||
|
|
||||||
|
Ready-to-run scripts in the `scripts/` directory covering the most common workflows.
|
||||||
|
|
||||||
|
**`scripts/collect_dataset.py`** — collect a dynamics dataset while playing the game:
|
||||||
|
```bash
|
||||||
|
python scripts/collect_dataset.py --steps 1000 --delta 10 --out reactor_dataset.pkl
|
||||||
|
# Ctrl-C to stop early; data is saved on exit
|
||||||
|
# Merge a previous session: --merge previous.pkl
|
||||||
|
```
|
||||||
|
|
||||||
|
**`scripts/train_sac.py`** — train a SAC + HER goal-conditioned policy on the kNN-GP simulator:
|
||||||
|
```bash
|
||||||
|
python scripts/train_sac.py
|
||||||
|
# Expects /tmp/reactor_knn.pkl and /tmp/nucon_dataset.pkl
|
||||||
|
# Saves trained policy to /tmp/sac_nucon_knn.zip
|
||||||
|
```
|
||||||
|
This script is the most elaborate end-to-end example: it loads a pre-fitted kNN-GP model, seeds episode resets from dataset states, uses delta actions and an uncertainty penalty, and configures SAC + HER for fast sim training.
|
||||||
|
|
||||||
## Testing
|
## Testing
|
||||||
|
|
||||||
NuCon includes a test suite to verify its functionality and compatibility with the Nucleares game.
|
NuCon includes a test suite to verify its functionality and compatibility with the Nucleares game.
|
||||||
|
|||||||
@@ -0,0 +1,27 @@
|
|||||||
|
<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 512 512" width="512" height="512">
|
||||||
|
<!-- Background -->
|
||||||
|
<rect width="512" height="512" fill="#0d0d0d" rx="72"/>
|
||||||
|
|
||||||
|
<!-- Outer pressure vessel ring -->
|
||||||
|
<circle cx="256" cy="256" r="185" fill="none" stroke="#3ab5f0" stroke-width="10"/>
|
||||||
|
|
||||||
|
<!-- Inner core ring -->
|
||||||
|
<circle cx="256" cy="256" r="65" fill="none" stroke="#3ab5f0" stroke-width="10"/>
|
||||||
|
|
||||||
|
<!-- Central nucleus -->
|
||||||
|
<circle cx="256" cy="256" r="18" fill="#3ab5f0"/>
|
||||||
|
|
||||||
|
<!-- 6 control rods: gap inside core, line from core to vessel -->
|
||||||
|
<!-- θ=0° (top) -->
|
||||||
|
<line x1="256" y1="191" x2="256" y2="71" stroke="#3ab5f0" stroke-width="9" stroke-linecap="round"/>
|
||||||
|
<!-- θ=60° -->
|
||||||
|
<line x1="312" y1="223" x2="416" y2="163" stroke="#3ab5f0" stroke-width="9" stroke-linecap="round"/>
|
||||||
|
<!-- θ=120° -->
|
||||||
|
<line x1="312" y1="289" x2="416" y2="349" stroke="#3ab5f0" stroke-width="9" stroke-linecap="round"/>
|
||||||
|
<!-- θ=180° (bottom) -->
|
||||||
|
<line x1="256" y1="321" x2="256" y2="441" stroke="#3ab5f0" stroke-width="9" stroke-linecap="round"/>
|
||||||
|
<!-- θ=240° -->
|
||||||
|
<line x1="200" y1="289" x2="96" y2="349" stroke="#3ab5f0" stroke-width="9" stroke-linecap="round"/>
|
||||||
|
<!-- θ=300° -->
|
||||||
|
<line x1="200" y1="223" x2="96" y2="163" stroke="#3ab5f0" stroke-width="9" stroke-linecap="round"/>
|
||||||
|
</svg>
|
||||||
|
After Width: | Height: | Size: 1.3 KiB |
+231
-74
@@ -18,13 +18,15 @@ Actors = {
|
|||||||
# --- NN-based dynamics model ---
|
# --- NN-based dynamics model ---
|
||||||
|
|
||||||
class ReactorDynamicsNet(nn.Module):
|
class ReactorDynamicsNet(nn.Module):
|
||||||
def __init__(self, input_dim, output_dim):
|
def __init__(self, input_dim, output_dim, dropout=0.3):
|
||||||
super(ReactorDynamicsNet, self).__init__()
|
super(ReactorDynamicsNet, self).__init__()
|
||||||
self.network = nn.Sequential(
|
self.network = nn.Sequential(
|
||||||
nn.Linear(input_dim + 1, 128), # +1 for time_delta
|
nn.Linear(input_dim + 1, 128), # +1 for time_delta
|
||||||
nn.ReLU(),
|
nn.ReLU(),
|
||||||
|
nn.Dropout(dropout),
|
||||||
nn.Linear(128, 128),
|
nn.Linear(128, 128),
|
||||||
nn.ReLU(),
|
nn.ReLU(),
|
||||||
|
nn.Dropout(dropout),
|
||||||
nn.Linear(128, output_dim)
|
nn.Linear(128, output_dim)
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -33,23 +35,81 @@ class ReactorDynamicsNet(nn.Module):
|
|||||||
return self.network(x)
|
return self.network(x)
|
||||||
|
|
||||||
class ReactorDynamicsModel(nn.Module):
|
class ReactorDynamicsModel(nn.Module):
|
||||||
|
"""
|
||||||
|
NN dynamics model predicting per-second rates of change (like ReactorKNNModel).
|
||||||
|
|
||||||
|
Inputs are z-score normalised; outputs are normalised rates.
|
||||||
|
forward() returns absolute next-state dict: cur + predicted_rate * time_delta.
|
||||||
|
forward_with_uncertainty() returns (next_state, 0.0) — no uncertainty estimate.
|
||||||
|
"""
|
||||||
def __init__(self, input_params: List[str], output_params: List[str]):
|
def __init__(self, input_params: List[str], output_params: List[str]):
|
||||||
super(ReactorDynamicsModel, self).__init__()
|
super(ReactorDynamicsModel, self).__init__()
|
||||||
self.input_params = input_params
|
self.input_params = input_params
|
||||||
self.output_params = output_params
|
self.output_params = output_params
|
||||||
self.net = ReactorDynamicsNet(len(input_params), len(output_params))
|
self.net = ReactorDynamicsNet(len(input_params), len(output_params))
|
||||||
|
# Normalisation stats set by fit()
|
||||||
|
self.register_buffer('_in_mean', torch.zeros(len(input_params)))
|
||||||
|
self.register_buffer('_in_std', torch.ones(len(input_params)))
|
||||||
|
self.register_buffer('_rate_mean', torch.zeros(len(output_params)))
|
||||||
|
self.register_buffer('_rate_std', torch.ones(len(output_params)))
|
||||||
|
|
||||||
def _state_dict_to_tensor(self, state_dict):
|
def fit_normalisation(self, dataset):
|
||||||
return torch.tensor([state_dict[p] for p in self.input_params], dtype=torch.float32)
|
"""Compute and store normalisation stats from a dataset."""
|
||||||
|
in_vecs, rate_vecs = [], []
|
||||||
|
for state, _action, next_state, dt in dataset:
|
||||||
|
if dt <= 0:
|
||||||
|
continue
|
||||||
|
in_vecs.append([state.get(p, 0.0) for p in self.input_params])
|
||||||
|
rate_vecs.append([(next_state.get(p, 0.0) - state.get(p, 0.0)) / dt
|
||||||
|
for p in self.output_params])
|
||||||
|
ins = np.array(in_vecs, dtype=np.float32)
|
||||||
|
rates = np.array(rate_vecs, dtype=np.float32)
|
||||||
|
in_std = ins.std(0)
|
||||||
|
r_std = rates.std(0)
|
||||||
|
self._in_mean.copy_(torch.from_numpy(ins.mean(0)))
|
||||||
|
self._in_std.copy_(torch.from_numpy(np.where(in_std < 1e-6, 1.0, in_std)))
|
||||||
|
self._rate_mean.copy_(torch.from_numpy(rates.mean(0)))
|
||||||
|
self._rate_std.copy_(torch.from_numpy(np.where(r_std < 1e-6, 1.0, r_std)))
|
||||||
|
|
||||||
def _tensor_to_state_dict(self, tensor):
|
def _normalise_input(self, t: torch.Tensor) -> torch.Tensor:
|
||||||
return {p: tensor[i].item() for i, p in enumerate(self.output_params)}
|
return (t - self._in_mean) / self._in_std
|
||||||
|
|
||||||
|
def _denormalise_rate(self, t: torch.Tensor) -> torch.Tensor:
|
||||||
|
return t * self._rate_std + self._rate_mean
|
||||||
|
|
||||||
def forward(self, state_dict, time_delta):
|
def forward(self, state_dict, time_delta):
|
||||||
state_tensor = self._state_dict_to_tensor(state_dict).unsqueeze(0)
|
return self.forward_with_uncertainty(state_dict, time_delta)[0]
|
||||||
time_delta_tensor = torch.tensor([time_delta], dtype=torch.float32).unsqueeze(0)
|
|
||||||
predicted_tensor = self.net(state_tensor, time_delta_tensor)
|
def forward_with_uncertainty(self, state_dict, time_delta, mc_samples=3):
|
||||||
return self._tensor_to_state_dict(predicted_tensor.squeeze(0))
|
"""MC-Dropout uncertainty: run mc_samples stochastic forward passes.
|
||||||
|
|
||||||
|
Uncertainty is the mean normalised std across output dims, clipped to [0, 1].
|
||||||
|
0 = very confident (low variance), ~1 = high variance / OOD.
|
||||||
|
"""
|
||||||
|
s = torch.tensor([state_dict.get(p, 0.0) for p in self.input_params],
|
||||||
|
dtype=torch.float32).unsqueeze(0)
|
||||||
|
s_norm = self._normalise_input(s)
|
||||||
|
dt_t = torch.tensor([[time_delta]], dtype=torch.float32)
|
||||||
|
|
||||||
|
# Keep dropout active for uncertainty sampling
|
||||||
|
self.net.train()
|
||||||
|
with torch.no_grad():
|
||||||
|
samples = torch.stack([self.net(s_norm, dt_t).squeeze(0)
|
||||||
|
for _ in range(mc_samples)]) # (mc_samples, out_dim)
|
||||||
|
self.net.eval()
|
||||||
|
|
||||||
|
rate_norm_mean = samples.mean(0)
|
||||||
|
rate_norm_std = samples.std(0)
|
||||||
|
|
||||||
|
rate = self._denormalise_rate(rate_norm_mean)
|
||||||
|
cur = torch.tensor([state_dict.get(p, 0.0) for p in self.output_params],
|
||||||
|
dtype=torch.float32)
|
||||||
|
predicted = cur + rate * time_delta
|
||||||
|
pred_dict = {p: float(predicted[i]) for i, p in enumerate(self.output_params)}
|
||||||
|
|
||||||
|
# Uncertainty: mean coefficient of variation in normalised space, clipped to [0,1]
|
||||||
|
uncertainty = float(rate_norm_std.mean().clamp(0.0, 1.0))
|
||||||
|
return pred_dict, uncertainty
|
||||||
|
|
||||||
# --- kNN-based dynamics model ---
|
# --- kNN-based dynamics model ---
|
||||||
|
|
||||||
@@ -95,14 +155,17 @@ class ReactorKNNModel:
|
|||||||
self._raw_states = np.array(raw)
|
self._raw_states = np.array(raw)
|
||||||
self._rates = np.array(rates)
|
self._rates = np.array(rates)
|
||||||
self._mean = self._raw_states.mean(axis=0)
|
self._mean = self._raw_states.mean(axis=0)
|
||||||
self._std = self._raw_states.std(axis=0) + 1e-8
|
raw_std = self._raw_states.std(axis=0)
|
||||||
|
# Dimensions with zero variance in the training data carry no distance information.
|
||||||
|
# Use inf so they contribute 0 to normalised L2 (i.e., are ignored in kNN lookup).
|
||||||
|
self._std = np.where(raw_std < 1e-6, np.inf, raw_std)
|
||||||
self._states = (self._raw_states - self._mean) / self._std
|
self._states = (self._raw_states - self._mean) / self._std
|
||||||
|
|
||||||
def _lookup(self, state_dict: Dict):
|
def _lookup(self, s: np.ndarray):
|
||||||
"""Return (s_norm, idx, k) for the k nearest neighbours."""
|
"""Return (s_norm, idx, k) for the k nearest neighbours. s is a raw (d_in,) array."""
|
||||||
s = np.array([state_dict[p] for p in self.input_params], dtype=np.float32)
|
|
||||||
s_norm = (s - self._mean) / self._std
|
s_norm = (s - self._mean) / self._std
|
||||||
dists = np.linalg.norm(self._states - s_norm, axis=1)
|
diff = self._states - s_norm # (n, d_in) broadcast
|
||||||
|
dists = np.einsum('ij,ij->i', diff, diff) # squared L2, faster than linalg.norm
|
||||||
k = min(self.k, len(dists))
|
k = min(self.k, len(dists))
|
||||||
idx = np.argpartition(dists, k - 1)[:k]
|
idx = np.argpartition(dists, k - 1)[:k]
|
||||||
return s_norm, idx, k
|
return s_norm, idx, k
|
||||||
@@ -122,22 +185,22 @@ class ReactorKNNModel:
|
|||||||
if self._states is None:
|
if self._states is None:
|
||||||
raise ValueError("Model not fitted. Call fit(dataset) first.")
|
raise ValueError("Model not fitted. Call fit(dataset) first.")
|
||||||
|
|
||||||
s_norm, idx, k = self._lookup(state_dict)
|
s = np.array([state_dict[p] for p in self.input_params], dtype=np.float32)
|
||||||
|
s_norm, idx, k = self._lookup(s)
|
||||||
X = self._states[idx] # (k, d_in)
|
X = self._states[idx] # (k, d_in)
|
||||||
Y = self._rates[idx] # (k, d_out)
|
Y = self._rates[idx] # (k, d_out)
|
||||||
|
|
||||||
# RBF kernel (vectorised): k(a,b) = exp(-0.5 ||a-b||^2)
|
# RBF kernel: k(a,b) = exp(-0.5 ||a-b||^2)
|
||||||
def rbf_matrix(A, B):
|
def rbf(A, B):
|
||||||
diff = A[:, None, :] - B[None, :, :] # (|A|, |B|, d)
|
diff = A[:, None, :] - B[None, :, :]
|
||||||
return np.exp(-0.5 * (diff ** 2).sum(axis=-1)) # (|A|, |B|)
|
return np.exp(-0.5 * np.einsum('ijk,ijk->ij', diff, diff))
|
||||||
|
|
||||||
K = rbf_matrix(X, X) + 1e-4 * np.eye(k) # (k, k)
|
K = rbf(X, X) + 1e-4 * np.eye(k)
|
||||||
k_star = rbf_matrix(s_norm[None, :], X)[0] # (k,)
|
k_star = rbf(s_norm[None, :], X)[0]
|
||||||
|
|
||||||
K_inv = np.linalg.inv(K)
|
K_inv = np.linalg.inv(K)
|
||||||
mean_rates = k_star @ K_inv @ Y # (d_out,)
|
mean_rates = k_star @ K_inv @ Y
|
||||||
|
|
||||||
# Posterior variance (scalar, shared across all output dims)
|
|
||||||
var = max(0.0, 1.0 - float(k_star @ K_inv @ k_star))
|
var = max(0.0, 1.0 - float(k_star @ K_inv @ k_star))
|
||||||
std = float(np.sqrt(var))
|
std = float(np.sqrt(var))
|
||||||
|
|
||||||
@@ -147,26 +210,67 @@ class ReactorKNNModel:
|
|||||||
pred_dict = {p: float(predicted[i]) for i, p in enumerate(self.output_params)}
|
pred_dict = {p: float(predicted[i]) for i, p in enumerate(self.output_params)}
|
||||||
return pred_dict, std
|
return pred_dict, std
|
||||||
|
|
||||||
|
# --- Mixture model ---
|
||||||
|
|
||||||
|
class MixtureModel:
|
||||||
|
"""Combines two dynamics models, selecting based on kNN uncertainty.
|
||||||
|
|
||||||
|
Uses knn_model when its uncertainty is below threshold (it's confident /
|
||||||
|
near training data). Falls back to nn_model when kNN is OOD.
|
||||||
|
|
||||||
|
Both models must implement forward_with_uncertainty(state_dict, time_delta).
|
||||||
|
input_params / output_params are taken from knn_model.
|
||||||
|
"""
|
||||||
|
def __init__(self, knn_model, nn_model):
|
||||||
|
self.knn_model = knn_model
|
||||||
|
self.nn_model = nn_model
|
||||||
|
self.input_params = knn_model.input_params
|
||||||
|
self.output_params = knn_model.output_params
|
||||||
|
|
||||||
|
def forward(self, state_dict, time_delta):
|
||||||
|
return self.forward_with_uncertainty(state_dict, time_delta)[0]
|
||||||
|
|
||||||
|
def forward_with_uncertainty(self, state_dict, time_delta):
|
||||||
|
knn_pred, knn_u = self.knn_model.forward_with_uncertainty(state_dict, time_delta)
|
||||||
|
nn_pred, nn_u = self.nn_model.forward_with_uncertainty(state_dict, time_delta)
|
||||||
|
w_knn = 1.0 - knn_u # high when kNN is confident
|
||||||
|
w_nn = knn_u # high when kNN is OOD
|
||||||
|
blended = {p: w_knn * knn_pred[p] + w_nn * nn_pred[p]
|
||||||
|
for p in self.output_params}
|
||||||
|
uncertainty = w_knn * knn_u + w_nn * nn_u # weighted uncertainty
|
||||||
|
return blended, uncertainty
|
||||||
|
|
||||||
|
|
||||||
# --- Learner ---
|
# --- Learner ---
|
||||||
|
|
||||||
class NuconModelLearner:
|
class NuconModelLearner:
|
||||||
def __init__(self, nucon=None, actor='null', dataset_path='nucon_dataset.pkl',
|
def __init__(self, nucon=None, actor='null', dataset_path='nucon_dataset.pkl',
|
||||||
time_delta: Union[float, Tuple[float, float]] = 1.0,
|
time_delta: Union[float, Tuple[float, float]] = 1.0,
|
||||||
model_type: str = 'nn', knn_k: int = 5,
|
|
||||||
include_valve_states: bool = False):
|
include_valve_states: bool = False):
|
||||||
self.nucon = Nucon() if nucon is None else nucon
|
self.nucon = Nucon() if nucon is None else nucon
|
||||||
self.actor = Actors[actor](self.nucon) if actor in Actors else actor
|
self.actor = Actors[actor](self.nucon) if actor in Actors else actor
|
||||||
self.dataset = self.load_dataset(dataset_path) or []
|
self.dataset = self.load_dataset(dataset_path) or []
|
||||||
self.dataset_path = dataset_path
|
self.dataset_path = dataset_path
|
||||||
self.include_valve_states = include_valve_states
|
self.include_valve_states = include_valve_states
|
||||||
|
self.model = None
|
||||||
|
self.optimizer = None
|
||||||
|
|
||||||
# Exclude params with no physics signal
|
# Exclude params with no physics signal
|
||||||
_JUNK_PARAMS = frozenset({'GAME_VERSION', 'TIME', 'TIME_STAMP', 'TIME_DAY',
|
_JUNK_PARAMS = frozenset({'GAME_VERSION', 'TIME', 'TIME_STAMP', 'TIME_DAY',
|
||||||
'ALARMS_ACTIVE', 'FUN_IS_ENABLED', 'GAME_SIM_SPEED'})
|
'ALARMS_ACTIVE', 'FUN_IS_ENABLED', 'GAME_SIM_SPEED'})
|
||||||
candidate_params = {k: p for k, p in self.nucon.get_all_readable().items()
|
candidate_params = {k: p for k, p in self.nucon.get_all_readable().items()
|
||||||
if k not in _JUNK_PARAMS and p.param_type != str}
|
if k not in _JUNK_PARAMS and p.param_type != str}
|
||||||
# Filter out params that return None (subsystem not installed)
|
# Filter out params that return None (subsystem not installed).
|
||||||
|
# Retry until the game is reachable.
|
||||||
|
import requests as _requests
|
||||||
|
while True:
|
||||||
|
try:
|
||||||
test_state = {k: self.nucon.get(k) for k in candidate_params}
|
test_state = {k: self.nucon.get(k) for k in candidate_params}
|
||||||
|
break
|
||||||
|
except (_requests.exceptions.ConnectionError,
|
||||||
|
_requests.exceptions.Timeout):
|
||||||
|
print("Waiting for game to be reachable…")
|
||||||
|
time.sleep(5)
|
||||||
self.readable_params = [k for k in candidate_params if test_state[k] is not None]
|
self.readable_params = [k for k in candidate_params if test_state[k] is not None]
|
||||||
self.non_writable_params = [k for k in self.readable_params
|
self.non_writable_params = [k for k in self.readable_params
|
||||||
if not self.nucon.get_all_readable()[k].is_writable]
|
if not self.nucon.get_all_readable()[k].is_writable]
|
||||||
@@ -179,15 +283,6 @@ class NuconModelLearner:
|
|||||||
self.readable_params = self.readable_params + self.valve_keys
|
self.readable_params = self.readable_params + self.valve_keys
|
||||||
# valve positions are input-only (not predicted as outputs)
|
# valve positions are input-only (not predicted as outputs)
|
||||||
|
|
||||||
if model_type == 'nn':
|
|
||||||
self.model = ReactorDynamicsModel(self.readable_params, self.non_writable_params)
|
|
||||||
self.optimizer = optim.Adam(self.model.parameters())
|
|
||||||
elif model_type == 'knn':
|
|
||||||
self.model = ReactorKNNModel(self.readable_params, self.non_writable_params, k=knn_k)
|
|
||||||
self.optimizer = None
|
|
||||||
else:
|
|
||||||
raise ValueError(f"Unknown model_type '{model_type}'. Use 'nn' or 'knn'.")
|
|
||||||
|
|
||||||
if isinstance(time_delta, (int, float)):
|
if isinstance(time_delta, (int, float)):
|
||||||
self.time_delta = lambda: time_delta
|
self.time_delta = lambda: time_delta
|
||||||
elif isinstance(time_delta, tuple) and len(time_delta) == 2:
|
elif isinstance(time_delta, tuple) and len(time_delta) == 2:
|
||||||
@@ -211,33 +306,65 @@ class NuconModelLearner:
|
|||||||
state[key] = valves.get(name, {}).get('Value', 0.0)
|
state[key] = valves.get(name, {}).get('Value', 0.0)
|
||||||
return state
|
return state
|
||||||
|
|
||||||
def collect_data(self, num_steps):
|
def collect_data(self, num_steps, save_every=10):
|
||||||
"""
|
"""
|
||||||
Collect state-transition tuples from the live game.
|
Collect state-transition tuples from the live game.
|
||||||
|
|
||||||
Sleeps wall_time = target_game_delta / sim_speed so that each stored
|
Sleeps wall_time = target_game_delta / sim_speed so that each stored
|
||||||
game_delta is uniform regardless of the game's simulation speed setting.
|
game_delta is uniform regardless of the game's simulation speed setting.
|
||||||
|
|
||||||
|
Saves the dataset every ``save_every`` steps so a crash doesn't lose
|
||||||
|
everything. On a connection error the step is skipped and collection
|
||||||
|
resumes once the game is reachable again (retries every 5 s).
|
||||||
"""
|
"""
|
||||||
state = self._get_state()
|
import requests as _requests
|
||||||
for _ in range(num_steps):
|
|
||||||
|
def get_state_with_retry():
|
||||||
|
while True:
|
||||||
|
try:
|
||||||
|
return self._get_state()
|
||||||
|
except (_requests.exceptions.ConnectionError,
|
||||||
|
_requests.exceptions.Timeout) as e:
|
||||||
|
print(f"Connection lost ({e}). Retrying in 5 s…")
|
||||||
|
time.sleep(5)
|
||||||
|
|
||||||
|
state = get_state_with_retry()
|
||||||
|
collected = 0
|
||||||
|
for i in range(num_steps):
|
||||||
action = self.actor(state)
|
action = self.actor(state)
|
||||||
for param_id, value in action.items():
|
for param_id, value in action.items():
|
||||||
|
try:
|
||||||
self.nucon.set(param_id, value)
|
self.nucon.set(param_id, value)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
target_game_delta = self.time_delta()
|
target_game_delta = self.time_delta()
|
||||||
|
try:
|
||||||
sim_speed = self.nucon.GAME_SIM_SPEED.value or 1.0
|
sim_speed = self.nucon.GAME_SIM_SPEED.value or 1.0
|
||||||
|
except Exception:
|
||||||
|
sim_speed = 1.0
|
||||||
time.sleep(target_game_delta / sim_speed)
|
time.sleep(target_game_delta / sim_speed)
|
||||||
next_state = self._get_state()
|
|
||||||
|
|
||||||
|
next_state = get_state_with_retry()
|
||||||
self.dataset.append((state, action, next_state, target_game_delta))
|
self.dataset.append((state, action, next_state, target_game_delta))
|
||||||
state = next_state
|
state = next_state
|
||||||
|
collected += 1
|
||||||
|
|
||||||
|
if collected % save_every == 0:
|
||||||
|
self.save_dataset()
|
||||||
|
print(f" {collected}/{num_steps} steps collected, dataset saved.")
|
||||||
|
|
||||||
self.save_dataset()
|
self.save_dataset()
|
||||||
|
print(f"Collection complete. {collected} steps, {len(self.dataset)} total samples.")
|
||||||
|
|
||||||
def train_model(self, batch_size=32, num_epochs=10, test_split=0.2):
|
def train_model(self, batch_size=32, num_epochs=10, test_split=0.2, lr=1e-3):
|
||||||
"""Train the NN model. For kNN, call fit_knn() instead."""
|
"""Train a neural-network dynamics model on the current dataset."""
|
||||||
if not isinstance(self.model, ReactorDynamicsModel):
|
if self.model is None:
|
||||||
raise ValueError("train_model() is for the NN model. Use fit_knn() for kNN.")
|
self.model = ReactorDynamicsModel(self.readable_params, self.non_writable_params)
|
||||||
|
elif not isinstance(self.model, ReactorDynamicsModel):
|
||||||
|
raise ValueError("A kNN model is already loaded. Create a new learner to train an NN.")
|
||||||
|
self.model.fit_normalisation(self.dataset)
|
||||||
|
self.optimizer = optim.Adam(self.model.parameters(), lr=lr, weight_decay=1e-4)
|
||||||
random.shuffle(self.dataset)
|
random.shuffle(self.dataset)
|
||||||
split_idx = int(len(self.dataset) * (1 - test_split))
|
split_idx = int(len(self.dataset) * (1 - test_split))
|
||||||
train_data = self.dataset[:split_idx]
|
train_data = self.dataset[:split_idx]
|
||||||
@@ -247,17 +374,19 @@ class NuconModelLearner:
|
|||||||
test_loss = self._test_epoch(test_data)
|
test_loss = self._test_epoch(test_data)
|
||||||
print(f"Epoch {epoch+1}/{num_epochs}, Train Loss: {train_loss:.4f}, Test Loss: {test_loss:.4f}")
|
print(f"Epoch {epoch+1}/{num_epochs}, Train Loss: {train_loss:.4f}, Test Loss: {test_loss:.4f}")
|
||||||
|
|
||||||
def fit_knn(self):
|
def fit_knn(self, k: int = 5):
|
||||||
"""Fit the kNN/GP model from the current dataset (instantaneous, no gradient steps)."""
|
"""Fit a kNN/GP dynamics model from the current dataset (instantaneous, no gradient steps)."""
|
||||||
if not isinstance(self.model, ReactorKNNModel):
|
if self.model is None:
|
||||||
raise ValueError("fit_knn() is for the kNN model. Use train_model() for NN.")
|
self.model = ReactorKNNModel(self.readable_params, self.non_writable_params, k=k)
|
||||||
|
elif not isinstance(self.model, ReactorKNNModel):
|
||||||
|
raise ValueError("An NN model is already loaded. Create a new learner to fit a kNN.")
|
||||||
self.model.fit(self.dataset)
|
self.model.fit(self.dataset)
|
||||||
print(f"kNN model fitted on {len(self.dataset)} samples.")
|
print(f"kNN model fitted on {len(self.dataset)} samples.")
|
||||||
|
|
||||||
def predict_with_uncertainty(self, state_dict: Dict, time_delta: float):
|
def predict_with_uncertainty(self, state_dict: Dict, time_delta: float):
|
||||||
"""Return (prediction_dict, uncertainty_std). Only available for kNN model."""
|
"""Return (prediction_dict, uncertainty_std). Only available after fit_knn()."""
|
||||||
if not isinstance(self.model, ReactorKNNModel):
|
if not isinstance(self.model, ReactorKNNModel):
|
||||||
raise ValueError("predict_with_uncertainty() requires model_type='knn'.")
|
raise ValueError("predict_with_uncertainty() requires a fitted kNN model (call fit_knn()).")
|
||||||
return self.model.forward_with_uncertainty(state_dict, time_delta)
|
return self.model.forward_with_uncertainty(state_dict, time_delta)
|
||||||
|
|
||||||
def drop_well_fitted(self, error_threshold: float):
|
def drop_well_fitted(self, error_threshold: float):
|
||||||
@@ -266,6 +395,8 @@ class NuconModelLearner:
|
|||||||
Keeps only hard/surprising transitions. Useful for NN training to focus
|
Keeps only hard/surprising transitions. Useful for NN training to focus
|
||||||
capacity on difficult regions of state space.
|
capacity on difficult regions of state space.
|
||||||
"""
|
"""
|
||||||
|
if self.model is None:
|
||||||
|
raise ValueError("No model fitted yet. Call train_model() or fit_knn() first.")
|
||||||
kept = []
|
kept = []
|
||||||
for state, action, next_state, time_delta in self.dataset:
|
for state, action, next_state, time_delta in self.dataset:
|
||||||
pred = self.model.forward(state, time_delta)
|
pred = self.model.forward(state, time_delta)
|
||||||
@@ -326,51 +457,73 @@ class NuconModelLearner:
|
|||||||
print(f"drop_redundant: kept {len(self.dataset)}, dropped {dropped} samples.")
|
print(f"drop_redundant: kept {len(self.dataset)}, dropped {dropped} samples.")
|
||||||
|
|
||||||
def _train_epoch(self, data, batch_size):
|
def _train_epoch(self, data, batch_size):
|
||||||
out_indices = [self.readable_params.index(p) if p in self.readable_params else None
|
self.model.train()
|
||||||
for p in self.non_writable_params]
|
|
||||||
total_loss = 0
|
total_loss = 0
|
||||||
|
n_batches = 0
|
||||||
for i in range(0, len(data), batch_size):
|
for i in range(0, len(data), batch_size):
|
||||||
batch = data[i:i+batch_size]
|
batch = [s for s in data[i:i+batch_size] if s[3] > 0]
|
||||||
|
if not batch:
|
||||||
|
continue
|
||||||
|
states = torch.tensor([[s[0].get(p, 0.0) for p in self.readable_params] for s in batch], dtype=torch.float32)
|
||||||
|
targets = torch.tensor([[(s[2].get(p, 0.0) - s[0].get(p, 0.0)) / s[3] for p in self.non_writable_params] for s in batch], dtype=torch.float32)
|
||||||
|
dts = torch.tensor([[s[3]] for s in batch], dtype=torch.float32)
|
||||||
|
s_norm = self.model._normalise_input(states)
|
||||||
|
rate_norm_pred = self.model.net(s_norm, dts)
|
||||||
|
rate_norm_target = (targets - self.model._rate_mean) / self.model._rate_std
|
||||||
self.optimizer.zero_grad()
|
self.optimizer.zero_grad()
|
||||||
loss = torch.tensor(0.0)
|
loss = torch.nn.functional.mse_loss(rate_norm_pred, rate_norm_target)
|
||||||
for state, _, next_state, time_delta in batch:
|
|
||||||
state_t = self.model._state_dict_to_tensor(state).unsqueeze(0)
|
|
||||||
td_t = torch.tensor([[time_delta]], dtype=torch.float32)
|
|
||||||
pred = self.model.net(state_t, td_t).squeeze(0)
|
|
||||||
target = torch.tensor([next_state[p] for p in self.non_writable_params],
|
|
||||||
dtype=torch.float32)
|
|
||||||
loss = loss + torch.nn.functional.mse_loss(pred, target)
|
|
||||||
loss = loss / len(batch)
|
|
||||||
loss.backward()
|
loss.backward()
|
||||||
self.optimizer.step()
|
self.optimizer.step()
|
||||||
total_loss += loss.item()
|
total_loss += loss.item()
|
||||||
return total_loss / max(1, len(data) // batch_size)
|
n_batches += 1
|
||||||
|
self.model.eval()
|
||||||
|
return total_loss / max(1, n_batches)
|
||||||
|
|
||||||
def _test_epoch(self, data):
|
def _test_epoch(self, data):
|
||||||
total_loss = 0.0
|
total_loss = 0.0
|
||||||
|
n = 0
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
for state, _, next_state, time_delta in data:
|
for state, _, next_state, dt in data:
|
||||||
state_t = self.model._state_dict_to_tensor(state).unsqueeze(0)
|
if dt <= 0:
|
||||||
td_t = torch.tensor([[time_delta]], dtype=torch.float32)
|
continue
|
||||||
pred = self.model.net(state_t, td_t).squeeze(0)
|
s_t = torch.tensor([[state.get(p, 0.0) for p in self.readable_params]], dtype=torch.float32)
|
||||||
target = torch.tensor([next_state[p] for p in self.non_writable_params],
|
s_norm = self.model._normalise_input(s_t)
|
||||||
dtype=torch.float32)
|
dt_t = torch.tensor([[dt]], dtype=torch.float32)
|
||||||
total_loss += torch.nn.functional.mse_loss(pred, target).item()
|
rate_norm_pred = self.model.net(s_norm, dt_t).squeeze(0)
|
||||||
return total_loss / len(data)
|
target = torch.tensor([(next_state.get(p, 0.0) - state.get(p, 0.0)) / dt
|
||||||
|
for p in self.non_writable_params], dtype=torch.float32)
|
||||||
|
rate_norm_target = (target - self.model._rate_mean) / self.model._rate_std
|
||||||
|
total_loss += torch.nn.functional.mse_loss(rate_norm_pred, rate_norm_target).item()
|
||||||
|
n += 1
|
||||||
|
return total_loss / max(1, n)
|
||||||
|
|
||||||
def save_model(self, path):
|
def save_model(self, path):
|
||||||
|
if self.model is None:
|
||||||
|
raise ValueError("No model to save. Call train_model() or fit_knn() first.")
|
||||||
if isinstance(self.model, ReactorDynamicsModel):
|
if isinstance(self.model, ReactorDynamicsModel):
|
||||||
torch.save(self.model.state_dict(), path)
|
torch.save({
|
||||||
|
'state_dict': self.model.state_dict(),
|
||||||
|
'input_params': self.model.input_params,
|
||||||
|
'output_params': self.model.output_params,
|
||||||
|
}, path)
|
||||||
else:
|
else:
|
||||||
with open(path, 'wb') as f:
|
with open(path, 'wb') as f:
|
||||||
pickle.dump(self.model, f)
|
pickle.dump(self.model, f)
|
||||||
|
|
||||||
def load_model(self, path):
|
def load_model(self, path):
|
||||||
if isinstance(self.model, ReactorDynamicsModel):
|
if path.endswith('.pkl'):
|
||||||
self.model.load_state_dict(torch.load(path))
|
|
||||||
else:
|
|
||||||
with open(path, 'rb') as f:
|
with open(path, 'rb') as f:
|
||||||
self.model = pickle.load(f)
|
self.model = pickle.load(f)
|
||||||
|
else:
|
||||||
|
checkpoint = torch.load(path, weights_only=False)
|
||||||
|
if isinstance(checkpoint, dict) and 'state_dict' in checkpoint:
|
||||||
|
m = ReactorDynamicsModel(checkpoint['input_params'], checkpoint['output_params'])
|
||||||
|
m.load_state_dict(checkpoint['state_dict'])
|
||||||
|
self.model = m
|
||||||
|
else:
|
||||||
|
# legacy plain state dict
|
||||||
|
self.model = ReactorDynamicsModel(self.readable_params, self.non_writable_params)
|
||||||
|
self.model.load_state_dict(checkpoint)
|
||||||
|
|
||||||
def save_dataset(self, path=None):
|
def save_dataset(self, path=None):
|
||||||
path = path or self.dataset_path
|
path = path or self.dataset_path
|
||||||
@@ -386,6 +539,10 @@ class NuconModelLearner:
|
|||||||
|
|
||||||
def merge_datasets(self, other_dataset_path):
|
def merge_datasets(self, other_dataset_path):
|
||||||
other_dataset = self.load_dataset(other_dataset_path)
|
other_dataset = self.load_dataset(other_dataset_path)
|
||||||
if other_dataset:
|
if not isinstance(other_dataset, list):
|
||||||
|
raise ValueError(
|
||||||
|
f"'{other_dataset_path}' does not contain a dataset (got {type(other_dataset).__name__}). "
|
||||||
|
f"Pass a dataset .pkl file, not a model file."
|
||||||
|
)
|
||||||
self.dataset.extend(other_dataset)
|
self.dataset.extend(other_dataset)
|
||||||
self.save_dataset()
|
self.save_dataset()
|
||||||
|
|||||||
+458
-79
@@ -1,86 +1,169 @@
|
|||||||
|
import inspect
|
||||||
import gymnasium as gym
|
import gymnasium as gym
|
||||||
from gymnasium import spaces
|
from gymnasium import spaces
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import time
|
import time
|
||||||
from typing import Dict, Any
|
from typing import Dict, Any, Callable, List, Optional
|
||||||
|
from enum import Enum
|
||||||
from nucon import Nucon, BreakerStatus, PumpStatus, PumpDryStatus, PumpOverloadStatus
|
from nucon import Nucon, BreakerStatus, PumpStatus, PumpDryStatus, PumpOverloadStatus
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Reward / objective helpers
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
def _alarm_penalty(obs):
|
||||||
|
"""Penalty proportional to number of active alarms. Only meaningful when running against the real game."""
|
||||||
|
raw = obs.get('ALARMS_ACTIVE', '')
|
||||||
|
if not raw or not raw.strip():
|
||||||
|
return 0.0
|
||||||
|
return -float(len(raw.split(',')))
|
||||||
|
|
||||||
Objectives = {
|
Objectives = {
|
||||||
"null": lambda obs: 0,
|
"null": lambda obs: 0,
|
||||||
"max_power": lambda obs: obs["GENERATOR_0_KW"] + obs["GENERATOR_1_KW"] + obs["GENERATOR_2_KW"],
|
"max_power": lambda obs: obs["GENERATOR_0_KW"] + obs["GENERATOR_1_KW"] + obs["GENERATOR_2_KW"],
|
||||||
"episode_time": lambda obs: obs["EPISODE_TIME"],
|
"episode_time": lambda obs: obs["EPISODE_TIME"],
|
||||||
|
"alarm_penalty": _alarm_penalty,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
def _uncertainty_penalty(start=0.3, scale=1.0, mode='l2'):
|
||||||
|
excess = lambda obs: max(0.0, obs.get('SIM_UNCERTAINTY', 0.0) - start)
|
||||||
|
if mode == 'l2':
|
||||||
|
return lambda obs: -scale * excess(obs) ** 2
|
||||||
|
elif mode == 'linear':
|
||||||
|
return lambda obs: -scale * excess(obs)
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Unknown mode '{mode}'. Use 'l2' or 'linear'.")
|
||||||
|
|
||||||
|
def _uncertainty_abort(threshold=0.7):
|
||||||
|
return lambda obs: 1.0 if obs.get('SIM_UNCERTAINTY', 0.0) >= threshold else 0.0
|
||||||
|
|
||||||
Parameterized_Objectives = {
|
Parameterized_Objectives = {
|
||||||
"target_temperature": lambda goal_temp: lambda obs: -((obs["CORE_TEMP"] - goal_temp) ** 2),
|
"target_temperature": lambda goal_temp: lambda obs: -((obs["CORE_TEMP"] - goal_temp) ** 2),
|
||||||
"target_gap": lambda goal_gap: lambda obs: -((obs["CORE_TEMP"] - obs["CORE_TEMP_MIN"] - goal_gap) ** 2),
|
"target_gap": lambda goal_gap: lambda obs: -((obs["CORE_TEMP"] - obs["CORE_TEMP_MIN"] - goal_gap) ** 2),
|
||||||
"temp_below": lambda max_temp: lambda obs: -(np.clip(obs["CORE_TEMP"] - max_temp, 0, np.inf) ** 2),
|
"temp_below": lambda max_temp: lambda obs: -(np.clip(obs["CORE_TEMP"] - max_temp, 0, np.inf) ** 2),
|
||||||
"temp_above": lambda min_temp: lambda obs: -(np.clip(min_temp - obs["CORE_TEMP"], 0, np.inf) ** 2),
|
"temp_above": lambda min_temp: lambda obs: -(np.clip(min_temp - obs["CORE_TEMP"], 0, np.inf) ** 2),
|
||||||
|
"temp_below_linear": lambda max_temp: lambda obs: -np.clip(obs["CORE_TEMP"] - max_temp, 0, np.inf),
|
||||||
|
"temp_above_linear": lambda min_temp: lambda obs: -np.clip(min_temp - obs["CORE_TEMP"], 0, np.inf),
|
||||||
"constant": lambda constant: lambda obs: constant,
|
"constant": lambda constant: lambda obs: constant,
|
||||||
|
"uncertainty_penalty": _uncertainty_penalty, # (start, scale, mode) -> (obs) -> float
|
||||||
}
|
}
|
||||||
|
|
||||||
|
Parameterized_Terminators = {
|
||||||
|
"uncertainty_abort": _uncertainty_abort, # (threshold,) -> (obs) -> float
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Internal helpers
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
def _build_flat_action_space(nucon, obs_param_set=None, delta_action_scale=None):
|
||||||
|
"""Return (Box, ordered_param_ids, param_ranges).
|
||||||
|
|
||||||
|
If delta_action_scale is set, the action space is [-1, 1]^n and actions are
|
||||||
|
treated as normalised deltas: actual_delta = action * delta_action_scale * (max - min).
|
||||||
|
Otherwise the action space spans [min_val, max_val] per param (absolute values).
|
||||||
|
"""
|
||||||
|
params = []
|
||||||
|
lows, highs, ranges = [], [], []
|
||||||
|
for param_id, param in nucon.get_all_writable().items():
|
||||||
|
if not param.is_readable or param.is_cheat:
|
||||||
|
continue
|
||||||
|
if obs_param_set is not None and param_id not in obs_param_set:
|
||||||
|
continue
|
||||||
|
if param.min_val is None or param.max_val is None:
|
||||||
|
continue # SAC requires finite action bounds
|
||||||
|
sp = _build_param_space(param)
|
||||||
|
if sp is None:
|
||||||
|
continue
|
||||||
|
params.append(param_id)
|
||||||
|
lows.append(sp.low[0])
|
||||||
|
highs.append(sp.high[0])
|
||||||
|
ranges.append(sp.high[0] - sp.low[0])
|
||||||
|
if delta_action_scale is not None:
|
||||||
|
n = len(params)
|
||||||
|
box = spaces.Box(low=-np.ones(n, dtype=np.float32),
|
||||||
|
high=np.ones(n, dtype=np.float32), dtype=np.float32)
|
||||||
|
else:
|
||||||
|
box = spaces.Box(low=np.array(lows, dtype=np.float32),
|
||||||
|
high=np.array(highs, dtype=np.float32), dtype=np.float32)
|
||||||
|
return box, params, np.array(lows, dtype=np.float32), np.array(ranges, dtype=np.float32)
|
||||||
|
|
||||||
|
|
||||||
|
def _unflatten_action(flat_action, param_ids):
|
||||||
|
return {pid: float(flat_action[i]) for i, pid in enumerate(param_ids)}
|
||||||
|
|
||||||
|
|
||||||
|
def _build_param_space(param):
|
||||||
|
"""Return a gymnasium Box for a single NuconParameter, or None if unsupported."""
|
||||||
|
if param.param_type in (float, int):
|
||||||
|
lo = param.min_val if param.min_val is not None else -np.inf
|
||||||
|
hi = param.max_val if param.max_val is not None else np.inf
|
||||||
|
return spaces.Box(low=lo, high=hi, shape=(1,), dtype=np.float32)
|
||||||
|
elif param.param_type == bool:
|
||||||
|
return spaces.Box(low=0, high=1, shape=(1,), dtype=np.float32)
|
||||||
|
elif param.param_type == str:
|
||||||
|
return None
|
||||||
|
elif issubclass(param.param_type, Enum):
|
||||||
|
return spaces.Box(low=0, high=len(param.param_type) - 1, shape=(1,), dtype=np.float32)
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _apply_action(nucon, action):
|
||||||
|
for param_id, value in action.items():
|
||||||
|
param = nucon._parameters[param_id]
|
||||||
|
v = float(np.asarray(value).flat[0])
|
||||||
|
if param.param_type == bool:
|
||||||
|
value = v >= 0.5 # [0,1] space: above midpoint → True
|
||||||
|
elif issubclass(param.param_type, Enum):
|
||||||
|
value = param.param_type(int(v))
|
||||||
|
else:
|
||||||
|
value = param.param_type(v)
|
||||||
|
if param.min_val is not None and param.max_val is not None:
|
||||||
|
value = param.param_type(np.clip(value, param.min_val, param.max_val))
|
||||||
|
nucon.set(param, value)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# NuconEnv
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
class NuconEnv(gym.Env):
|
class NuconEnv(gym.Env):
|
||||||
metadata = {'render_modes': ['human']}
|
metadata = {'render_modes': ['human']}
|
||||||
|
|
||||||
def __init__(self, nucon=None, simulator=None, render_mode=None, seconds_per_step=5, objectives=['null'], terminators=['null'], objective_weights=None, terminate_above=0):
|
def __init__(self, nucon=None, simulator=None, render_mode=None, seconds_per_step=5,
|
||||||
|
objectives=['null'], terminators=['null'], objective_weights=None, terminate_above=0):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
|
||||||
self.render_mode = render_mode
|
self.render_mode = render_mode
|
||||||
self.seconds_per_step = seconds_per_step
|
self.seconds_per_step = seconds_per_step
|
||||||
if objective_weights is None:
|
if objective_weights is None:
|
||||||
objective_weights = [1.0 for objective in objectives]
|
objective_weights = [1.0 for _ in objectives]
|
||||||
self.objective_weights = objective_weights
|
self.objective_weights = objective_weights
|
||||||
self.terminate_above = terminate_above
|
self.terminate_above = terminate_above
|
||||||
self.simulator = simulator
|
self.simulator = simulator
|
||||||
|
|
||||||
if nucon is None:
|
if nucon is None:
|
||||||
if simulator:
|
nucon = Nucon(port=simulator.port) if simulator else Nucon()
|
||||||
nucon = Nucon(port=simulator.port)
|
|
||||||
else:
|
|
||||||
nucon = Nucon()
|
|
||||||
self.nucon = nucon
|
self.nucon = nucon
|
||||||
|
|
||||||
# Define observation space
|
# Observation space — SIM_UNCERTAINTY included when a simulator is present
|
||||||
obs_spaces = {'EPISODE_TIME': spaces.Box(low=0, high=np.inf, shape=(1,), dtype=np.float32)}
|
obs_spaces = {'EPISODE_TIME': spaces.Box(low=0, high=np.inf, shape=(1,), dtype=np.float32)}
|
||||||
|
if simulator is not None:
|
||||||
|
obs_spaces['SIM_UNCERTAINTY'] = spaces.Box(low=0.0, high=1.0, shape=(1,), dtype=np.float32)
|
||||||
for param_id, param in self.nucon.get_all_readable().items():
|
for param_id, param in self.nucon.get_all_readable().items():
|
||||||
if param.param_type == float:
|
sp = _build_param_space(param)
|
||||||
obs_spaces[param_id] = spaces.Box(low=param.min_val or -np.inf, high=param.max_val or np.inf, shape=(1,), dtype=np.float32)
|
if sp is not None:
|
||||||
elif param.param_type == int:
|
obs_spaces[param_id] = sp
|
||||||
if param.min_val is not None and param.max_val is not None:
|
|
||||||
obs_spaces[param_id] = spaces.Box(low=param.min_val, high=param.max_val, shape=(1,), dtype=np.float32)
|
|
||||||
else:
|
|
||||||
obs_spaces[param_id] = spaces.Box(low=-np.inf, high=np.inf, shape=(1,), dtype=np.float32)
|
|
||||||
elif param.param_type == bool:
|
|
||||||
obs_spaces[param_id] = spaces.Box(low=0, high=1, shape=(1,), dtype=np.float32)
|
|
||||||
elif issubclass(param.param_type, Enum):
|
|
||||||
obs_spaces[param_id] = spaces.Box(low=0, high=1, shape=(len(param.param_type),), dtype=np.float32)
|
|
||||||
else:
|
|
||||||
raise ValueError(f"Unsupported observation parameter type: {param.param_type}")
|
|
||||||
|
|
||||||
self.observation_space = spaces.Dict(obs_spaces)
|
self.observation_space = spaces.Dict(obs_spaces)
|
||||||
|
|
||||||
# Define action space
|
self.action_space, self._action_params, self._action_lows, self._action_ranges = \
|
||||||
action_spaces = {}
|
_build_flat_action_space(self.nucon)
|
||||||
for param_id, param in self.nucon.get_all_writable().items():
|
|
||||||
if param.param_type == float:
|
|
||||||
action_spaces[param_id] = spaces.Box(low=param.min_val or -np.inf, high=param.max_val or np.inf, shape=(1,), dtype=np.float32)
|
|
||||||
elif param.param_type == int:
|
|
||||||
if param.min_val is not None and param.max_val is not None:
|
|
||||||
action_spaces[param_id] = spaces.Box(low=param.min_val, high=param.max_val, shape=(1,), dtype=np.float32)
|
|
||||||
else:
|
|
||||||
action_spaces[param_id] = spaces.Box(low=-np.inf, high=np.inf, shape=(1,), dtype=np.float32)
|
|
||||||
elif param.param_type == bool:
|
|
||||||
action_spaces[param_id] = spaces.Box(low=0, high=1, shape=(1,), dtype=np.float32)
|
|
||||||
elif issubclass(param.param_type, Enum):
|
|
||||||
action_spaces[param_id] = spaces.Box(low=0, high=1, shape=(len(param.param_type),), dtype=np.float32)
|
|
||||||
else:
|
|
||||||
raise ValueError(f"Unsupported action parameter type: {param.param_type}")
|
|
||||||
|
|
||||||
self.action_space = spaces.Dict(action_spaces)
|
|
||||||
|
|
||||||
self.objectives = []
|
self.objectives = []
|
||||||
self.terminators = []
|
self.terminators = []
|
||||||
|
|
||||||
for objective in objectives:
|
for objective in objectives:
|
||||||
if objective in Objectives:
|
if objective in Objectives:
|
||||||
self.objectives.append(Objectives[objective])
|
self.objectives.append(Objectives[objective])
|
||||||
@@ -88,7 +171,6 @@ class NuconEnv(gym.Env):
|
|||||||
self.objectives.append(objective)
|
self.objectives.append(objective)
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"Unsupported objective: {objective}")
|
raise ValueError(f"Unsupported objective: {objective}")
|
||||||
|
|
||||||
for terminator in terminators:
|
for terminator in terminators:
|
||||||
if terminator in Objectives:
|
if terminator in Objectives:
|
||||||
self.terminators.append(Objectives[terminator])
|
self.terminators.append(Objectives[terminator])
|
||||||
@@ -97,75 +179,351 @@ class NuconEnv(gym.Env):
|
|||||||
else:
|
else:
|
||||||
raise ValueError(f"Unsupported terminator: {terminator}")
|
raise ValueError(f"Unsupported terminator: {terminator}")
|
||||||
|
|
||||||
def _get_obs(self):
|
def _get_obs(self, sim_uncertainty=None):
|
||||||
obs = {}
|
obs = {}
|
||||||
for param_id, param in self.nucon.get_all_readable().items():
|
for param_id, param in self.nucon.get_all_readable().items():
|
||||||
|
if param.param_type == str or param_id not in self.observation_space.spaces:
|
||||||
|
continue
|
||||||
value = self.nucon.get(param_id)
|
value = self.nucon.get(param_id)
|
||||||
if isinstance(value, Enum):
|
if isinstance(value, Enum):
|
||||||
value = value.value
|
value = value.value
|
||||||
obs[param_id] = value
|
obs[param_id] = value
|
||||||
obs["EPISODE_TIME"] = self._total_steps * self.seconds_per_step
|
obs['EPISODE_TIME'] = self._total_steps * self.seconds_per_step
|
||||||
|
if 'SIM_UNCERTAINTY' in self.observation_space.spaces:
|
||||||
|
obs['SIM_UNCERTAINTY'] = sim_uncertainty if sim_uncertainty is not None else 0.0
|
||||||
return obs
|
return obs
|
||||||
|
|
||||||
def _get_info(self):
|
def _get_info(self, obs):
|
||||||
info = {'objectives': {}, 'objectives_weighted': {}}
|
info = {'objectives': {}, 'objectives_weighted': {}}
|
||||||
for objective, weight in zip(self.objectives, self.objective_weights):
|
for objective, weight in zip(self.objectives, self.objective_weights):
|
||||||
obj = objective(self._get_obs())
|
obj = objective(obs)
|
||||||
info['objectives'][objective.__name__] = obj
|
name = getattr(objective, '__name__', repr(objective))
|
||||||
info['objectives_weighted'][objective.__name__] = obj * weight
|
info['objectives'][name] = obj
|
||||||
|
info['objectives_weighted'][name] = obj * weight
|
||||||
return info
|
return info
|
||||||
|
|
||||||
def reset(self, seed=None, options=None):
|
def reset(self, seed=None, options=None):
|
||||||
super().reset(seed=seed)
|
super().reset(seed=seed)
|
||||||
|
|
||||||
self._total_steps = 0
|
self._total_steps = 0
|
||||||
observation = self._get_obs()
|
observation = self._get_obs()
|
||||||
info = self._get_info()
|
return observation, self._get_info(observation)
|
||||||
|
|
||||||
return observation, info
|
|
||||||
|
|
||||||
def step(self, action):
|
def step(self, action):
|
||||||
# Apply the action to the Nucon system
|
_apply_action(self.nucon, _unflatten_action(action, self._action_params))
|
||||||
for param_id, value in action.items():
|
|
||||||
param = next(p for p in self.nucon if p.id == param_id)
|
|
||||||
if issubclass(param.param_type, Enum):
|
|
||||||
value = param.param_type(value)
|
|
||||||
if param.min_val is not None and param.max_val is not None:
|
|
||||||
value = np.clip(value, param.min_val, param.max_val)
|
|
||||||
self.nucon.set(param, value)
|
|
||||||
|
|
||||||
observation = self._get_obs()
|
# Advance sim (or sleep) — get uncertainty for obs injection
|
||||||
terminated = np.sum([terminator(observation) for terminator in self.terminators]) > self.terminate_above
|
|
||||||
truncated = False
|
truncated = False
|
||||||
info = self._get_info()
|
uncertainty = None
|
||||||
reward = sum(obj for obj in info['objectives_weighted'].values())
|
if self.simulator:
|
||||||
|
uncertainty = self.simulator.update(self.seconds_per_step, return_uncertainty=True)
|
||||||
|
else:
|
||||||
|
sim_speed = self.nucon.GAME_SIM_SPEED.value or 1.0
|
||||||
|
time.sleep(self.seconds_per_step / sim_speed)
|
||||||
|
|
||||||
self._total_steps += 1
|
self._total_steps += 1
|
||||||
if self.simulator:
|
observation = self._get_obs(sim_uncertainty=uncertainty)
|
||||||
self.simulator.update(self.seconds_per_step)
|
info = self._get_info(observation)
|
||||||
else:
|
reward = sum(obj for obj in info['objectives_weighted'].values())
|
||||||
time.sleep(self.seconds_per_step)
|
terminated = np.sum([t(observation) for t in self.terminators]) > self.terminate_above
|
||||||
return observation, reward, terminated, truncated, info
|
return observation, reward, terminated, truncated, info
|
||||||
|
|
||||||
def render(self):
|
def render(self):
|
||||||
if self.render_mode == "human":
|
|
||||||
pass
|
pass
|
||||||
|
|
||||||
def close(self):
|
def close(self):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
def _flatten_action(self, action):
|
|
||||||
return np.concatenate([v.flatten() for v in action.values()])
|
|
||||||
|
|
||||||
def _unflatten_action(self, flat_action):
|
|
||||||
return {k: v.reshape(1, -1) for k, v in self.action_space.items()}
|
|
||||||
|
|
||||||
def _flatten_observation(self, observation):
|
def _flatten_observation(self, observation):
|
||||||
return np.concatenate([v.flatten() for v in observation.values()])
|
return np.concatenate([np.asarray(v).flatten() for v in observation.values()])
|
||||||
|
|
||||||
def _unflatten_observation(self, flat_observation):
|
|
||||||
return {k: v.reshape(1, -1) for k, v in self.observation_space.items()}
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# NuconGoalEnv
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
class NuconGoalEnv(gym.Env):
|
||||||
|
"""
|
||||||
|
Goal-conditioned reactor environment compatible with SB3 HER (Hindsight Experience Replay).
|
||||||
|
|
||||||
|
Observation is a Dict with three keys:
|
||||||
|
- 'observation': all readable non-goal, non-str params + SIM_UNCERTAINTY (when sim active)
|
||||||
|
- 'achieved_goal': current values of goal_params, normalised to [0, 1] within goal_range
|
||||||
|
- 'desired_goal': target values sampled each episode, normalised to [0, 1]
|
||||||
|
|
||||||
|
``SIM_UNCERTAINTY`` in 'observation' lets reward_fn / terminators reference uncertainty directly.
|
||||||
|
|
||||||
|
reward_fn signature: ``(achieved, desired)`` or ``(achieved, desired, obs)`` — the 3-arg form
|
||||||
|
receives the full observation dict (including SIM_UNCERTAINTY) for uncertainty-aware shaping.
|
||||||
|
|
||||||
|
Usage with SB3 HER::
|
||||||
|
|
||||||
|
from stable_baselines3 import SAC
|
||||||
|
from stable_baselines3.common.buffers import HerReplayBuffer
|
||||||
|
from nucon.rl import NuconGoalEnv, UncertaintyPenalty, UncertaintyAbort
|
||||||
|
|
||||||
|
env = NuconGoalEnv(
|
||||||
|
goal_params=['GENERATOR_0_KW', 'GENERATOR_1_KW', 'GENERATOR_2_KW'],
|
||||||
|
goal_range={'GENERATOR_0_KW': (0, 1200), 'GENERATOR_1_KW': (0, 1200), 'GENERATOR_2_KW': (0, 1200)},
|
||||||
|
tolerance=0.05,
|
||||||
|
simulator=simulator,
|
||||||
|
# uncertainty-aware reward: penalise OOD, abort if too far out
|
||||||
|
reward_fn=lambda ag, dg, obs: (
|
||||||
|
-(np.linalg.norm(ag - dg) ** 2)
|
||||||
|
- 2.0 * max(0, obs.get('SIM_UNCERTAINTY', 0) - 0.3) ** 2
|
||||||
|
),
|
||||||
|
terminators=[UncertaintyAbort(threshold=0.7)],
|
||||||
|
)
|
||||||
|
model = SAC('MultiInputPolicy', env, replay_buffer_class=HerReplayBuffer)
|
||||||
|
model.learn(total_timesteps=500_000)
|
||||||
|
"""
|
||||||
|
|
||||||
|
metadata = {'render_modes': ['human']}
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
goal_params,
|
||||||
|
goal_range=None,
|
||||||
|
reward_fn=None,
|
||||||
|
tolerance=None,
|
||||||
|
nucon=None,
|
||||||
|
simulator=None,
|
||||||
|
render_mode=None,
|
||||||
|
seconds_per_step=5,
|
||||||
|
terminators=None,
|
||||||
|
terminate_above=0,
|
||||||
|
additional_objectives=None,
|
||||||
|
additional_objective_weights=None,
|
||||||
|
obs_params=None,
|
||||||
|
action_params=None,
|
||||||
|
init_states=None,
|
||||||
|
delta_action_scale=None,
|
||||||
|
goal_sampling_std=None,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
self.render_mode = render_mode
|
||||||
|
self.seconds_per_step = seconds_per_step
|
||||||
|
self._delta_action_scale = delta_action_scale
|
||||||
|
self.terminate_above = terminate_above
|
||||||
|
self.simulator = simulator
|
||||||
|
self.goal_params = list(goal_params)
|
||||||
|
self.tolerance = tolerance
|
||||||
|
|
||||||
|
if nucon is None:
|
||||||
|
nucon = Nucon(port=simulator.port) if simulator else Nucon()
|
||||||
|
self.nucon = nucon
|
||||||
|
|
||||||
|
all_readable = self.nucon.get_all_readable()
|
||||||
|
for pid in self.goal_params:
|
||||||
|
if pid not in all_readable:
|
||||||
|
raise ValueError(f"Goal param '{pid}' is not a readable parameter")
|
||||||
|
|
||||||
|
goal_range = goal_range or {}
|
||||||
|
self._goal_low = np.array([
|
||||||
|
goal_range.get(pid, (all_readable[pid].min_val or 0.0, all_readable[pid].max_val or 1.0))[0]
|
||||||
|
for pid in self.goal_params
|
||||||
|
], dtype=np.float32)
|
||||||
|
self._goal_high = np.array([
|
||||||
|
goal_range.get(pid, (all_readable[pid].min_val or 0.0, all_readable[pid].max_val or 1.0))[1]
|
||||||
|
for pid in self.goal_params
|
||||||
|
], dtype=np.float32)
|
||||||
|
self._goal_range = self._goal_high - self._goal_low
|
||||||
|
self._goal_range[self._goal_range == 0] = 1.0
|
||||||
|
|
||||||
|
# Detect reward_fn arity for backward compat (2-arg vs 3-arg)
|
||||||
|
self._reward_fn = reward_fn
|
||||||
|
if reward_fn is not None:
|
||||||
|
n_args = len(inspect.signature(reward_fn).parameters)
|
||||||
|
self._reward_fn_wants_obs = n_args >= 3
|
||||||
|
else:
|
||||||
|
self._reward_fn_wants_obs = False
|
||||||
|
|
||||||
|
# Observation params: model.input_params defines the canonical list — the same set is
|
||||||
|
# used whether training in sim or deploying to the real game (the game simply has more
|
||||||
|
# params available; we query only the subset we care about).
|
||||||
|
# Explicit obs_params overrides everything (use when deploying to real game without sim).
|
||||||
|
# SB3 HER requires observation to be a flat Box, not a nested Dict.
|
||||||
|
goal_set = set(self.goal_params)
|
||||||
|
self._obs_with_uncertainty = simulator is not None
|
||||||
|
if obs_params is not None:
|
||||||
|
base_params = [p for p in obs_params if p not in goal_set]
|
||||||
|
elif simulator is not None and hasattr(simulator, 'model') and simulator.model is not None:
|
||||||
|
base_params = [p for p in simulator.model.input_params
|
||||||
|
if p not in goal_set and p in all_readable
|
||||||
|
and _build_param_space(all_readable[p]) is not None]
|
||||||
|
else:
|
||||||
|
base_params = [p for p, param in all_readable.items()
|
||||||
|
if p not in goal_set and _build_param_space(param) is not None]
|
||||||
|
# SIM_UNCERTAINTY is not in _obs_params — it's not available at deployment on the real game
|
||||||
|
self._obs_params = base_params
|
||||||
|
|
||||||
|
n_goals = len(self.goal_params)
|
||||||
|
self.observation_space = spaces.Dict({
|
||||||
|
'observation': spaces.Box(low=-np.inf, high=np.inf,
|
||||||
|
shape=(len(self._obs_params),), dtype=np.float32),
|
||||||
|
'achieved_goal': spaces.Box(low=0.0, high=1.0, shape=(n_goals,), dtype=np.float32),
|
||||||
|
'desired_goal': spaces.Box(low=0.0, high=1.0, shape=(n_goals,), dtype=np.float32),
|
||||||
|
})
|
||||||
|
|
||||||
|
# Action space: writable params within the obs param set, or an explicit override list.
|
||||||
|
action_set = set(action_params) if action_params is not None else set(base_params)
|
||||||
|
self.action_space, self._action_params, self._action_lows, self._action_ranges = \
|
||||||
|
_build_flat_action_space(self.nucon, action_set, delta_action_scale)
|
||||||
|
|
||||||
|
self._terminators = terminators or []
|
||||||
|
_objs = additional_objectives or []
|
||||||
|
self._objectives = [Objectives[o] if isinstance(o, str) else o for o in _objs]
|
||||||
|
self._objective_weights = additional_objective_weights or [1.0] * len(self._objectives)
|
||||||
|
self._init_states = init_states # list of state dicts to sample on reset
|
||||||
|
self._goal_sampling_std = goal_sampling_std # Gaussian std in normalised goal space; None → uniform
|
||||||
|
self._desired_goal = np.zeros(n_goals, dtype=np.float32)
|
||||||
|
self._total_steps = 0
|
||||||
|
|
||||||
|
def compute_reward(self, achieved_goal, desired_goal, info):
|
||||||
|
"""Dense negative L2, sparse with tolerance, or custom reward_fn."""
|
||||||
|
obs_named = info.get('obs_named', {}) if isinstance(info, dict) else {}
|
||||||
|
if self._reward_fn is not None:
|
||||||
|
if self._reward_fn_wants_obs:
|
||||||
|
return self._reward_fn(achieved_goal, desired_goal, obs_named)
|
||||||
|
return self._reward_fn(achieved_goal, desired_goal)
|
||||||
|
dist = np.linalg.norm(achieved_goal - desired_goal, axis=-1)
|
||||||
|
if self.tolerance is not None:
|
||||||
|
return (dist <= self.tolerance).astype(np.float32) - 1.0
|
||||||
|
return -dist
|
||||||
|
|
||||||
|
def _read_goal_values(self):
|
||||||
|
raw = np.array([self.nucon.get(pid) or 0.0 for pid in self.goal_params], dtype=np.float32)
|
||||||
|
return np.clip((raw - self._goal_low) / self._goal_range, 0.0, 1.0)
|
||||||
|
|
||||||
|
def _read_obs(self, sim_uncertainty=None):
|
||||||
|
"""Return (gym_obs_dict, reward_obs_dict).
|
||||||
|
|
||||||
|
When a simulator is attached, reads directly from sim.parameters (no HTTP).
|
||||||
|
Otherwise falls back to a single batch HTTP request.
|
||||||
|
"""
|
||||||
|
def _to_float(v):
|
||||||
|
if v is None:
|
||||||
|
return 0.0
|
||||||
|
return float(v.value if isinstance(v, Enum) else v)
|
||||||
|
|
||||||
|
if self.simulator is not None:
|
||||||
|
# Direct in-process read — no HTTP overhead
|
||||||
|
def _get(pid):
|
||||||
|
return _to_float(self.simulator.get(pid))
|
||||||
|
else:
|
||||||
|
raw = self.nucon._batch_query(self._obs_params + self.goal_params)
|
||||||
|
all_params = self.nucon.get_all_readable()
|
||||||
|
def _get(pid):
|
||||||
|
try:
|
||||||
|
v = self.nucon._parse_value(all_params[pid], raw.get(pid, '0'))
|
||||||
|
return _to_float(v)
|
||||||
|
except Exception:
|
||||||
|
return 0.0
|
||||||
|
|
||||||
|
reward_obs = {}
|
||||||
|
if self._obs_with_uncertainty:
|
||||||
|
reward_obs['SIM_UNCERTAINTY'] = float(sim_uncertainty) if sim_uncertainty is not None else 0.0
|
||||||
|
for pid in self._obs_params:
|
||||||
|
reward_obs[pid] = _get(pid)
|
||||||
|
|
||||||
|
obs_vec = np.array([reward_obs[p] for p in self._obs_params], dtype=np.float32)
|
||||||
|
goal_raw = np.array([_get(p) for p in self.goal_params], dtype=np.float32)
|
||||||
|
achieved = np.clip((goal_raw - self._goal_low) / self._goal_range, 0.0, 1.0)
|
||||||
|
gym_obs = {'observation': obs_vec, 'achieved_goal': achieved,
|
||||||
|
'desired_goal': self._desired_goal.copy()}
|
||||||
|
return gym_obs, reward_obs
|
||||||
|
|
||||||
|
def reset(self, seed=None, options=None):
|
||||||
|
super().reset(seed=seed)
|
||||||
|
self._total_steps = 0
|
||||||
|
rng = np.random.default_rng(seed)
|
||||||
|
if self._init_states is not None and self.simulator is not None:
|
||||||
|
state = self._init_states[rng.integers(len(self._init_states))]
|
||||||
|
for k, v in state.items():
|
||||||
|
try:
|
||||||
|
self.simulator.set(k, v, force=True)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
if self._goal_sampling_std is not None:
|
||||||
|
# Sample goal as Gaussian delta from current state — usually a small change,
|
||||||
|
# occasionally a large one.
|
||||||
|
current = np.array([
|
||||||
|
float(self.simulator.get(p) if self.simulator else 0.0)
|
||||||
|
for p in self.goal_params
|
||||||
|
], dtype=np.float32)
|
||||||
|
current_norm = np.clip((current - self._goal_low) / self._goal_range, 0.0, 1.0)
|
||||||
|
delta = rng.normal(0.0, self._goal_sampling_std, size=len(self.goal_params))
|
||||||
|
self._desired_goal = np.clip(current_norm + delta, 0.0, 1.0).astype(np.float32)
|
||||||
|
else:
|
||||||
|
self._desired_goal = rng.uniform(0.0, 1.0, size=len(self.goal_params)).astype(np.float32)
|
||||||
|
gym_obs, _ = self._read_obs()
|
||||||
|
return gym_obs, {}
|
||||||
|
|
||||||
|
def step(self, action):
|
||||||
|
flat = np.asarray(action, dtype=np.float32)
|
||||||
|
if self._delta_action_scale is not None:
|
||||||
|
# Compute absolute values from deltas, reading current state
|
||||||
|
if self.simulator is None:
|
||||||
|
raw_current = self.nucon._batch_query(self._action_params)
|
||||||
|
all_params = self.nucon.get_all_readable()
|
||||||
|
absolute = {}
|
||||||
|
for i, pid in enumerate(self._action_params):
|
||||||
|
param = self.nucon._parameters[pid]
|
||||||
|
if param.param_type == bool:
|
||||||
|
absolute[pid] = 1.0 if flat[i] > 0 else 0.0
|
||||||
|
else:
|
||||||
|
if self.simulator is not None:
|
||||||
|
v = self.simulator.get(pid)
|
||||||
|
current = float(v.value if isinstance(v, Enum) else v) if v is not None else 0.0
|
||||||
|
else:
|
||||||
|
try:
|
||||||
|
v = self.nucon._parse_value(all_params[pid], raw_current.get(pid, '0'))
|
||||||
|
current = float(v.value if isinstance(v, Enum) else v)
|
||||||
|
except Exception:
|
||||||
|
current = 0.0
|
||||||
|
delta = float(flat[i]) * self._delta_action_scale * self._action_ranges[i]
|
||||||
|
absolute[pid] = float(np.clip(current + delta,
|
||||||
|
self._action_lows[i],
|
||||||
|
self._action_lows[i] + self._action_ranges[i]))
|
||||||
|
else:
|
||||||
|
absolute = _unflatten_action(flat, self._action_params)
|
||||||
|
|
||||||
|
if self.simulator is not None:
|
||||||
|
# Write directly to sim — skip HTTP entirely
|
||||||
|
for pid, val in absolute.items():
|
||||||
|
try:
|
||||||
|
self.simulator.set(pid, val, force=True)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
else:
|
||||||
|
_apply_action(self.nucon, absolute)
|
||||||
|
|
||||||
|
if self.simulator:
|
||||||
|
uncertainty = self.simulator.update(self.seconds_per_step, return_uncertainty=True)
|
||||||
|
else:
|
||||||
|
sim_speed = self.nucon.GAME_SIM_SPEED.value or 1.0
|
||||||
|
time.sleep(self.seconds_per_step / sim_speed)
|
||||||
|
uncertainty = None
|
||||||
|
|
||||||
|
self._total_steps += 1
|
||||||
|
gym_obs, reward_obs = self._read_obs(sim_uncertainty=uncertainty)
|
||||||
|
info = {'achieved_goal': gym_obs['achieved_goal'], 'desired_goal': gym_obs['desired_goal'],
|
||||||
|
'obs_named': reward_obs}
|
||||||
|
reward = float(self.compute_reward(gym_obs['achieved_goal'], gym_obs['desired_goal'], info))
|
||||||
|
reward += sum(w * o(reward_obs) for o, w in zip(self._objectives, self._objective_weights))
|
||||||
|
terminated = any(t(reward_obs) > self.terminate_above for t in self._terminators)
|
||||||
|
return gym_obs, reward, terminated, False, info
|
||||||
|
|
||||||
|
def render(self):
|
||||||
|
pass
|
||||||
|
|
||||||
|
def close(self):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Registration
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
def register_nucon_envs():
|
def register_nucon_envs():
|
||||||
gym.register(
|
gym.register(
|
||||||
@@ -181,7 +539,28 @@ def register_nucon_envs():
|
|||||||
gym.register(
|
gym.register(
|
||||||
id='Nucon-safe_max_power-v0',
|
id='Nucon-safe_max_power-v0',
|
||||||
entry_point='nucon.rl:NuconEnv',
|
entry_point='nucon.rl:NuconEnv',
|
||||||
kwargs={'seconds_per_step': 5, 'objectives': [Parameterized_Objectives['temp_above'](min_temp=310), Parameterized_Objectives['temp_below'](max_temp=365), 'max_power'], 'objective_weights': [1, 10, 1/100_000]}
|
kwargs={'seconds_per_step': 5,
|
||||||
|
'objectives': [Parameterized_Objectives['temp_above'](min_temp=310),
|
||||||
|
Parameterized_Objectives['temp_below'](max_temp=365), 'max_power'],
|
||||||
|
'objective_weights': [1, 10, 1/100_000]}
|
||||||
|
)
|
||||||
|
gym.register(
|
||||||
|
id='Nucon-goal_power-v0',
|
||||||
|
entry_point='nucon.rl:NuconGoalEnv',
|
||||||
|
kwargs={
|
||||||
|
'goal_params': ['GENERATOR_0_KW', 'GENERATOR_1_KW', 'GENERATOR_2_KW'],
|
||||||
|
'goal_range': {'GENERATOR_0_KW': (0.0, 1200.0), 'GENERATOR_1_KW': (0.0, 1200.0), 'GENERATOR_2_KW': (0.0, 1200.0)},
|
||||||
|
'seconds_per_step': 5,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
gym.register(
|
||||||
|
id='Nucon-goal_temp-v0',
|
||||||
|
entry_point='nucon.rl:NuconGoalEnv',
|
||||||
|
kwargs={
|
||||||
|
'goal_params': ['CORE_TEMP'],
|
||||||
|
'goal_range': {'CORE_TEMP': (280.0, 380.0)},
|
||||||
|
'seconds_per_step': 5,
|
||||||
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
register_nucon_envs()
|
register_nucon_envs()
|
||||||
+62
-17
@@ -5,7 +5,8 @@ from flask import Flask, request, jsonify
|
|||||||
from nucon import Nucon, ParameterEnum, PumpStatus, PumpDryStatus, PumpOverloadStatus, BreakerStatus
|
from nucon import Nucon, ParameterEnum, PumpStatus, PumpDryStatus, PumpOverloadStatus, BreakerStatus
|
||||||
import threading
|
import threading
|
||||||
import torch
|
import torch
|
||||||
from nucon.model import ReactorDynamicsModel
|
from nucon.model import ReactorDynamicsModel, ReactorKNNModel
|
||||||
|
import pickle
|
||||||
|
|
||||||
class OperatingState(Enum):
|
class OperatingState(Enum):
|
||||||
# Tuple indicates a range of values, while list indicates a set of possible values
|
# Tuple indicates a range of values, while list indicates a set of possible values
|
||||||
@@ -165,6 +166,8 @@ class NuconSimulator:
|
|||||||
def __init__(self, host: str = 'localhost', port: int = 8786):
|
def __init__(self, host: str = 'localhost', port: int = 8786):
|
||||||
self._nucon = Nucon()
|
self._nucon = Nucon()
|
||||||
self.parameters = self.Parameters(self._nucon)
|
self.parameters = self.Parameters(self._nucon)
|
||||||
|
self.host = host
|
||||||
|
self.port = port
|
||||||
self.time = 0.0
|
self.time = 0.0
|
||||||
self.allow_all_writes = False
|
self.allow_all_writes = False
|
||||||
self.set_state(OperatingState.OFFLINE)
|
self.set_state(OperatingState.OFFLINE)
|
||||||
@@ -212,38 +215,72 @@ class NuconSimulator:
|
|||||||
def set_allow_all_writes(self, allow: bool) -> None:
|
def set_allow_all_writes(self, allow: bool) -> None:
|
||||||
self.allow_all_writes = allow
|
self.allow_all_writes = allow
|
||||||
|
|
||||||
def update(self, time_step: float) -> None:
|
def update(self, time_step: float, return_uncertainty: bool = False):
|
||||||
self._update_reactor_state(time_step)
|
"""Advance the simulator by time_step game-seconds.
|
||||||
|
|
||||||
|
If return_uncertainty=True and a kNN model is loaded, returns the GP
|
||||||
|
posterior std for this step (0 = on known data, ~1 = OOD).
|
||||||
|
Always returns None when using an NN model.
|
||||||
|
"""
|
||||||
|
uncertainty = self._update_reactor_state(time_step, return_uncertainty=return_uncertainty)
|
||||||
self.time += time_step
|
self.time += time_step
|
||||||
|
return uncertainty
|
||||||
|
|
||||||
|
def set_model(self, model) -> None:
|
||||||
|
"""Set a pre-loaded ReactorDynamicsModel or ReactorKNNModel directly."""
|
||||||
|
self.model = model
|
||||||
|
if isinstance(model, ReactorDynamicsModel):
|
||||||
|
self.model.eval()
|
||||||
|
|
||||||
def load_model(self, model_path: str) -> None:
|
def load_model(self, model_path: str) -> None:
|
||||||
|
"""Load a model from a file. .pkl → ReactorKNNModel, otherwise → ReactorDynamicsModel (torch)."""
|
||||||
try:
|
try:
|
||||||
|
if model_path.endswith('.pkl'):
|
||||||
|
with open(model_path, 'rb') as f:
|
||||||
|
self.model = pickle.load(f)
|
||||||
|
print(f"kNN model loaded from {model_path}")
|
||||||
|
else:
|
||||||
|
# Reconstruct shell from the saved state dict; input/output params
|
||||||
|
# are stored inside the checkpoint.
|
||||||
|
checkpoint = torch.load(model_path, weights_only=False)
|
||||||
|
if isinstance(checkpoint, dict) and 'input_params' in checkpoint:
|
||||||
|
self.model = ReactorDynamicsModel(checkpoint['input_params'], checkpoint['output_params'])
|
||||||
|
self.model.load_state_dict(checkpoint['state_dict'])
|
||||||
|
else:
|
||||||
|
# Legacy: plain state dict — fall back using sim readable/non-writable lists
|
||||||
self.model = ReactorDynamicsModel(self.readable_params, self.non_writable_params)
|
self.model = ReactorDynamicsModel(self.readable_params, self.non_writable_params)
|
||||||
self.model.load_state_dict(torch.load(model_path))
|
self.model.load_state_dict(checkpoint)
|
||||||
self.model.eval() # Set the model to evaluation mode
|
self.model.eval()
|
||||||
print(f"Model loaded successfully from {model_path}")
|
print(f"NN model loaded from {model_path}")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
print(f"Error loading model: {str(e)}")
|
print(f"Error loading model: {str(e)}")
|
||||||
self.model = None
|
self.model = None
|
||||||
|
|
||||||
def _update_reactor_state(self, time_step: float) -> None:
|
def _update_reactor_state(self, time_step: float, return_uncertainty: bool = False):
|
||||||
if not self.model:
|
if not self.model:
|
||||||
raise ValueError("Model not set. Please load a model using load_model() method.")
|
raise ValueError("Model not set. Please load a model using load_model() or set_model().")
|
||||||
|
|
||||||
|
# Build state dict using only the params the model knows about
|
||||||
|
params = self.parameters
|
||||||
state = {}
|
state = {}
|
||||||
for param in self.readable_params:
|
for param_id in self.model.input_params:
|
||||||
value = self.get(param)
|
value = getattr(params, param_id, None)
|
||||||
if isinstance(value, Enum):
|
if isinstance(value, Enum):
|
||||||
value = value.value
|
value = value.value
|
||||||
state[param] = value
|
state[param_id] = 0.0 if value is None else value
|
||||||
|
|
||||||
# Use the model to predict the next state
|
# Forward pass
|
||||||
with torch.no_grad():
|
uncertainty = None
|
||||||
next_state = self.model(state, time_step)
|
if return_uncertainty:
|
||||||
|
next_state, uncertainty = self.model.forward_with_uncertainty(state, time_step)
|
||||||
|
else:
|
||||||
|
next_state = self.model.forward(state, time_step)
|
||||||
|
|
||||||
# Update the simulator's state
|
# Write outputs directly — bypass sim.set() type-checking overhead
|
||||||
for param, value in next_state.items():
|
for param_id, value in next_state.items():
|
||||||
self.set(param, value)
|
setattr(params, param_id, value)
|
||||||
|
|
||||||
|
return uncertainty
|
||||||
|
|
||||||
def set_state(self, state: OperatingState) -> None:
|
def set_state(self, state: OperatingState) -> None:
|
||||||
self._sample_parameters_from_state(state)
|
self._sample_parameters_from_state(state)
|
||||||
@@ -286,6 +323,14 @@ class NuconSimulator:
|
|||||||
|
|
||||||
try:
|
try:
|
||||||
value = self.get(variable)
|
value = self.get(variable)
|
||||||
|
if value is None:
|
||||||
|
param = self._nucon[variable]
|
||||||
|
if param.enum_type is not None:
|
||||||
|
value = next(iter(param.enum_type)).value # first enum member's int value
|
||||||
|
else:
|
||||||
|
value = param.param_type() # int()->0, float()->0.0, bool()->False
|
||||||
|
if isinstance(value, Enum):
|
||||||
|
value = value.value
|
||||||
return str(value), 200
|
return str(value), 200
|
||||||
except (KeyError, AttributeError):
|
except (KeyError, AttributeError):
|
||||||
return jsonify({"error": f"Unknown variable: {variable}"}), 404
|
return jsonify({"error": f"Unknown variable: {variable}"}), 404
|
||||||
|
|||||||
@@ -0,0 +1,55 @@
|
|||||||
|
"""Collect a dynamics dataset from the running Nucleares game.
|
||||||
|
|
||||||
|
Play the game normally while this script runs in the background.
|
||||||
|
It records state transitions every `time_delta` game-seconds and
|
||||||
|
saves them incrementally so nothing is lost if you quit early.
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
python scripts/collect_dataset.py # default settings
|
||||||
|
python scripts/collect_dataset.py --steps 2000 --delta 5 # faster sampling
|
||||||
|
python scripts/collect_dataset.py --out my_dataset.pkl
|
||||||
|
|
||||||
|
The saved dataset is a list of (state_before, action_dict, state_after, time_delta)
|
||||||
|
tuples compatible with NuconModelLearner.fit_knn() and train_model().
|
||||||
|
|
||||||
|
Tips for good data:
|
||||||
|
- Cover a range of operating states: startup, ramp, steady-state, shutdown.
|
||||||
|
- Vary individual rod bank positions, pump speeds, and MSCV setpoints.
|
||||||
|
- Collect at least 500 samples for kNN-GP; 5000+ for the NN backend.
|
||||||
|
- Merge multiple sessions with NuconModelLearner.merge_datasets().
|
||||||
|
"""
|
||||||
|
import argparse
|
||||||
|
import pickle
|
||||||
|
from nucon.model import NuconModelLearner
|
||||||
|
|
||||||
|
parser = argparse.ArgumentParser()
|
||||||
|
parser.add_argument('--steps', type=int, default=1000,
|
||||||
|
help='Number of samples to collect (default: 1000)')
|
||||||
|
parser.add_argument('--delta', type=float, default=10.0,
|
||||||
|
help='Game-seconds between samples (default: 10.0)')
|
||||||
|
parser.add_argument('--out', default='reactor_dataset.pkl',
|
||||||
|
help='Output path for dataset (default: reactor_dataset.pkl)')
|
||||||
|
parser.add_argument('--merge', default=None,
|
||||||
|
help='Existing dataset to merge into before saving')
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
learner = NuconModelLearner(
|
||||||
|
time_delta=args.delta,
|
||||||
|
dataset_path=args.out,
|
||||||
|
)
|
||||||
|
|
||||||
|
if args.merge:
|
||||||
|
learner.merge_datasets(args.merge)
|
||||||
|
print(f"Merged existing dataset from {args.merge} ({len(learner.dataset)} samples)")
|
||||||
|
|
||||||
|
print(f"Collecting {args.steps} samples (Δt={args.delta}s each) → {args.out}")
|
||||||
|
print("Play the game — vary rod positions, pump speeds, and operating states.")
|
||||||
|
print("Press Ctrl-C to stop early; data collected so far will be saved.")
|
||||||
|
|
||||||
|
try:
|
||||||
|
learner.collect_data(num_steps=args.steps)
|
||||||
|
except KeyboardInterrupt:
|
||||||
|
print("\nInterrupted — saving collected data...")
|
||||||
|
|
||||||
|
learner.save_dataset(args.out)
|
||||||
|
print(f"Saved {len(learner.dataset)} samples to {args.out}")
|
||||||
@@ -0,0 +1,739 @@
|
|||||||
|
"""Classical PID-based reactor controller with curses TUI.
|
||||||
|
|
||||||
|
Architecture:
|
||||||
|
Core control (shared):
|
||||||
|
- Rod PID: keeps CORE_TEMP at setpoint via ROD_BANK_POS_0_ORDERED
|
||||||
|
|
||||||
|
Per-train control (trains 1/2/3, 0-indexed as 0/1/2 in param names):
|
||||||
|
- Primary pump: not touched; warns in TUI if far from suggested 65%
|
||||||
|
- MSCV PI: drives train power output, gated on steam availability
|
||||||
|
- Secondary pump feedforward: half of steam outlet + level PID
|
||||||
|
- Bypass: hold at 0
|
||||||
|
|
||||||
|
Auxiliary:
|
||||||
|
- Vacuum pump: on continuously; turned off only during retention tank drain
|
||||||
|
- Condenser circulation pump: fixed 25% (prevents overcooling of return water)
|
||||||
|
- Retention tank: drain via ejector return valve when > 75%, stop at 50%
|
||||||
|
- Condenser fill: run FREIGHT_PUMP_CONDENSER below 45%, stop at 60%
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
python3.14 scripts/reactor_control.py --trains 3 --target 50000
|
||||||
|
python3.14 scripts/reactor_control.py --trains 1 3 --target 30000 40000
|
||||||
|
python3.14 scripts/reactor_control.py --trains 1 2 3 --target 20000 20000 20000
|
||||||
|
|
||||||
|
TUI keys:
|
||||||
|
0 Select core (then +/- adjusts temp setpoint ±5°C)
|
||||||
|
1 / 2 / 3 Select train (then +/- adjusts target power ±5 MW; + adds if absent)
|
||||||
|
d Remove selected train from control
|
||||||
|
g Toggle grid-demand following
|
||||||
|
q / Esc Quit
|
||||||
|
"""
|
||||||
|
import argparse
|
||||||
|
import curses
|
||||||
|
import time
|
||||||
|
import numpy as np
|
||||||
|
from enum import Enum
|
||||||
|
from nucon import Nucon
|
||||||
|
|
||||||
|
parser = argparse.ArgumentParser()
|
||||||
|
parser.add_argument('--trains', type=int, nargs='+', default=[3])
|
||||||
|
parser.add_argument('--target', type=float, nargs='+', default=[50_000])
|
||||||
|
parser.add_argument('--temp-setpoint', type=float, default=330.0)
|
||||||
|
parser.add_argument('--dt', type=float, default=5.0)
|
||||||
|
parser.add_argument('--grid-follow', action='store_true',
|
||||||
|
help='Auto-set train targets from grid demand')
|
||||||
|
parser.add_argument('--grid-buffer', type=float, default=10.0,
|
||||||
|
help='Extra MW above grid demand when grid-following (default: 5)')
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
if len(args.target) == 1:
|
||||||
|
targets = {t: args.target[0] for t in args.trains}
|
||||||
|
else:
|
||||||
|
if len(args.target) != len(args.trains):
|
||||||
|
raise ValueError("--target must have 1 value or one per --trains entry")
|
||||||
|
targets = dict(zip(args.trains, args.target))
|
||||||
|
|
||||||
|
nucon = Nucon()
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# PID controller
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
class PID:
|
||||||
|
def __init__(self, kp, ki, kd, out_min, out_max, integral_max=None):
|
||||||
|
self.kp, self.ki, self.kd = kp, ki, kd
|
||||||
|
self.out_min, self.out_max = out_min, out_max
|
||||||
|
self.integral_max = integral_max or (out_max - out_min)
|
||||||
|
self._integral = 0.0
|
||||||
|
self._prev_error = None
|
||||||
|
|
||||||
|
def step(self, error, dt):
|
||||||
|
self._integral = np.clip(self._integral + error * dt,
|
||||||
|
-self.integral_max, self.integral_max)
|
||||||
|
derivative = 0.0 if self._prev_error is None else (error - self._prev_error) / dt
|
||||||
|
self._prev_error = error
|
||||||
|
return float(np.clip(
|
||||||
|
self.kp * error + self.ki * self._integral + self.kd * derivative,
|
||||||
|
self.out_min, self.out_max))
|
||||||
|
|
||||||
|
def reset(self):
|
||||||
|
self._integral = 0.0
|
||||||
|
self._prev_error = None
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Helpers
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
_all_readable = None
|
||||||
|
def _get_all_readable():
|
||||||
|
global _all_readable
|
||||||
|
if _all_readable is None:
|
||||||
|
_all_readable = nucon.get_all_readable()
|
||||||
|
return _all_readable
|
||||||
|
|
||||||
|
def set_param(param_id, value):
|
||||||
|
param = nucon._parameters[param_id]
|
||||||
|
v = float(np.clip(value, param.min_val or 0, param.max_val or 100))
|
||||||
|
nucon.set(param, v)
|
||||||
|
return v
|
||||||
|
|
||||||
|
def read_state(param_ids):
|
||||||
|
all_r = _get_all_readable()
|
||||||
|
raw = nucon._batch_query([p for p in param_ids if p in all_r])
|
||||||
|
state = {}
|
||||||
|
for p in param_ids:
|
||||||
|
if p not in all_r:
|
||||||
|
state[p] = 0.0
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
v = nucon._parse_value(all_r[p], raw.get(p, '0'))
|
||||||
|
state[p] = float(v.value if isinstance(v, Enum) else v)
|
||||||
|
except Exception:
|
||||||
|
state[p] = 0.0
|
||||||
|
return state
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Per-train controller
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
class TrainController:
|
||||||
|
"""Controls one train (steam gen N + turbine N + generator N)."""
|
||||||
|
|
||||||
|
def __init__(self, train_num, target_kw):
|
||||||
|
self.n = train_num
|
||||||
|
self.i = train_num - 1
|
||||||
|
self.target_kw = target_kw
|
||||||
|
|
||||||
|
self.mscv_pid = PID(kp=0.00002, ki=0.000002, kd=0.0,
|
||||||
|
out_min=-0.3, out_max=0.2, integral_max=3.0)
|
||||||
|
self._prev_steam_out = None
|
||||||
|
self.sec_pid = PID(kp=0.0005, ki=0.00005, kd=0.001,
|
||||||
|
out_min=-2.0, out_max=2.0, integral_max=3.0)
|
||||||
|
self.sec_level_target = 25_000.0
|
||||||
|
|
||||||
|
self.prim_pump = float(nucon.get(f'COOLANT_CORE_CIRCULATION_PUMP_{self.i}_ORDERED_SPEED') or 50.0)
|
||||||
|
self.PRIM_PUMP_SUGGESTED = 65.0 # warn in TUI if far from this
|
||||||
|
self.mscv = 9.0
|
||||||
|
self.sec_pump = 40.0
|
||||||
|
|
||||||
|
set_param(f'STEAM_TURBINE_{self.i}_BYPASS_ORDERED', 0.0)
|
||||||
|
|
||||||
|
self._params = [
|
||||||
|
f'STEAM_GEN_{self.i}_OUTLET',
|
||||||
|
f'MSCV_{self.i}_OPENING_ACTUAL',
|
||||||
|
f'STEAM_TURBINE_{self.i}_RPM',
|
||||||
|
f'STEAM_TURBINE_{self.i}_BYPASS_ACTUAL',
|
||||||
|
f'GENERATOR_{self.i}_KW',
|
||||||
|
f'COOLANT_CORE_CIRCULATION_PUMP_{self.i}_ORDERED_SPEED',
|
||||||
|
f'COOLANT_SEC_CIRCULATION_PUMP_{self.i}_ORDERED_SPEED',
|
||||||
|
f'COOLANT_SEC_{self.i}_LIQUID_VOLUME',
|
||||||
|
]
|
||||||
|
|
||||||
|
def params(self):
|
||||||
|
return self._params
|
||||||
|
|
||||||
|
def step(self, s, dt):
|
||||||
|
steam_out = s[f'STEAM_GEN_{self.i}_OUTLET']
|
||||||
|
power_kw = s[f'GENERATOR_{self.i}_KW']
|
||||||
|
power_error = self.target_kw - power_kw
|
||||||
|
|
||||||
|
# Dead-band: don't adjust MSCV when within 3% of target (avoid hunting)
|
||||||
|
if abs(power_error) < 0.03 * self.target_kw:
|
||||||
|
mscv_delta = 0.0
|
||||||
|
self.mscv_pid.reset()
|
||||||
|
else:
|
||||||
|
mscv_delta = self.mscv_pid.step(power_error, dt)
|
||||||
|
steam_rose = (self._prev_steam_out is None or
|
||||||
|
steam_out >= self._prev_steam_out - 1.0)
|
||||||
|
if mscv_delta > 0 and not steam_rose:
|
||||||
|
mscv_delta = 0.0
|
||||||
|
self._prev_steam_out = steam_out
|
||||||
|
# Cap only prevents opening further — don't force MSCV down as steam fluctuates.
|
||||||
|
mscv_max = max(steam_out / 8.0, 1.0)
|
||||||
|
new_mscv = self.mscv + mscv_delta
|
||||||
|
if mscv_delta > 0:
|
||||||
|
new_mscv = min(new_mscv, mscv_max)
|
||||||
|
self.mscv = float(np.clip(new_mscv, 0.5, 100.0))
|
||||||
|
set_param(f'MSCV_{self.i}_OPENING_ORDERED', self.mscv)
|
||||||
|
|
||||||
|
self.prim_pump = s.get(f'COOLANT_CORE_CIRCULATION_PUMP_{self.i}_ORDERED_SPEED', self.prim_pump)
|
||||||
|
|
||||||
|
sec_ff = steam_out / 2.0
|
||||||
|
level = s[f'COOLANT_SEC_{self.i}_LIQUID_VOLUME']
|
||||||
|
level_error = self.sec_level_target - level
|
||||||
|
sec_corr = self.sec_pid.step(level_error, dt)
|
||||||
|
sec_target = float(np.clip(sec_ff + sec_corr, 5.0, 100.0))
|
||||||
|
self.sec_pump += 0.3 * (sec_target - self.sec_pump)
|
||||||
|
set_param(f'COOLANT_SEC_CIRCULATION_PUMP_{self.i}_ORDERED_SPEED', self.sec_pump)
|
||||||
|
|
||||||
|
if s[f'STEAM_TURBINE_{self.i}_BYPASS_ACTUAL'] > 1.0:
|
||||||
|
set_param(f'STEAM_TURBINE_{self.i}_BYPASS_ORDERED', 0.0)
|
||||||
|
|
||||||
|
return power_kw, power_error, steam_out, level, level_error
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Global controller state
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
TEMP_MAX = 410.0
|
||||||
|
ROD_INTERVAL = 6
|
||||||
|
ROD_TIERS = [
|
||||||
|
(3.0, 0.1),
|
||||||
|
(8.0, 0.4),
|
||||||
|
(15.0, 0.8),
|
||||||
|
(float('inf'), 1.2),
|
||||||
|
]
|
||||||
|
rod_pos = float(nucon.get('ROD_BANK_POS_0_ACTUAL') or 85.0)
|
||||||
|
rod_cycle = 0
|
||||||
|
rod_integral = 0.0
|
||||||
|
|
||||||
|
train_controllers = {t: TrainController(t, targets[t]) for t in args.trains}
|
||||||
|
|
||||||
|
core_params = [
|
||||||
|
'CORE_TEMP', 'ROD_BANK_POS_0_ACTUAL',
|
||||||
|
'CORE_STATE_CRITICALITY',
|
||||||
|
'VACUUM_RETENTION_TANK_VOLUME',
|
||||||
|
'CONDENSER_VOLUME', 'CONDENSER_VAPOR_VOLUME',
|
||||||
|
'CONDENSER_VACUUM', # vacuum level % — monitor for pump health
|
||||||
|
'POWER_DEMAND_MW',
|
||||||
|
'CORE_PRIMARY_CIRCUIT_COOLING_TANK_VOLUME', # pressurizer water volume
|
||||||
|
'COOLANT_CORE_PRIMARY_LOOP_LEVEL', # overall primary loop fill %
|
||||||
|
'FREIGHT_PUMP_FEEDWATER_ACTIVE',
|
||||||
|
]
|
||||||
|
|
||||||
|
RETENTION_MAX = 40_000.0
|
||||||
|
RETENTION_HI = 0.75 * RETENTION_MAX
|
||||||
|
RETENTION_MID = 0.50 * RETENTION_MAX
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Pressurizer / primary circuit constants
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
PRSR_VALVE = 'Valvula_Pressurizer_Spray'
|
||||||
|
# CORE_PRIMARY_CIRCUIT_COOLING_TANK_VOLUME is the pressurizer water volume.
|
||||||
|
# Observed: 106030 = 60% → max ≈ 176717
|
||||||
|
PRSR_VOL_MAX = 176_717.0
|
||||||
|
PRSR_LEVEL_LO = 50.0 # % — open spray valve below this
|
||||||
|
PRSR_LEVEL_CLOSE = 60.0 # % — close spray valve once level recovers
|
||||||
|
PRSR_LEVEL_HI = 70.0 # % — op range high (informational)
|
||||||
|
PRIM_FILL_LO = 80.0 # % — start feedwater pump below this (uses COOLANT_CORE_PRIMARY_LOOP_LEVEL)
|
||||||
|
PRIM_FILL_HI = 90.0 # % — stop feedwater pump above this
|
||||||
|
|
||||||
|
# Initialise aux state from live game values so restarts are seamless.
|
||||||
|
_init = read_state([
|
||||||
|
'VACUUM_RETENTION_TANK_VOLUME',
|
||||||
|
'STEAM_EJECTOR_CONDENSER_RETURN_VALVE_ACTUAL',
|
||||||
|
'CONDENSER_VOLUME', 'CONDENSER_VAPOR_VOLUME',
|
||||||
|
'FREIGHT_PUMP_CONDENSER_ACTIVE',
|
||||||
|
'CONDENSER_VACUUM_PUMP_ACTIVE',
|
||||||
|
'CONDENSER_CIRCULATION_PUMP_ACTIVE',
|
||||||
|
])
|
||||||
|
_ret_vol_init = _init.get('VACUUM_RETENTION_TANK_VOLUME', 0.0)
|
||||||
|
_ret_valve_init = _init.get('STEAM_EJECTOR_CONDENSER_RETURN_VALVE_ACTUAL', 0.0)
|
||||||
|
ret_valve = _ret_valve_init
|
||||||
|
ret_draining = (_ret_valve_init > 0.5 and _ret_vol_init > RETENTION_MID)
|
||||||
|
if _ret_valve_init > 0.5 and not ret_draining:
|
||||||
|
set_param('STEAM_EJECTOR_CONDENSER_RETURN_VALVE', 0.0)
|
||||||
|
ret_valve = 0.0
|
||||||
|
ret_prev_vol = _ret_vol_init
|
||||||
|
|
||||||
|
_cond_vol_init = _init.get('CONDENSER_VOLUME', 0.0)
|
||||||
|
_cond_vap_init = _init.get('CONDENSER_VAPOR_VOLUME', 0.0)
|
||||||
|
_cond_tot_init = _cond_vol_init + _cond_vap_init
|
||||||
|
_cond_pct_init = (_cond_vol_init / _cond_tot_init * 100.0) if _cond_tot_init > 0 else 0.0
|
||||||
|
_cond_pump_init = bool(_init.get('FREIGHT_PUMP_CONDENSER_ACTIVE', False))
|
||||||
|
if _cond_pump_init and _cond_pct_init >= 60.0:
|
||||||
|
nucon.set(nucon._parameters['FREIGHT_PUMP_CONDENSER_SWITCH'], False)
|
||||||
|
cond_pump_on = False
|
||||||
|
elif not _cond_pump_init and _cond_pct_init < 45.0:
|
||||||
|
nucon.set(nucon._parameters['FREIGHT_PUMP_CONDENSER_SWITCH'], True)
|
||||||
|
cond_pump_on = True
|
||||||
|
else:
|
||||||
|
cond_pump_on = _cond_pump_init
|
||||||
|
|
||||||
|
# Vacuum pump — keep on continuously; turn off only during retention tank drain.
|
||||||
|
# (Opening the return valve breaks the suction path so the pump has no effect.)
|
||||||
|
vac_pump_on = bool(_init.get('CONDENSER_VACUUM_PUMP_ACTIVE', False))
|
||||||
|
if not vac_pump_on:
|
||||||
|
nucon.set(nucon._parameters['CONDENSER_VACUUM_PUMP_START_STOP'], True)
|
||||||
|
vac_pump_on = True
|
||||||
|
|
||||||
|
# Condenser circulation pump — run at moderate speed to prevent overcooling
|
||||||
|
# (manual §Stabilization: "prevent excessive cooling of the coolant returning to the evaporator").
|
||||||
|
_cond_circ_on = bool(_init.get('CONDENSER_CIRCULATION_PUMP_ACTIVE', False))
|
||||||
|
if not _cond_circ_on:
|
||||||
|
nucon.set(nucon._parameters['CONDENSER_CIRCULATION_PUMP_SWITCH'], True)
|
||||||
|
set_param('CONDENSER_CIRCULATION_PUMP_ORDERED_SPEED', 25.0)
|
||||||
|
|
||||||
|
# Pressurizer spray valve — init from live state
|
||||||
|
_prsr_live = read_state(['CORE_PRIMARY_CIRCUIT_COOLING_TANK_VOLUME', 'COOLANT_CORE_PRIMARY_LOOP_LEVEL', 'FREIGHT_PUMP_FEEDWATER_ACTIVE'])
|
||||||
|
_prsr_level = _prsr_live.get('CORE_PRIMARY_CIRCUIT_COOLING_TANK_VOLUME', PRSR_VOL_MAX * 0.6) / PRSR_VOL_MAX * 100.0
|
||||||
|
_prsr_valve = nucon.get_valve(PRSR_VALVE)
|
||||||
|
_prsr_open = _prsr_valve.get('IsOpened', False) or _prsr_valve.get('Value', 0) > 50
|
||||||
|
prsr_spraying = _prsr_open and _prsr_level < PRSR_LEVEL_CLOSE
|
||||||
|
if _prsr_open and not prsr_spraying:
|
||||||
|
nucon.close_valve(PRSR_VALVE)
|
||||||
|
|
||||||
|
feedwater_on = bool(_prsr_live.get('FREIGHT_PUMP_FEEDWATER_ACTIVE', False))
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# TUI helpers
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
def _bar(pct, width=18):
|
||||||
|
pct = max(0.0, min(100.0, pct))
|
||||||
|
filled = int(pct / 100.0 * width)
|
||||||
|
return '█' * filled + '░' * (width - filled)
|
||||||
|
|
||||||
|
def _safe_addstr(scr, row, col, text, attr=0):
|
||||||
|
H, W = scr.getmaxyx()
|
||||||
|
if row < 0 or row >= H:
|
||||||
|
return
|
||||||
|
if col < 0:
|
||||||
|
text = text[-col:]
|
||||||
|
col = 0
|
||||||
|
if col >= W:
|
||||||
|
return
|
||||||
|
text = text[:W - col]
|
||||||
|
try:
|
||||||
|
scr.addstr(row, col, text, attr)
|
||||||
|
except curses.error:
|
||||||
|
pass
|
||||||
|
|
||||||
|
def _hline(scr, row, char='─'):
|
||||||
|
H, W = scr.getmaxyx()
|
||||||
|
if 0 <= row < H:
|
||||||
|
_safe_addstr(scr, row, 0, char * (W - 1))
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Main TUI loop
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
def run_controller(stdscr):
|
||||||
|
global rod_pos, rod_cycle, rod_integral
|
||||||
|
global ret_valve, ret_draining, ret_prev_vol, cond_pump_on
|
||||||
|
global prsr_spraying, feedwater_on
|
||||||
|
global vac_pump_on
|
||||||
|
global train_controllers, targets
|
||||||
|
|
||||||
|
curses.curs_set(0)
|
||||||
|
stdscr.nodelay(True)
|
||||||
|
curses.start_color()
|
||||||
|
curses.use_default_colors()
|
||||||
|
curses.init_pair(1, curses.COLOR_GREEN, -1) # good / normal
|
||||||
|
curses.init_pair(2, curses.COLOR_YELLOW, -1) # warning
|
||||||
|
curses.init_pair(3, curses.COLOR_RED, -1) # alarm
|
||||||
|
curses.init_pair(4, curses.COLOR_CYAN, -1) # selected
|
||||||
|
curses.init_pair(5, curses.COLOR_WHITE, curses.COLOR_BLUE) # title bar
|
||||||
|
|
||||||
|
GREEN = curses.color_pair(1)
|
||||||
|
YELLOW = curses.color_pair(2)
|
||||||
|
RED = curses.color_pair(3)
|
||||||
|
CYAN = curses.color_pair(4)
|
||||||
|
TITLE = curses.color_pair(5)
|
||||||
|
BOLD = curses.A_BOLD
|
||||||
|
REV = curses.A_REVERSE
|
||||||
|
|
||||||
|
SELECTION_ORDER = [0, 1, 2, 3, 4] # 0=core 1/2/3=trains 4=grid
|
||||||
|
selected_train = args.trains[0] if args.trains else 1
|
||||||
|
temp_setpoint = args.temp_setpoint # mutable; adjustable from TUI
|
||||||
|
temp_auto = True # auto-adjust setpoint to meet total power demand
|
||||||
|
grid_follow = args.grid_follow
|
||||||
|
# Per-train power caps: manual max each train should carry (used for proportional distribution)
|
||||||
|
grid_caps = {t: tc.target_kw for t, tc in train_controllers.items()}
|
||||||
|
cycle = 0
|
||||||
|
train_data = {}
|
||||||
|
# display state — updated each control cycle, read by draw() at any time
|
||||||
|
disp = dict(s={}, dynamic_setpoint=temp_setpoint, temp_auto=temp_auto, criticality=0.0,
|
||||||
|
ret_pct=0.0, ret_draining=False, ret_valve=0.0,
|
||||||
|
cond_pct=0.0, cond_pump_on=False,
|
||||||
|
vac_pump_on=vac_pump_on,
|
||||||
|
prsr_level=_prsr_level, prsr_spraying=prsr_spraying,
|
||||||
|
prim_level=_prsr_live.get('COOLANT_CORE_PRIMARY_LOOP_LEVEL', 100.0),
|
||||||
|
feedwater_on=feedwater_on,
|
||||||
|
grid_follow=grid_follow, grid_demand_kw=0.0)
|
||||||
|
|
||||||
|
def rebuild_all_params():
|
||||||
|
p = list(core_params)
|
||||||
|
for tc in train_controllers.values():
|
||||||
|
p += tc.params()
|
||||||
|
return p
|
||||||
|
|
||||||
|
all_params = rebuild_all_params()
|
||||||
|
|
||||||
|
def handle_key(key):
|
||||||
|
nonlocal selected_train, all_params, temp_setpoint, temp_auto, grid_follow, grid_caps
|
||||||
|
if key in (ord('q'), 27):
|
||||||
|
return True # signal quit
|
||||||
|
# Direct selection by number
|
||||||
|
elif key in (ord('0'), ord('1'), ord('2'), ord('3')):
|
||||||
|
selected_train = key - ord('0')
|
||||||
|
elif key == ord('g'):
|
||||||
|
selected_train = 4
|
||||||
|
# Up/Down cycle through selections
|
||||||
|
elif key == curses.KEY_UP:
|
||||||
|
idx = SELECTION_ORDER.index(selected_train) if selected_train in SELECTION_ORDER else 0
|
||||||
|
selected_train = SELECTION_ORDER[(idx - 1) % len(SELECTION_ORDER)]
|
||||||
|
elif key == curses.KEY_DOWN:
|
||||||
|
idx = SELECTION_ORDER.index(selected_train) if selected_train in SELECTION_ORDER else 0
|
||||||
|
selected_train = SELECTION_ORDER[(idx + 1) % len(SELECTION_ORDER)]
|
||||||
|
# Right/+ increase Left/- decrease
|
||||||
|
elif key in (ord('+'), ord('='), curses.KEY_RIGHT):
|
||||||
|
if selected_train == 0 and not temp_auto:
|
||||||
|
temp_setpoint = min(round(temp_setpoint / 5.0) * 5.0 + 5.0, 375.0)
|
||||||
|
elif selected_train == 4:
|
||||||
|
args.grid_buffer = min(args.grid_buffer + 1.0, 100.0)
|
||||||
|
elif selected_train in train_controllers:
|
||||||
|
if grid_follow:
|
||||||
|
cur = grid_caps.get(selected_train, train_controllers[selected_train].target_kw)
|
||||||
|
grid_caps[selected_train] = min(round(cur / 5_000) * 5_000 + 5_000, 100_000) # snap+step
|
||||||
|
else:
|
||||||
|
tc = train_controllers[selected_train]
|
||||||
|
tc.target_kw = min(round(tc.target_kw / 5_000) * 5_000 + 5_000, 100_000)
|
||||||
|
targets[selected_train] = tc.target_kw
|
||||||
|
grid_caps[selected_train] = tc.target_kw
|
||||||
|
elif selected_train in (1, 2, 3):
|
||||||
|
targets[selected_train] = 5_000
|
||||||
|
train_controllers[selected_train] = TrainController(selected_train, 5_000)
|
||||||
|
grid_caps[selected_train] = 5_000
|
||||||
|
all_params = rebuild_all_params()
|
||||||
|
elif key in (ord('-'), curses.KEY_LEFT):
|
||||||
|
if selected_train == 0 and not temp_auto:
|
||||||
|
temp_setpoint = max(round(temp_setpoint / 5.0) * 5.0 - 5.0, 250.0)
|
||||||
|
elif selected_train == 4:
|
||||||
|
args.grid_buffer = max(args.grid_buffer - 1.0, 0.0)
|
||||||
|
elif selected_train in train_controllers:
|
||||||
|
if grid_follow:
|
||||||
|
cur = grid_caps.get(selected_train, train_controllers[selected_train].target_kw)
|
||||||
|
grid_caps[selected_train] = max(round(cur / 5_000) * 5_000 - 5_000, 0)
|
||||||
|
else:
|
||||||
|
tc = train_controllers[selected_train]
|
||||||
|
tc.target_kw = max(round(tc.target_kw / 5_000) * 5_000 - 5_000, 0)
|
||||||
|
targets[selected_train] = tc.target_kw
|
||||||
|
grid_caps[selected_train] = tc.target_kw
|
||||||
|
elif key == ord('d'):
|
||||||
|
if selected_train == 0:
|
||||||
|
temp_auto = not temp_auto
|
||||||
|
disp['temp_auto'] = temp_auto
|
||||||
|
elif selected_train == 4:
|
||||||
|
grid_follow = not grid_follow
|
||||||
|
disp['grid_follow'] = grid_follow
|
||||||
|
elif selected_train in train_controllers:
|
||||||
|
del train_controllers[selected_train]
|
||||||
|
grid_caps.pop(selected_train, None)
|
||||||
|
if selected_train in targets:
|
||||||
|
del targets[selected_train]
|
||||||
|
all_params = rebuild_all_params()
|
||||||
|
selected_train = 0 if not train_controllers else list(train_controllers.keys())[0]
|
||||||
|
return False
|
||||||
|
|
||||||
|
def draw():
|
||||||
|
s = disp['s']
|
||||||
|
dynamic_setpoint = disp['dynamic_setpoint']
|
||||||
|
criticality = disp['criticality']
|
||||||
|
ret_pct = disp['ret_pct']
|
||||||
|
ret_draining = disp['ret_draining']
|
||||||
|
ret_valve = disp['ret_valve']
|
||||||
|
cond_pct = disp['cond_pct']
|
||||||
|
cond_pump_on = disp['cond_pump_on']
|
||||||
|
if not s:
|
||||||
|
return
|
||||||
|
stdscr.erase()
|
||||||
|
H, W = stdscr.getmaxyx()
|
||||||
|
row = 0
|
||||||
|
|
||||||
|
title = f" NUCLEARES CONTROLLER ─ Cycle {cycle:5d} ─ dt={args.dt:.0f}s "
|
||||||
|
_safe_addstr(stdscr, row, 0, title.ljust(W - 1), TITLE | BOLD)
|
||||||
|
row += 1
|
||||||
|
|
||||||
|
_hline(stdscr, row); row += 1
|
||||||
|
core_sel = (selected_train == 0)
|
||||||
|
core_attr = CYAN | BOLD if core_sel else BOLD
|
||||||
|
temp_auto_ = disp['temp_auto']
|
||||||
|
core_temp = s.get('CORE_TEMP', 0.0)
|
||||||
|
temp_color = RED if core_temp > 370 else YELLOW if core_temp > 355 else GREEN
|
||||||
|
scram_str = ' !! SCRAM !!' if core_temp > TEMP_MAX else ''
|
||||||
|
auto_str = 'AUTO' if temp_auto_ else 'MAN '
|
||||||
|
auto_color = GREEN if temp_auto_ else YELLOW
|
||||||
|
_safe_addstr(stdscr, row, 2, '◆ CORE' + (' ◀' if core_sel else ''), core_attr)
|
||||||
|
_safe_addstr(stdscr, row, 10, f'[{auto_str}]', auto_color | BOLD)
|
||||||
|
_safe_addstr(stdscr, row, 16, 'Temp: ', BOLD)
|
||||||
|
_safe_addstr(stdscr, row, 22, f'{core_temp:6.1f}°C', temp_color | BOLD)
|
||||||
|
sp_color = RED if dynamic_setpoint < 306 or dynamic_setpoint > 375 else 0
|
||||||
|
_safe_addstr(stdscr, row, 32, f'sp=', 0)
|
||||||
|
_safe_addstr(stdscr, row, 35, f'{dynamic_setpoint:.0f}°C', sp_color | BOLD)
|
||||||
|
_safe_addstr(stdscr, row, 40,
|
||||||
|
f' Rod: {s.get("ROD_BANK_POS_0_ACTUAL", 0):5.1f} '
|
||||||
|
f'Crit: {criticality:+.3f}{scram_str}')
|
||||||
|
row += 1
|
||||||
|
|
||||||
|
for t in (1, 2, 3):
|
||||||
|
_hline(stdscr, row); row += 1
|
||||||
|
is_sel = (t == selected_train)
|
||||||
|
is_active = (t in train_controllers)
|
||||||
|
tc = train_controllers.get(t)
|
||||||
|
sel_attr = CYAN | BOLD if is_sel else 0
|
||||||
|
label = f'◆ TRAIN {t}' + (' ◀' if is_sel else '')
|
||||||
|
_safe_addstr(stdscr, row, 2, label, sel_attr | BOLD)
|
||||||
|
if is_active and t in train_data:
|
||||||
|
power_kw, power_error, steam_out, level, level_error = train_data[t]
|
||||||
|
pwr_pct = power_kw / tc.target_kw * 100.0 if tc.target_kw > 0 else 0.0
|
||||||
|
pwr_color = GREEN if abs(power_error) < 2000 else YELLOW if abs(power_error) < 8000 else RED
|
||||||
|
cap = grid_caps.get(t, tc.target_kw)
|
||||||
|
gf = disp['grid_follow']
|
||||||
|
tgt_str = (f'tgt={tc.target_kw/1000:.1f}/{cap/1000:.0f}MW'
|
||||||
|
if gf and abs(tc.target_kw - cap) > 500
|
||||||
|
else f'tgt={tc.target_kw/1000:.0f}MW')
|
||||||
|
_safe_addstr(stdscr, row, 16, 'Power: ', BOLD)
|
||||||
|
_safe_addstr(stdscr, row, 23, f'{power_kw/1000:5.1f} MW', pwr_color | BOLD)
|
||||||
|
_safe_addstr(stdscr, row, 32,
|
||||||
|
f'[{_bar(pwr_pct, 14)}] {power_error/1000:+5.1f}MW {tgt_str}')
|
||||||
|
row += 1
|
||||||
|
prim_warn = abs(tc.prim_pump - tc.PRIM_PUMP_SUGGESTED) > 10
|
||||||
|
prim_attr = YELLOW if prim_warn else 0
|
||||||
|
prim_str = f'{tc.prim_pump:3.0f}%{"!" if prim_warn else " "}'
|
||||||
|
_safe_addstr(stdscr, row, 16,
|
||||||
|
f'Steam: {steam_out:5.1f} MSCV: {tc.mscv:4.1f} Prim: ')
|
||||||
|
_safe_addstr(stdscr, row, 51, prim_str, prim_attr)
|
||||||
|
_safe_addstr(stdscr, row, 56,
|
||||||
|
f' Sec: {tc.sec_pump:3.0f}% Lvl: {level:.0f} (Δ{level_error:+.0f})')
|
||||||
|
elif not is_active:
|
||||||
|
hint = ' (+/Up to add)' if is_sel else ''
|
||||||
|
_safe_addstr(stdscr, row, 16, f'not controlled{hint}',
|
||||||
|
YELLOW if is_sel else 0)
|
||||||
|
row += 1
|
||||||
|
|
||||||
|
_hline(stdscr, row); row += 1
|
||||||
|
gf = disp['grid_follow']
|
||||||
|
gdkw = disp['grid_demand_kw']
|
||||||
|
total_cap = sum(grid_caps.get(t, tc.target_kw) for t, tc in train_controllers.items())
|
||||||
|
grid_sel = (selected_train == 4)
|
||||||
|
grid_attr = CYAN | BOLD if grid_sel else BOLD
|
||||||
|
gf_color = GREEN | BOLD if gf else (CYAN | BOLD if grid_sel else 0)
|
||||||
|
_safe_addstr(stdscr, row, 2, '◆ GRID' + (' ◀' if grid_sel else ''), grid_attr)
|
||||||
|
_safe_addstr(stdscr, row, 16, f'Demand: {gdkw/1000:5.1f} MW', BOLD)
|
||||||
|
if gf:
|
||||||
|
target_total = gdkw + args.grid_buffer * 1000.0
|
||||||
|
_safe_addstr(stdscr, row, 34,
|
||||||
|
f' AUTO buf={args.grid_buffer:.0f}MW '
|
||||||
|
f'→{target_total/1000:.1f}/{total_cap/1000:.0f}MW total', gf_color)
|
||||||
|
else:
|
||||||
|
_safe_addstr(stdscr, row, 34,
|
||||||
|
f' off buf={args.grid_buffer:.0f}MW cap={total_cap/1000:.0f}MW', gf_color)
|
||||||
|
row += 1
|
||||||
|
|
||||||
|
_hline(stdscr, row); row += 1
|
||||||
|
ret_color = RED if ret_pct > 75 else YELLOW if ret_pct > 60 else GREEN
|
||||||
|
_safe_addstr(stdscr, row, 2, '◆ RETENTION TANK ', BOLD)
|
||||||
|
_safe_addstr(stdscr, row, 20, f'[{_bar(ret_pct, 20)}]', ret_color)
|
||||||
|
_safe_addstr(stdscr, row, 43, f' {ret_pct:4.0f}%')
|
||||||
|
_safe_addstr(stdscr, row, 49,
|
||||||
|
f' DRAINING valve={ret_valve:.0f}%' if ret_draining else ' OK',
|
||||||
|
YELLOW if ret_draining else GREEN)
|
||||||
|
row += 1
|
||||||
|
cond_vac_ = s.get('CONDENSER_VACUUM', 0.0)
|
||||||
|
cond_color = RED if cond_pct < 25 else YELLOW if cond_pct < 40 else GREEN
|
||||||
|
vac_on_ = disp.get('vac_pump_on', True)
|
||||||
|
vac_color = (RED if cond_vac_ < 50 else YELLOW if cond_vac_ < 80 else GREEN) if vac_on_ else YELLOW
|
||||||
|
_safe_addstr(stdscr, row, 2, '◆ CONDENSER FILL ', BOLD)
|
||||||
|
_safe_addstr(stdscr, row, 20, f'[{_bar(cond_pct, 20)}]', cond_color)
|
||||||
|
_safe_addstr(stdscr, row, 43, f' {cond_pct:4.0f}%')
|
||||||
|
_safe_addstr(stdscr, row, 49, ' PUMP ON' if cond_pump_on else ' OK',
|
||||||
|
YELLOW if cond_pump_on else GREEN)
|
||||||
|
_safe_addstr(stdscr, row, 60,
|
||||||
|
f' VAC:{"OFF" if not vac_on_ else f"{cond_vac_:.0f}%"}',
|
||||||
|
vac_color)
|
||||||
|
row += 1
|
||||||
|
prsr_level_ = disp['prsr_level']
|
||||||
|
prsr_spray_ = disp['prsr_spraying']
|
||||||
|
feedwater_ = disp['feedwater_on']
|
||||||
|
prsr_color = RED if prsr_level_ < 40 or prsr_level_ > 80 else YELLOW if prsr_level_ < PRSR_LEVEL_LO or prsr_level_ > PRSR_LEVEL_HI else GREEN
|
||||||
|
_safe_addstr(stdscr, row, 2, '◆ PRESSURIZER ', BOLD)
|
||||||
|
_safe_addstr(stdscr, row, 20, f'[{_bar(prsr_level_, 20)}]', prsr_color)
|
||||||
|
_safe_addstr(stdscr, row, 43, f' {prsr_level_:4.1f}%')
|
||||||
|
_safe_addstr(stdscr, row, 49, ' SPRAY ON' if prsr_spray_ else ' OK',
|
||||||
|
YELLOW if prsr_spray_ else GREEN)
|
||||||
|
row += 1
|
||||||
|
prim_level_ = disp.get('prim_level', 100.0)
|
||||||
|
prim_color = RED if prim_level_ < 70 else YELLOW if prim_level_ < PRIM_FILL_LO else GREEN
|
||||||
|
_safe_addstr(stdscr, row, 2, '◆ PRIMARY VESSEL ', BOLD)
|
||||||
|
_safe_addstr(stdscr, row, 20, f'[{_bar(prim_level_, 20)}]', prim_color)
|
||||||
|
_safe_addstr(stdscr, row, 43, f' {prim_level_:4.1f}%')
|
||||||
|
_safe_addstr(stdscr, row, 49, ' FW PUMP ON' if feedwater_ else ' OK',
|
||||||
|
YELLOW if feedwater_ else GREEN)
|
||||||
|
row += 1
|
||||||
|
|
||||||
|
if selected_train == 0:
|
||||||
|
adj_hint = f'←/→ sp {disp["dynamic_setpoint"]:.0f}°C±5' if not disp['temp_auto'] else f'sp={disp["dynamic_setpoint"]:.0f}°C (auto)'
|
||||||
|
d_hint = f' [d] auto {"OFF" if disp["temp_auto"] else "ON"}'
|
||||||
|
elif selected_train == 4:
|
||||||
|
adj_hint = f'←/→ buf {args.grid_buffer:.0f}MW±1'
|
||||||
|
d_hint = ' [d] toggle auto'
|
||||||
|
elif disp['grid_follow']:
|
||||||
|
cap = grid_caps.get(selected_train, 0)
|
||||||
|
adj_hint = f'←/→ max {cap/1000:.0f}MW±5'
|
||||||
|
d_hint = ' [d] remove'
|
||||||
|
else:
|
||||||
|
adj_hint = '←/→ target ±5MW'
|
||||||
|
d_hint = ' [d] remove'
|
||||||
|
_safe_addstr(stdscr, H - 1, 0,
|
||||||
|
f' [↑↓] select [0-3/g] jump {adj_hint}{d_hint} [q] quit '.ljust(W - 1),
|
||||||
|
REV)
|
||||||
|
stdscr.refresh()
|
||||||
|
|
||||||
|
while True:
|
||||||
|
t0 = time.time()
|
||||||
|
s = read_state(all_params)
|
||||||
|
cycle += 1
|
||||||
|
# ---- Rod control ----
|
||||||
|
temp_error = s['CORE_TEMP'] - temp_setpoint
|
||||||
|
criticality = s.get('CORE_STATE_CRITICALITY', 0.0)
|
||||||
|
rod_cycle += 1
|
||||||
|
if s['CORE_TEMP'] > TEMP_MAX:
|
||||||
|
rod_pos = 100.0
|
||||||
|
for tc in train_controllers.values():
|
||||||
|
tc.prim_pump = 90.0
|
||||||
|
set_param(f'COOLANT_CORE_CIRCULATION_PUMP_{tc.i}_ORDERED_SPEED', 90.0)
|
||||||
|
else:
|
||||||
|
urgent = temp_error > 5.0 or criticality > 0.3
|
||||||
|
if urgent or rod_cycle >= ROD_INTERVAL:
|
||||||
|
if rod_cycle >= ROD_INTERVAL:
|
||||||
|
rod_cycle = 0
|
||||||
|
abs_err = abs(temp_error)
|
||||||
|
max_step = next(lim for thresh, lim in ROD_TIERS if abs_err <= thresh)
|
||||||
|
if urgent and rod_cycle != 0:
|
||||||
|
max_step = min(max_step, 0.25)
|
||||||
|
if not urgent:
|
||||||
|
rod_integral = float(np.clip(rod_integral + 0.002 * temp_error, -3.0, 3.0))
|
||||||
|
else:
|
||||||
|
rod_integral *= 0.5
|
||||||
|
raw_delta = 0.04 * temp_error + 1.0 * criticality + rod_integral
|
||||||
|
rod_delta = float(np.clip(raw_delta, -max_step, max_step))
|
||||||
|
rod_pos = float(np.clip(s['ROD_BANK_POS_0_ACTUAL'] + rod_delta, 0.0, 100.0))
|
||||||
|
set_param('ROD_BANK_POS_0_ORDERED', rod_pos)
|
||||||
|
|
||||||
|
# ---- Grid-demand following ----
|
||||||
|
grid_demand_kw = s.get('POWER_DEMAND_MW', 0.0) * 1000.0
|
||||||
|
if grid_follow and train_controllers:
|
||||||
|
total_target_kw = grid_demand_kw + args.grid_buffer * 1000.0
|
||||||
|
# Distribute proportionally to each train's manual cap; never exceed cap
|
||||||
|
total_cap = sum(grid_caps.get(t, tc.target_kw) for t, tc in train_controllers.items())
|
||||||
|
if total_cap > 0:
|
||||||
|
for t, tc in train_controllers.items():
|
||||||
|
cap = grid_caps.get(t, tc.target_kw)
|
||||||
|
share = total_target_kw * (cap / total_cap)
|
||||||
|
tc.target_kw = float(np.clip(share, 0.0, cap))
|
||||||
|
|
||||||
|
# ---- Per-train control ----
|
||||||
|
for t, tc in train_controllers.items():
|
||||||
|
res = tc.step(s, args.dt)
|
||||||
|
train_data[t] = res
|
||||||
|
|
||||||
|
# ---- Auto temp setpoint ----
|
||||||
|
if temp_auto and train_data:
|
||||||
|
total_error = sum(train_data[t][1] for t in train_data) # sum of power_errors
|
||||||
|
sp_delta = float(np.clip(total_error * 0.00002, -0.5, 0.5))
|
||||||
|
temp_setpoint = float(np.clip(temp_setpoint + sp_delta, 306.0, 375.0))
|
||||||
|
|
||||||
|
# ---- Aux: retention tank ----
|
||||||
|
ret_vol = s.get('VACUUM_RETENTION_TANK_VOLUME', 0.0)
|
||||||
|
ret_pct = ret_vol / RETENTION_MAX * 100.0
|
||||||
|
if ret_draining and ret_vol <= RETENTION_MID:
|
||||||
|
ret_draining = False
|
||||||
|
ret_valve = 0.0
|
||||||
|
set_param('STEAM_EJECTOR_CONDENSER_RETURN_VALVE', 0.0)
|
||||||
|
# Drain complete — restart vacuum pump
|
||||||
|
if not vac_pump_on:
|
||||||
|
nucon.set(nucon._parameters['CONDENSER_VACUUM_PUMP_START_STOP'], True)
|
||||||
|
vac_pump_on = True
|
||||||
|
elif ret_vol > RETENTION_HI:
|
||||||
|
if not ret_draining:
|
||||||
|
# Starting drain — stop vacuum pump.
|
||||||
|
# The ejector return valve bypasses the suction path so the pump has no effect
|
||||||
|
# and wastes power; turn it off for the duration of the drain.
|
||||||
|
nucon.set(nucon._parameters['CONDENSER_VACUUM_PUMP_START_STOP'], False)
|
||||||
|
vac_pump_on = False
|
||||||
|
ret_draining = True
|
||||||
|
if ret_prev_vol is not None and ret_vol >= ret_prev_vol - 50.0:
|
||||||
|
ret_valve = min(ret_valve + 1.0, 50.0)
|
||||||
|
set_param('STEAM_EJECTOR_CONDENSER_RETURN_VALVE', ret_valve)
|
||||||
|
elif ret_draining:
|
||||||
|
set_param('STEAM_EJECTOR_CONDENSER_RETURN_VALVE', ret_valve)
|
||||||
|
ret_prev_vol = ret_vol
|
||||||
|
|
||||||
|
# ---- Aux: condenser fill ----
|
||||||
|
cond_vol = s.get('CONDENSER_VOLUME', 0.0)
|
||||||
|
cond_vap = s.get('CONDENSER_VAPOR_VOLUME', 0.0)
|
||||||
|
cond_tot = cond_vol + cond_vap
|
||||||
|
cond_pct = (cond_vol / cond_tot * 100.0) if cond_tot > 0 else 0.0
|
||||||
|
if not cond_pump_on and cond_pct < 45.0:
|
||||||
|
cond_pump_on = True
|
||||||
|
nucon.set(nucon._parameters['FREIGHT_PUMP_CONDENSER_SWITCH'], True)
|
||||||
|
elif cond_pump_on and cond_pct >= 60.0:
|
||||||
|
cond_pump_on = False
|
||||||
|
nucon.set(nucon._parameters['FREIGHT_PUMP_CONDENSER_SWITCH'], False)
|
||||||
|
|
||||||
|
# ---- Aux: pressurizer spray valve (level 50-70%) ----
|
||||||
|
prsr_level = s.get('CORE_PRIMARY_CIRCUIT_COOLING_TANK_VOLUME', PRSR_VOL_MAX * 0.6) / PRSR_VOL_MAX * 100.0
|
||||||
|
if not prsr_spraying and prsr_level < PRSR_LEVEL_LO:
|
||||||
|
prsr_spraying = True
|
||||||
|
nucon.open_valve(PRSR_VALVE)
|
||||||
|
elif prsr_spraying and prsr_level >= PRSR_LEVEL_CLOSE:
|
||||||
|
prsr_spraying = False
|
||||||
|
nucon.close_valve(PRSR_VALVE)
|
||||||
|
elif not prsr_spraying:
|
||||||
|
# Valve should be at rest — power off actuator if it's reached closed position
|
||||||
|
_vs = nucon.get_valve(PRSR_VALVE)
|
||||||
|
if _vs.get('IsClosed') and _vs.get('Actuator') != 'OFF':
|
||||||
|
nucon.off_valve(PRSR_VALVE)
|
||||||
|
|
||||||
|
# ---- Aux: primary circuit feedwater (overall loop fill > 80%) ----
|
||||||
|
prim_level = s.get('COOLANT_CORE_PRIMARY_LOOP_LEVEL', 100.0)
|
||||||
|
if not feedwater_on and prim_level < PRIM_FILL_LO:
|
||||||
|
feedwater_on = True
|
||||||
|
nucon.set(nucon._parameters['FREIGHT_PUMP_FEEDWATER_SWITCH'], True)
|
||||||
|
elif feedwater_on and prim_level >= PRIM_FILL_HI:
|
||||||
|
feedwater_on = False
|
||||||
|
nucon.set(nucon._parameters['FREIGHT_PUMP_FEEDWATER_SWITCH'], False)
|
||||||
|
|
||||||
|
# ---- Update display state and redraw ----
|
||||||
|
disp.update(s=s, dynamic_setpoint=temp_setpoint, temp_auto=temp_auto,
|
||||||
|
criticality=criticality,
|
||||||
|
ret_pct=ret_pct, ret_draining=ret_draining, ret_valve=ret_valve,
|
||||||
|
cond_pct=cond_pct, cond_pump_on=cond_pump_on,
|
||||||
|
vac_pump_on=vac_pump_on,
|
||||||
|
prsr_level=prsr_level, prsr_spraying=prsr_spraying,
|
||||||
|
prim_level=prim_level, feedwater_on=feedwater_on,
|
||||||
|
grid_follow=grid_follow, grid_demand_kw=grid_demand_kw)
|
||||||
|
draw()
|
||||||
|
|
||||||
|
# ---- Poll input + redraw at 50 ms intervals for the rest of the cycle ----
|
||||||
|
sim_speed = nucon.GAME_SIM_SPEED.value or 1.0
|
||||||
|
deadline = t0 + args.dt / sim_speed
|
||||||
|
stdscr.timeout(50)
|
||||||
|
while time.time() < deadline:
|
||||||
|
key = stdscr.getch()
|
||||||
|
if key == -1:
|
||||||
|
continue
|
||||||
|
if handle_key(key):
|
||||||
|
return
|
||||||
|
disp['dynamic_setpoint'] = temp_setpoint
|
||||||
|
disp['temp_auto'] = temp_auto
|
||||||
|
disp['grid_follow'] = grid_follow
|
||||||
|
draw()
|
||||||
|
stdscr.timeout(-1)
|
||||||
|
|
||||||
|
curses.wrapper(run_controller)
|
||||||
@@ -0,0 +1,151 @@
|
|||||||
|
"""SAC + HER training on kNN-GP simulator.
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
python3.14 train_sac.py
|
||||||
|
python3.14 train_sac.py --load /tmp/sac_nucon_knn # hot-start from previous run
|
||||||
|
|
||||||
|
Requirements:
|
||||||
|
- NuCon game running (for parameter metadata)
|
||||||
|
- /tmp/reactor_knn.pkl (kNN-GP model)
|
||||||
|
- /tmp/nucon_dataset.pkl (500-sample dataset for init_states)
|
||||||
|
"""
|
||||||
|
import argparse
|
||||||
|
import pickle
|
||||||
|
import torch
|
||||||
|
from gymnasium.wrappers import TimeLimit
|
||||||
|
from stable_baselines3 import SAC
|
||||||
|
from stable_baselines3.her.her_replay_buffer import HerReplayBuffer
|
||||||
|
from stable_baselines3.common.callbacks import CheckpointCallback
|
||||||
|
|
||||||
|
from nucon.sim import NuconSimulator
|
||||||
|
from nucon.model import ReactorDynamicsModel, MixtureModel
|
||||||
|
from nucon.rl import NuconGoalEnv, Parameterized_Objectives, Parameterized_Terminators
|
||||||
|
|
||||||
|
parser = argparse.ArgumentParser()
|
||||||
|
parser.add_argument('--load', default=None, help='Path to existing model to hot-start from')
|
||||||
|
parser.add_argument('--steps', type=int, default=50_000, help='Total timesteps (default: 50000)')
|
||||||
|
parser.add_argument('--out', default='/tmp/sac_nucon_knn', help='Output path for saved model')
|
||||||
|
parser.add_argument('--model', default='/tmp/reactor_knn.pkl', help='Dynamics model (.pkl for kNN, .pt for NN)')
|
||||||
|
parser.add_argument('--model2', default=None, help='Second dynamics model for mixture (optional)')
|
||||||
|
parser.add_argument('--dataset', default='/tmp/nucon_dataset.pkl', help='Dataset for init states')
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Load dynamics model(s) and dataset
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
def _load_model(path):
|
||||||
|
if path.endswith('.pt'):
|
||||||
|
ckpt = torch.load(path, weights_only=False)
|
||||||
|
m = ReactorDynamicsModel(ckpt['input_params'], ckpt['output_params'])
|
||||||
|
m.load_state_dict(ckpt['state_dict'])
|
||||||
|
m.eval()
|
||||||
|
return m
|
||||||
|
with open(path, 'rb') as f:
|
||||||
|
return pickle.load(f)
|
||||||
|
|
||||||
|
dynamics_model = _load_model(args.model)
|
||||||
|
if args.model2:
|
||||||
|
dynamics_model = MixtureModel(dynamics_model, _load_model(args.model2))
|
||||||
|
|
||||||
|
with open(args.dataset, 'rb') as f:
|
||||||
|
dataset = pickle.load(f)
|
||||||
|
|
||||||
|
# Seed resets to in-distribution states from dataset
|
||||||
|
init_states = [s for _, _, s, _ in dataset]
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Build sim + env
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
sim = NuconSimulator(port=8786)
|
||||||
|
sim.set_model(dynamics_model)
|
||||||
|
|
||||||
|
BATCH_SIZE = 2048
|
||||||
|
MAX_EPISODE_STEPS = 200
|
||||||
|
|
||||||
|
GENERATORS = ['GENERATOR_0_KW', 'GENERATOR_1_KW', 'GENERATOR_2_KW']
|
||||||
|
POWER_RANGE = {g: (0.0, 100_000.0) for g in GENERATORS} # per-generator kW; ~100 MW upper bound
|
||||||
|
|
||||||
|
# Curated obs: physically relevant features for power control (~25 dims vs ~260 full)
|
||||||
|
OBS_PARAMS = [
|
||||||
|
'CORE_TEMP', 'CORE_PRESSURE', 'CORE_STATE_CRITICALITY', 'CORE_WEAR', 'CORE_INTEGRITY',
|
||||||
|
'ROD_BANK_POS_0_ACTUAL', 'ROD_BANK_POS_0_ORDERED',
|
||||||
|
'COOLANT_CORE_FLOW_SPEED', 'COOLANT_CORE_VESSEL_TEMPERATURE',
|
||||||
|
'COOLANT_CORE_PRESSURE', 'COOLANT_CORE_QUANTITY_IN_VESSEL',
|
||||||
|
'STEAM_TURBINE_0_RPM', 'STEAM_TURBINE_0_TEMPERATURE', 'STEAM_TURBINE_0_PRESSURE',
|
||||||
|
'STEAM_TURBINE_1_RPM', 'STEAM_TURBINE_1_TEMPERATURE', 'STEAM_TURBINE_1_PRESSURE',
|
||||||
|
'STEAM_TURBINE_2_RPM', 'STEAM_TURBINE_2_TEMPERATURE', 'STEAM_TURBINE_2_PRESSURE',
|
||||||
|
'GENERATOR_0_V', 'GENERATOR_1_V', 'GENERATOR_2_V',
|
||||||
|
]
|
||||||
|
|
||||||
|
env = NuconGoalEnv(
|
||||||
|
goal_params=GENERATORS,
|
||||||
|
goal_range=POWER_RANGE,
|
||||||
|
seconds_per_step=10,
|
||||||
|
simulator=sim,
|
||||||
|
obs_params=OBS_PARAMS,
|
||||||
|
additional_objectives=[
|
||||||
|
Parameterized_Objectives['uncertainty_penalty'](start=0.3),
|
||||||
|
Parameterized_Objectives['temp_below_linear'](max_temp=420),
|
||||||
|
],
|
||||||
|
additional_objective_weights=[1.0, 0.01],
|
||||||
|
init_states=init_states,
|
||||||
|
delta_action_scale=0.05,
|
||||||
|
goal_sampling_std=0.15, # Gaussian delta in normalised space (~180 kW typical)
|
||||||
|
)
|
||||||
|
|
||||||
|
env = TimeLimit(env, max_episode_steps=MAX_EPISODE_STEPS)
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# SAC + HER
|
||||||
|
# learning_starts = batch_size: wait for batch_size complete (short) episodes
|
||||||
|
# before the first gradient step. As the policy learns to stay in-dist, episodes
|
||||||
|
# will get longer and HER has more transitions to relabel.
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
if args.load:
|
||||||
|
print(f"Hot-starting from {args.load}")
|
||||||
|
model = SAC.load(args.load, env=env, device='auto',
|
||||||
|
custom_objects={'learning_rate': 3e-4, 'batch_size': BATCH_SIZE,
|
||||||
|
'tau': 0.005, 'gamma': 0.98,
|
||||||
|
'train_freq': 64, 'gradient_steps': 8,
|
||||||
|
'learning_starts': MAX_EPISODE_STEPS,
|
||||||
|
'ent_coef': 0.1})
|
||||||
|
else:
|
||||||
|
model = SAC(
|
||||||
|
'MultiInputPolicy',
|
||||||
|
env,
|
||||||
|
replay_buffer_class=HerReplayBuffer,
|
||||||
|
replay_buffer_kwargs={
|
||||||
|
'n_sampled_goal': 4,
|
||||||
|
'goal_selection_strategy': 'future',
|
||||||
|
},
|
||||||
|
verbose=1,
|
||||||
|
learning_rate=3e-4,
|
||||||
|
batch_size=BATCH_SIZE,
|
||||||
|
tau=0.005,
|
||||||
|
gamma=0.98,
|
||||||
|
train_freq=64,
|
||||||
|
gradient_steps=8,
|
||||||
|
learning_starts=BATCH_SIZE,
|
||||||
|
ent_coef=0.1, # fixed; auto-tuning diverges on this many action dims
|
||||||
|
device='auto',
|
||||||
|
)
|
||||||
|
|
||||||
|
checkpoint_cb = CheckpointCallback(
|
||||||
|
save_freq=10_000,
|
||||||
|
save_path=args.out + '_checkpoints/',
|
||||||
|
name_prefix='sac',
|
||||||
|
)
|
||||||
|
|
||||||
|
import json, os
|
||||||
|
|
||||||
|
config = {'obs_params': OBS_PARAMS}
|
||||||
|
for save_dir in [args.out + '_checkpoints/', os.path.dirname(args.out) or '.']:
|
||||||
|
os.makedirs(save_dir, exist_ok=True)
|
||||||
|
with open(os.path.join(save_dir, 'config.json'), 'w') as f:
|
||||||
|
json.dump(config, f)
|
||||||
|
|
||||||
|
model.learn(total_timesteps=args.steps, callback=checkpoint_cb)
|
||||||
|
model.save(args.out)
|
||||||
|
with open(args.out + '.json', 'w') as f:
|
||||||
|
json.dump(config, f)
|
||||||
|
print(f"Saved to {args.out}.zip")
|
||||||
Reference in New Issue
Block a user