From f16e1508a0a58b8bec39fdd114430cef96bbfdca Mon Sep 17 00:00:00 2001 From: Enis Lorenz Date: Thu, 15 Jan 2026 15:45:30 +0100 Subject: [PATCH 1/2] Add carry-over state support and per-episode metrics. test_sanity_env now checks for determinism and carry-over continuitiy. Add carry-over state handling in the refactored environment by initializing timeline state on construction and skipping resets when enabled, while separating cumulative from episode-scoped metrics. Update MetricsTracker to track episode_* counters, propagate those counters through job assignment and baseline steps, and ensure episode summaries use per-episode values. Extend sanity tests and metrics usage to align with the refactor (metrics.current_hour, carry-over continuity checks). --- src/baseline.py | 15 ++-- src/environment.py | 151 +++++++++++++++++++++++++++++++++++----- src/job_management.py | 26 ++++--- src/metrics_tracker.py | 108 ++++++++++++++++------------ test/test_sanity_env.py | 62 +++++++++++++++-- 5 files changed, 280 insertions(+), 82 deletions(-) diff --git a/src/baseline.py b/src/baseline.py index 92024c8..c771235 100644 --- a/src/baseline.py +++ b/src/baseline.py @@ -38,8 +38,12 @@ def baseline_step(baseline_state, baseline_cores_available, baseline_running_job job_queue_2d, new_jobs_count, new_jobs_durations, new_jobs_nodes, new_jobs_cores, baseline_next_empty_slot ) - metrics.baseline_jobs_submitted += len(new_baseline_jobs) - metrics.baseline_jobs_rejected_queue_full += (new_jobs_count - len(new_baseline_jobs)) + metrics['baseline_jobs_submitted'] += len(new_baseline_jobs) + if 'episode_baseline_jobs_submitted' in metrics: + metrics['episode_baseline_jobs_submitted'] += len(new_baseline_jobs) + metrics['baseline_jobs_rejected_queue_full'] += (new_jobs_count - len(new_baseline_jobs)) + if 'episode_baseline_jobs_rejected_queue_full' in metrics: + metrics['episode_baseline_jobs_rejected_queue_full'] += (new_jobs_count - len(new_baseline_jobs)) _, baseline_next_empty_slot, _, next_job_id = assign_jobs_to_available_nodes( job_queue_2d, baseline_state['nodes'], baseline_cores_available, @@ -52,8 +56,11 @@ def baseline_step(baseline_state, baseline_cores_available, baseline_running_job num_unprocessed_jobs = np.sum(job_queue_2d[:, 0] > 0) # Track baseline max queue size - if num_unprocessed_jobs > metrics.baseline_max_queue_size_reached: - metrics.baseline_max_queue_size_reached = num_unprocessed_jobs + if num_unprocessed_jobs > metrics['baseline_max_queue_size_reached']: + metrics['baseline_max_queue_size_reached'] = num_unprocessed_jobs + if 'episode_baseline_max_queue_size_reached' in metrics: + if num_unprocessed_jobs > metrics['episode_baseline_max_queue_size_reached']: + metrics['episode_baseline_max_queue_size_reached'] = num_unprocessed_jobs baseline_state['job_queue'] = job_queue_2d.flatten() diff --git a/src/environment.py b/src/environment.py index 9129e73..af8f5be 100644 --- a/src/environment.py +++ b/src/environment.py @@ -80,7 +80,8 @@ def __init__(self, skip_plot_job_queue, steps_per_iteration, evaluation_mode=False, - workload_gen=None): + workload_gen=None, + carry_over_state=False): super().__init__() self.weights = weights @@ -119,6 +120,7 @@ def __init__(self, self.np_random = None self._seed = None self.workload_gen = workload_gen + self.carry_over_state = carry_over_state if self.external_durations: durations_sampler.init(self.external_durations) @@ -152,6 +154,9 @@ def __init__(self, # Initialize reward calculator self.reward_calculator = RewardCalculator(self.prices) + self.reset_timeline_state() + self.metrics.reset_episode_metrics() + # actions: - change number of available nodes: # action_type: 0: decrease, 1: maintain, 2: increase # action_magnitude: 0-MAX_CHANGE (+1ed in the action) @@ -192,21 +197,59 @@ def reset(self, seed=None, options=None): self.episode_idx += 1 # Reset metrics - self.metrics.reset_state_metrics() - - # Choose starting index in the external price series - if self.prices is not None and getattr(self.prices, "external_prices", None) is not None: - n_prices = len(self.prices.external_prices) - episode_span = EPISODE_HOURS - - # Episode k starts at hour k * episode_span (wrapping around the year) - start_index = (self.episode_idx * episode_span) % n_prices - if options and "price_start_index" in options: # For testing Purposes. Leave out 'options' to advance episode. - start_index = int(options["price_start_index"]) % n_prices - self.prices.reset(start_index=start_index) + if self.carry_over_state: + self.metrics.reset_episode_metrics() else: - # Synthetic prices or no external prices - self.prices.reset(start_index=0) + self.metrics.reset_state_metrics() + + self.metrics.current_hour = 0 + + if not self.carry_over_state: + # Choose starting index in the external price series + if self.prices is not None and getattr(self.prices, "external_prices", None) is not None: + n_prices = len(self.prices.external_prices) + episode_span = EPISODE_HOURS + + # Episode k starts at hour k * episode_span (wrapping around the year) + start_index = (self.episode_idx * episode_span) % n_prices + if options and "price_start_index" in options: # For testing Purposes. Leave out 'options' to advance episode. + start_index = int(options["price_start_index"]) % n_prices + self.prices.reset(start_index=start_index) + else: + # Synthetic prices or no external prices + self.prices.reset(start_index=0) + + self.state = { + # Initialize all nodes to be 'online but free' (0) + 'nodes': np.zeros(MAX_NODES, dtype=np.int32), + # Initialize job queue to be empty + 'job_queue': np.zeros((MAX_QUEUE_SIZE * 4), dtype=np.int32), + # Initialize predicted prices array + 'predicted_prices': self.prices.predicted_prices.copy(), + } + + self.baseline_state = { + 'nodes': np.zeros(MAX_NODES, dtype=np.int32), + 'job_queue': np.zeros((MAX_QUEUE_SIZE * 4), dtype=np.int32), + } + + self.cores_available = np.full(MAX_NODES, CORES_PER_NODE, dtype=np.int32) + self.baseline_cores_available = np.full(MAX_NODES, CORES_PER_NODE, dtype=np.int32) + + # Job tracking: { job_id: {'duration': remaining_hours, 'allocation': [(node_idx1, cores1), ...]}, ... } + self.running_jobs = {} + self.baseline_running_jobs = {} + + self.next_job_id = 0 # shared between baseline and normal jobs + + # Track next empty slot in job queue for O(1) insertion + self.next_empty_slot = 0 + self.baseline_next_empty_slot = 0 + + return self.state, {} + + def reset_timeline_state(self): + self.metrics.current_hour = 0 self.state = { # Initialize all nodes to be 'online but free' (0) @@ -224,7 +267,6 @@ def reset(self, seed=None, options=None): self.cores_available = np.full(MAX_NODES, CORES_PER_NODE, dtype=np.int32) self.baseline_cores_available = np.full(MAX_NODES, CORES_PER_NODE, dtype=np.int32) - # Job tracking: { job_id: {'duration': remaining_hours, 'allocation': [(node_idx1, cores1), ...]}, ... } self.running_jobs = {} self.baseline_running_jobs = {} @@ -235,11 +277,10 @@ def reset(self, seed=None, options=None): self.next_empty_slot = 0 self.baseline_next_empty_slot = 0 - return self.state, {} - def step(self, action): self.current_step += 1 self.metrics.current_hour += 1 + self.metrics.total_time_hours += 1 if self.metrics.current_hour == 1: self.current_episode += 1 self.env_print(Fore.GREEN + f"\n[[[ Starting episode: {self.current_episode}, step: {self.current_step}, hour: {self.metrics.current_hour}" + Fore.RESET) @@ -272,7 +313,9 @@ def step(self, action): new_jobs_nodes, new_jobs_cores, self.next_empty_slot ) self.metrics.jobs_submitted += len(new_jobs) + self.metrics.episode_jobs_submitted += len(new_jobs) self.metrics.jobs_rejected_queue_full += (new_jobs_count - len(new_jobs)) + self.metrics.episode_jobs_rejected_queue_full += (new_jobs_count - len(new_jobs)) self.env_print("nodes: ", np.array2string(self.state['nodes'], separator=' ', max_line_width=np.inf)) self.env_print(f"cores_available: {np.array2string(self.cores_available, separator=' ', max_line_width=np.inf)} ({np.sum(self.cores_available)})") @@ -288,11 +331,38 @@ def step(self, action): # Assign jobs to available nodes self.env_print(f"[4] Assigning jobs to available nodes...") + # Create metrics dict for job assignment + job_metrics = { + 'jobs_completed': self.metrics.jobs_completed, + 'total_job_wait_time': self.metrics.total_job_wait_time, + 'jobs_dropped': self.metrics.jobs_dropped, + 'dropped_this_episode': self.metrics.dropped_this_episode, + 'baseline_jobs_completed': self.metrics.baseline_jobs_completed, + 'baseline_total_job_wait_time': self.metrics.baseline_total_job_wait_time, + 'baseline_jobs_dropped': self.metrics.baseline_jobs_dropped, + 'baseline_dropped_this_episode': self.metrics.baseline_dropped_this_episode, + 'episode_jobs_completed': self.metrics.episode_jobs_completed, + 'episode_total_job_wait_time': self.metrics.episode_total_job_wait_time, + 'episode_jobs_dropped': self.metrics.episode_jobs_dropped, + 'episode_baseline_jobs_completed': self.metrics.episode_baseline_jobs_completed, + 'episode_baseline_total_job_wait_time': self.metrics.episode_baseline_total_job_wait_time, + 'episode_baseline_jobs_dropped': self.metrics.episode_baseline_jobs_dropped, + } + num_launched_jobs, self.next_empty_slot, num_dropped_this_step, self.next_job_id = assign_jobs_to_available_nodes( job_queue_2d, self.state['nodes'], self.cores_available, self.running_jobs, self.next_empty_slot, self.next_job_id, self.metrics, is_baseline=False ) + # Update metrics from job_metrics dict + self.metrics.jobs_completed = job_metrics['jobs_completed'] + self.metrics.total_job_wait_time = job_metrics['total_job_wait_time'] + self.metrics.jobs_dropped = job_metrics['jobs_dropped'] + self.metrics.dropped_this_episode = job_metrics['dropped_this_episode'] + self.metrics.episode_jobs_completed = job_metrics['episode_jobs_completed'] + self.metrics.episode_total_job_wait_time = job_metrics['episode_total_job_wait_time'] + self.metrics.episode_jobs_dropped = job_metrics['episode_jobs_dropped'] + self.env_print(f" {num_launched_jobs} jobs launched") # Calculate node utilization stats @@ -313,18 +383,53 @@ def step(self, action): # Track max queue size if num_unprocessed_jobs > self.metrics.max_queue_size_reached: self.metrics.max_queue_size_reached = num_unprocessed_jobs + if num_unprocessed_jobs > self.metrics.episode_max_queue_size_reached: + self.metrics.episode_max_queue_size_reached = num_unprocessed_jobs self.env_print(f"[5] Calculating reward...") # Baseline step + baseline_metrics = { + 'baseline_jobs_submitted': self.metrics.baseline_jobs_submitted, + 'baseline_jobs_rejected_queue_full': self.metrics.baseline_jobs_rejected_queue_full, + 'baseline_jobs_completed': self.metrics.baseline_jobs_completed, + 'baseline_total_job_wait_time': self.metrics.baseline_total_job_wait_time, + 'baseline_jobs_dropped': self.metrics.baseline_jobs_dropped, + 'baseline_dropped_this_episode': self.metrics.baseline_dropped_this_episode, + 'baseline_max_queue_size_reached': self.metrics.baseline_max_queue_size_reached, + 'episode_baseline_jobs_submitted': self.metrics.episode_baseline_jobs_submitted, + 'episode_baseline_jobs_rejected_queue_full': self.metrics.episode_baseline_jobs_rejected_queue_full, + 'episode_baseline_jobs_completed': self.metrics.episode_baseline_jobs_completed, + 'episode_baseline_total_job_wait_time': self.metrics.episode_baseline_total_job_wait_time, + 'episode_baseline_jobs_dropped': self.metrics.episode_baseline_jobs_dropped, + 'episode_baseline_max_queue_size_reached': self.metrics.episode_baseline_max_queue_size_reached, + } + baseline_cost, baseline_cost_off, self.baseline_next_empty_slot, self.next_job_id = baseline_step( self.baseline_state, self.baseline_cores_available, self.baseline_running_jobs, current_price, new_jobs_count, new_jobs_durations, new_jobs_nodes, new_jobs_cores, self.baseline_next_empty_slot, self.next_job_id, self.metrics, self.env_print ) + # Update metrics from baseline_metrics dict + self.metrics.baseline_jobs_submitted = baseline_metrics['baseline_jobs_submitted'] + self.metrics.baseline_jobs_rejected_queue_full = baseline_metrics['baseline_jobs_rejected_queue_full'] + self.metrics.baseline_jobs_completed = baseline_metrics['baseline_jobs_completed'] + self.metrics.baseline_total_job_wait_time = baseline_metrics['baseline_total_job_wait_time'] + self.metrics.baseline_jobs_dropped = baseline_metrics['baseline_jobs_dropped'] + self.metrics.baseline_dropped_this_episode = baseline_metrics['baseline_dropped_this_episode'] + self.metrics.baseline_max_queue_size_reached = baseline_metrics['baseline_max_queue_size_reached'] + self.metrics.episode_baseline_jobs_submitted = baseline_metrics['episode_baseline_jobs_submitted'] + self.metrics.episode_baseline_jobs_rejected_queue_full = baseline_metrics['episode_baseline_jobs_rejected_queue_full'] + self.metrics.episode_baseline_jobs_completed = baseline_metrics['episode_baseline_jobs_completed'] + self.metrics.episode_baseline_total_job_wait_time = baseline_metrics['episode_baseline_total_job_wait_time'] + self.metrics.episode_baseline_jobs_dropped = baseline_metrics['episode_baseline_jobs_dropped'] + self.metrics.episode_baseline_max_queue_size_reached = baseline_metrics['episode_baseline_max_queue_size_reached'] + self.metrics.baseline_cost += baseline_cost self.metrics.baseline_cost_off += baseline_cost_off + self.metrics.episode_baseline_cost += baseline_cost + self.metrics.episode_baseline_cost_off += baseline_cost_off step_reward, step_cost, eff_reward_norm, price_reward_norm, idle_penalty_norm, job_age_penalty_norm = self.reward_calculator.calculate( num_used_nodes, num_idle_nodes, current_price, average_future_price, @@ -334,6 +439,7 @@ def step(self, action): self.metrics.episode_reward += step_reward self.metrics.total_cost += step_cost + self.metrics.episode_total_cost += step_cost # Store normalized reward components for plotting self.metrics.eff_rewards.append(eff_reward_norm * 100) @@ -385,4 +491,11 @@ def step(self, action): self.env_print(Fore.GREEN + f"]]]" + Fore.RESET) - return self.state, step_reward, terminated, truncated, {} + info = { + "step_cost": float(step_cost), + "num_unprocessed_jobs": int(num_unprocessed_jobs), + "num_on_nodes": int(num_on_nodes), + "dropped_this_episode": int(getattr(self.metrics, "dropped_this_episode", 0)), + } + + return self.state, step_reward, terminated, truncated, info diff --git a/src/job_management.py b/src/job_management.py index 0fc0e34..54748c0 100644 --- a/src/job_management.py +++ b/src/job_management.py @@ -154,11 +154,17 @@ def assign_jobs_to_available_nodes(job_queue_2d, nodes, cores_available, running # Track job completion and wait time if is_baseline: - metrics.baseline_jobs_completed += 1 - metrics.baseline_total_job_wait_time += job_age + metrics['baseline_jobs_completed'] += 1 + metrics['baseline_total_job_wait_time'] += job_age + if 'episode_baseline_jobs_completed' in metrics: + metrics['episode_baseline_jobs_completed'] += 1 + metrics['episode_baseline_total_job_wait_time'] += job_age else: - metrics.jobs_completed += 1 - metrics.total_job_wait_time += job_age + metrics['jobs_completed'] += 1 + metrics['total_job_wait_time'] += job_age + if 'episode_jobs_completed' in metrics: + metrics['episode_jobs_completed'] += 1 + metrics['episode_total_job_wait_time'] += job_age num_processed_jobs += 1 continue @@ -176,11 +182,15 @@ def assign_jobs_to_available_nodes(job_queue_2d, nodes, cores_available, running num_dropped += 1 if is_baseline: - metrics.baseline_jobs_dropped += 1 - metrics.baseline_dropped_this_episode += 1 + metrics['baseline_jobs_dropped'] += 1 + metrics['baseline_dropped_this_episode'] += 1 + if 'episode_baseline_jobs_dropped' in metrics: + metrics['episode_baseline_jobs_dropped'] += 1 else: - metrics.jobs_dropped += 1 - metrics.dropped_this_episode += 1 + metrics['jobs_dropped'] += 1 + metrics['dropped_this_episode'] += 1 + if 'episode_jobs_dropped' in metrics: + metrics['episode_jobs_dropped'] += 1 else: job_queue_2d[job_idx][1] = new_age diff --git a/src/metrics_tracker.py b/src/metrics_tracker.py index 46ba1d1..aa8d163 100644 --- a/src/metrics_tracker.py +++ b/src/metrics_tracker.py @@ -6,27 +6,48 @@ class MetricsTracker: def __init__(self): """Initialize all metric counters.""" - self.reset_episode_metrics() + self.reset_cumulative_metrics() self.reset_state_metrics() # Cumulative metrics across all episodes self.episode_costs = [] - def reset_episode_metrics(self): + def reset_cumulative_metrics(self): """Reset metrics that persist across episodes.""" - # Job tracking metrics for agent (cumulative across episodes) + # Cost tracking (cumulative across episodes) + self.total_cost = 0 + self.baseline_cost = 0 + self.baseline_cost_off = 0 + + # Agent job metrics (cumulative across episodes) + self.jobs_submitted = 0 + self.jobs_completed = 0 + self.total_job_wait_time = 0 + self.max_queue_size_reached = 0 self.jobs_dropped = 0 self.jobs_rejected_queue_full = 0 - # Job tracking metrics for baseline (cumulative across episodes) + # Baseline job metrics (cumulative across episodes) + self.baseline_jobs_submitted = 0 + self.baseline_jobs_completed = 0 + self.baseline_total_job_wait_time = 0 + self.baseline_max_queue_size_reached = 0 self.baseline_jobs_dropped = 0 self.baseline_jobs_rejected_queue_full = 0 def reset_state_metrics(self): + """Reset timeline-dependent metrics.""" + self.current_hour = 0 + self.total_time_hours = 0 + self.reset_episode_metrics() + + def reset_episode_metrics(self): """Reset metrics at the start of each episode.""" # Episode-level metrics - self.current_hour = 0 self.episode_reward = 0 + self.episode_total_cost = 0 + self.episode_baseline_cost = 0 + self.episode_baseline_cost_off = 0 # Time series data for plotting self.on_nodes = [] @@ -39,22 +60,21 @@ def reset_state_metrics(self): self.idle_penalties = [] self.job_age_penalties = [] - # Cost tracking - self.total_cost = 0 - self.baseline_cost = 0 - self.baseline_cost_off = 0 - - # Agent job metrics - self.jobs_submitted = 0 - self.jobs_completed = 0 - self.total_job_wait_time = 0 - self.max_queue_size_reached = 0 - - # Baseline job metrics - self.baseline_jobs_submitted = 0 - self.baseline_jobs_completed = 0 - self.baseline_total_job_wait_time = 0 - self.baseline_max_queue_size_reached = 0 + # Agent job metrics (episode-level) + self.episode_jobs_submitted = 0 + self.episode_jobs_completed = 0 + self.episode_total_job_wait_time = 0 + self.episode_max_queue_size_reached = 0 + self.episode_jobs_dropped = 0 + self.episode_jobs_rejected_queue_full = 0 + + # Baseline job metrics (episode-level) + self.episode_baseline_jobs_submitted = 0 + self.episode_baseline_jobs_completed = 0 + self.episode_baseline_total_job_wait_time = 0 + self.episode_baseline_max_queue_size_reached = 0 + self.episode_baseline_jobs_dropped = 0 + self.episode_baseline_jobs_rejected_queue_full = 0 # Per-episode drop counters self.dropped_this_episode = 0 @@ -71,45 +91,45 @@ def record_episode_completion(self, current_episode): Dictionary with episode data """ # Calculate average wait times - avg_wait_time = self.total_job_wait_time / self.jobs_completed if self.jobs_completed > 0 else 0 - baseline_avg_wait_time = self.baseline_total_job_wait_time / self.baseline_jobs_completed if self.baseline_jobs_completed > 0 else 0 + avg_wait_time = self.episode_total_job_wait_time / self.episode_jobs_completed if self.episode_jobs_completed > 0 else 0 + baseline_avg_wait_time = self.episode_baseline_total_job_wait_time / self.episode_baseline_jobs_completed if self.episode_baseline_jobs_completed > 0 else 0 # Calculate completion rates - completion_rate = (self.jobs_completed / self.jobs_submitted * 100) if self.jobs_submitted > 0 else 0 - baseline_completion_rate = (self.baseline_jobs_completed / self.baseline_jobs_submitted * 100) if self.baseline_jobs_submitted > 0 else 0 + completion_rate = (self.episode_jobs_completed / self.episode_jobs_submitted * 100) if self.episode_jobs_submitted > 0 else 0 + baseline_completion_rate = (self.episode_baseline_jobs_completed / self.episode_baseline_jobs_submitted * 100) if self.episode_baseline_jobs_submitted > 0 else 0 - drop_rate = (self.jobs_dropped / self.jobs_submitted * 100) if self.jobs_submitted else 0.0 - baseline_drop_rate = (self.baseline_jobs_dropped / self.baseline_jobs_submitted * 100) if self.baseline_jobs_submitted else 0.0 + drop_rate = (self.episode_jobs_dropped / self.episode_jobs_submitted * 100) if self.episode_jobs_submitted else 0.0 + baseline_drop_rate = (self.episode_baseline_jobs_dropped / self.episode_baseline_jobs_submitted * 100) if self.episode_baseline_jobs_submitted else 0.0 episode_data = { 'episode': current_episode, - 'agent_cost': float(self.total_cost), - 'baseline_cost': float(self.baseline_cost), - 'baseline_cost_off': float(self.baseline_cost_off), - 'savings_vs_baseline': float(self.baseline_cost - self.total_cost), - 'savings_vs_baseline_off': float(self.baseline_cost_off - self.total_cost), - 'savings_pct_baseline': float(((self.baseline_cost - self.total_cost) / self.baseline_cost) * 100) if self.baseline_cost > 0 else 0, - 'savings_pct_baseline_off': float(((self.baseline_cost_off - self.total_cost) / self.baseline_cost_off) * 100) if self.baseline_cost_off > 0 else 0, + 'agent_cost': float(self.episode_total_cost), + 'baseline_cost': float(self.episode_baseline_cost), + 'baseline_cost_off': float(self.episode_baseline_cost_off), + 'savings_vs_baseline': float(self.episode_baseline_cost - self.episode_total_cost), + 'savings_vs_baseline_off': float(self.episode_baseline_cost_off - self.episode_total_cost), + 'savings_pct_baseline': float(((self.episode_baseline_cost - self.episode_total_cost) / self.episode_baseline_cost) * 100) if self.episode_baseline_cost > 0 else 0, + 'savings_pct_baseline_off': float(((self.episode_baseline_cost_off - self.episode_total_cost) / self.episode_baseline_cost_off) * 100) if self.episode_baseline_cost_off > 0 else 0, 'total_reward': float(self.episode_reward), # Agent job metrics - 'jobs_submitted': self.jobs_submitted, - 'jobs_completed': self.jobs_completed, + 'jobs_submitted': self.episode_jobs_submitted, + 'jobs_completed': self.episode_jobs_completed, 'avg_wait_time': float(avg_wait_time), 'completion_rate': float(completion_rate), - 'max_queue_size': self.max_queue_size_reached, + 'max_queue_size': self.episode_max_queue_size_reached, # Baseline job metrics - 'baseline_jobs_submitted': self.baseline_jobs_submitted, - 'baseline_jobs_completed': self.baseline_jobs_completed, + 'baseline_jobs_submitted': self.episode_baseline_jobs_submitted, + 'baseline_jobs_completed': self.episode_baseline_jobs_completed, 'baseline_avg_wait_time': float(baseline_avg_wait_time), 'baseline_completion_rate': float(baseline_completion_rate), - 'baseline_max_queue_size': self.baseline_max_queue_size_reached, + 'baseline_max_queue_size': self.episode_baseline_max_queue_size_reached, # Drop metrics - "jobs_dropped": self.jobs_dropped, + "jobs_dropped": self.episode_jobs_dropped, "drop_rate": float(drop_rate), - "jobs_rejected_queue_full": self.jobs_rejected_queue_full, - "baseline_jobs_dropped": self.baseline_jobs_dropped, + "jobs_rejected_queue_full": self.episode_jobs_rejected_queue_full, + "baseline_jobs_dropped": self.episode_baseline_jobs_dropped, "baseline_drop_rate": float(baseline_drop_rate), - "baseline_jobs_rejected_queue_full": self.baseline_jobs_rejected_queue_full, + "baseline_jobs_rejected_queue_full": self.episode_baseline_jobs_rejected_queue_full, } self.episode_costs.append(episode_data) return episode_data diff --git a/test/test_sanity_env.py b/test/test_sanity_env.py index 2a68af3..777e13f 100644 --- a/test/test_sanity_env.py +++ b/test/test_sanity_env.py @@ -10,6 +10,7 @@ ''' import argparse +import copy import numpy as np from gymnasium.utils.env_checker import check_env @@ -117,12 +118,12 @@ def determinism_test(make_env, seed, n_steps=200): def rollout(): env = make_env() # Pin external price window so determinism doesn't vary by episode. - obs, _ = env.reset(seed=seed, options={"price_start_index": 0}) + _obs, _ = env.reset(seed=seed, options={"price_start_index": 0}) traj = [] done = False i = 0 while not done and i < n_steps: - obs, r, term, trunc, info = env.step(actions[i]) + _obs, r, term, trunc, info = env.step(actions[i]) # record a small fingerprint traj.append(( float(r), @@ -140,6 +141,41 @@ def rollout(): b = rollout() assert a == b, "Determinism failed: same seed + same actions produced different trajectories" +def carry_over_test(make_env, seed, n_steps=1): + env = make_env() + obs, _ = env.reset(seed=seed) + env.action_space.seed(seed) + actions = [env.action_space.sample() for _ in range(n_steps)] + for action in actions: + obs, r, term, trunc, info = env.step(action) + if term or trunc: + break + + snapshot = { + "nodes": env.state["nodes"].copy(), + "job_queue": env.state["job_queue"].copy(), + "predicted_prices": env.state["predicted_prices"].copy(), + "cores_available": env.cores_available.copy(), + "running_jobs": copy.deepcopy(env.running_jobs), + "price_index": env.prices.price_index, + "next_job_id": env.next_job_id, + "next_empty_slot": env.next_empty_slot, + "current_hour": env.metrics.current_hour, + } + + env.reset(seed=seed) + + assert np.array_equal(env.state["nodes"], snapshot["nodes"]), "carry-over failed: nodes reset" + assert np.array_equal(env.state["job_queue"], snapshot["job_queue"]), "carry-over failed: job_queue reset" + assert np.array_equal(env.state["predicted_prices"], snapshot["predicted_prices"]), "carry-over failed: predicted_prices reset" + assert np.array_equal(env.cores_available, snapshot["cores_available"]), "carry-over failed: cores_available reset" + assert env.running_jobs == snapshot["running_jobs"], "carry-over failed: running_jobs reset" + assert env.prices.price_index == snapshot["price_index"], "carry-over failed: price_index reset" + assert env.next_job_id == snapshot["next_job_id"], "carry-over failed: next_job_id reset" + assert env.next_empty_slot == snapshot["next_empty_slot"], "carry-over failed: next_empty_slot reset" + assert env.metrics.current_hour == 0, "carry-over failed: current_hour not reset" + env.close() + # ----------------------------- # CLI + env construction @@ -172,6 +208,8 @@ def parse_args(): help="Where to sample the job from.") p.add_argument("--print-job-index", type=int, default=-1, help="Queue index to print (>=0), or -1 to print first active job.") + p.add_argument("--carry-over-state", action="store_true", + help="Carry over nodes/jobs/prices across episodes (timeline mode).") return p.parse_args() @@ -229,8 +267,9 @@ def norm_path(x): skip_plot_job_queue=True, steps_per_iteration=EPISODE_HOURS, # prevent plot cadence surprises evaluation_mode=False, - # plot_total_reward=False, - workload_gen=workload_gen + # plot_total_reward=False, + workload_gen=workload_gen, + carry_over_state=args.carry_over_state ) def maybe_print_job(env, obs, step_idx, every, kind="queue", job_index=-1): @@ -287,13 +326,17 @@ def reset(self, seed=None, options=None): options["price_start_index"] = 0 return super().reset(seed=seed, options=options) + def make_env_with_carry(carry_over_state, env_cls=ComputeClusterEnv): + local_args = argparse.Namespace(**vars(args)) + local_args.carry_over_state = carry_over_state + return make_env_from_args(local_args, env_cls=env_cls) # ------------------------------------- seed = 123 action = np.array([1, 0], dtype=np.int64) # "maintain, magnitude 1" effectively - env = make_env_from_args(args, env_cls=DeterministicPriceEnv) + env = make_env_with_carry(False, env_cls=DeterministicPriceEnv) o1, _ = env.reset(seed=seed) o1s, r1, t1, tr1, i1 = env.step(action) @@ -320,7 +363,7 @@ def cmp(name, a, b): # 1) Gym API compliance (optional) if args.check_gym: # Pin external price window so gym's determinism check is meaningful. - env = make_env_from_args(args, env_cls=DeterministicPriceEnv) + env = make_env_with_carry(False, env_cls=DeterministicPriceEnv) check_env(env, skip_render_check=True) env.close() print("[OK] gymnasium check_env passed") @@ -354,9 +397,14 @@ def cmp(name, a, b): # 3) Determinism (optional) if args.check_determinism: - determinism_test(lambda: make_env_from_args(args), seed=args.seed, n_steps=min(args.steps, 500)) + determinism_test(lambda: make_env_with_carry(False), seed=args.seed, n_steps=min(args.steps, 500)) print("[OK] determinism test passed") + # 4) Carry-over continuity (optional) + if args.carry_over_state: + carry_over_test(lambda: make_env_with_carry(True), seed=args.seed, n_steps=min(args.steps, 10)) + print("[OK] carry-over continuity test passed") + print("done") From 1a15eeac6bd9d161e2f586275a178f59c530a7c6 Mon Sep 17 00:00:00 2001 From: Enis Lorenz Date: Tue, 20 Jan 2026 16:41:33 +0100 Subject: [PATCH 2/2] Add: Carry over State. Some missing parts. --- src/plotter.py | 10 +++++----- train.py | 4 +++- train_iter.py | 21 ++++++++++++++++++++- 3 files changed, 28 insertions(+), 7 deletions(-) diff --git a/src/plotter.py b/src/plotter.py index 6ac922f..3ed9904 100644 --- a/src/plotter.py +++ b/src/plotter.py @@ -149,15 +149,15 @@ def add_panel(title, series, ylabel, ylim=None, overlay=None): ) # Reward components - if getattr(env, "plot_eff_reward", False): + if not getattr(env, "plot_eff_reward", False): add_panel("Efficiency reward (%)", getattr(env.metrics, "eff_rewards", None), "score", None) - if getattr(env, "plot_price_reward", False): + if not getattr(env, "plot_price_reward", False): add_panel("Price reward (%)", getattr(env.metrics, "price_rewards", None), "score", None) - if getattr(env, "plot_idle_penalty", False): + if not getattr(env, "plot_idle_penalty", False): add_panel("Idle penalty (%)", getattr(env.metrics, "idle_penalties", None), "score", None) - if getattr(env, "plot_job_age_penalty", False): + if not getattr(env, "plot_job_age_penalty", False): add_panel("Job-age penalty (%)", getattr(env.metrics, "job_age_penalties", None), "score", None) - if getattr(env, "plot_total_reward", False): + if not getattr(env, "plot_total_reward", False): add_panel("Total reward", getattr(env.metrics, "rewards", None), "reward", None) if not panels: diff --git a/train.py b/train.py index 91f71b4..2bdb01e 100644 --- a/train.py +++ b/train.py @@ -59,6 +59,7 @@ def main(): parser.add_argument("--wg-max-jobs-hour", type=int, default=1500, help="Cap jobs/hour for generator.") parser.add_argument("--plot-dashboard", action="store_true", help="Generate dashboard plot (per-hour panels + cumulative savings).") parser.add_argument("--dashboard-hours", type=int, default=24*14, help="Hours to show in dashboard time-series panels (default: 336).") + parser.add_argument("--carry-over-state", action="store_true", help="Carry over nodes/jobs/prices across episodes (timeline mode).") args = parser.parse_args() prices_file_path = args.prices @@ -136,7 +137,8 @@ def main(): skip_plot_job_queue=args.skip_plot_job_queue, steps_per_iteration=STEPS_PER_ITERATION, evaluation_mode=args.evaluate_savings, - workload_gen=workload_gen) + workload_gen=workload_gen, + carry_over_state=args.carry_over_state) env.reset() # Check if there are any saved models in models_dir diff --git a/train_iter.py b/train_iter.py index d9a0fc7..6094b99 100644 --- a/train_iter.py +++ b/train_iter.py @@ -95,6 +95,7 @@ def run( hourly_jobs, plot_dashboard=False, dashboard_hours=24 * 14, + carry_over_state=False, ): python_executable = sys.executable command = [ @@ -113,6 +114,8 @@ def run( ] if plot_dashboard: command += ["--plot-dashboard", "--dashboard-hours", str(dashboard_hours)] + if carry_over_state: + command += ["--carry-over-state"] print(f"executing: {command}") current_env = os.environ.copy() @@ -151,6 +154,7 @@ def main(): parser.add_argument("--iter-limit-per-step", type=int, help="Max number of training iterations per step (1 iteration = {TIMESTEPS} steps)") parser.add_argument("--plot-dashboard", action="store_true", help="Forward to train.py to generate dashboard plots.") parser.add_argument("--dashboard-hours", type=int, default=24*14, help="Forward to train.py.") + parser.add_argument("--carry-over-state", action="store_true", help="Forward to train.py to carry state across episodes.") parser.add_argument("--session", help="Session ID") @@ -175,7 +179,22 @@ def main(): for combo in combinations: efficiency_weight, price_weight, idle_weight, job_age_weight, drop_weight = combo print(f"Running with weights: efficiency={efficiency_weight}, price={price_weight}, idle={idle_weight}, job_age={job_age_weight}, drop={drop_weight}") - run(efficiency_weight, price_weight, idle_weight, job_age_weight, drop_weight, args.iter_limit_per_step, args.session, args.prices, args.job_durations, args.jobs, args.hourly_jobs,plot_dashboard=args.plot_dashboard,dashboard_hours=args.dashboard_hours) + run( + efficiency_weight, + price_weight, + idle_weight, + job_age_weight, + drop_weight, + args.iter_limit_per_step, + args.session, + args.prices, + args.job_durations, + args.jobs, + args.hourly_jobs, + plot_dashboard=args.plot_dashboard, + dashboard_hours=args.dashboard_hours, + carry_over_state=args.carry_over_state, + ) if __name__ == "__main__": main()