From 864bfc6edfbd3fd12f5314003fb2644afdb7953e Mon Sep 17 00:00:00 2001 From: Enis Lorenz Date: Mon, 16 Feb 2026 13:25:12 +0100 Subject: [PATCH 1/9] Add: workloadgen tester to plot distributions. --- test/test_inspect_workloadgen.py | 178 +++++++++++++++++++++++++++++++ 1 file changed, 178 insertions(+) create mode 100644 test/test_inspect_workloadgen.py diff --git a/test/test_inspect_workloadgen.py b/test/test_inspect_workloadgen.py new file mode 100644 index 0000000..9ec8a84 --- /dev/null +++ b/test/test_inspect_workloadgen.py @@ -0,0 +1,178 @@ +""" +Run with: +python inspect_workloadgen.py --arrivals poisson --poisson-lambda 200 --max-jobs-hour 1500 --hours 336 --plot +""" + +# inspect_workloadgen.py +import argparse +import hashlib + +import numpy as np +import matplotlib.pyplot as plt +from datetime import datetime +import os + +from src.workloadgen import WorkloadGenConfig, WorkloadGenerator + + +def digest_jobs_triplets(triplets): + """ + Stable digest to verify determinism. + + We digest (hour_idx, duration, nodes, cores_per_node) so the hash is robust against + future refactors that might change how jobs are flattened/stored. + """ + arr = np.array(triplets, dtype=np.int32) # (hour, duration, nodes, cores) + return hashlib.sha256(arr.tobytes()).hexdigest() + + +def summarize(name, x): + x = np.asarray(x) + if x.size == 0: + print(f"{name}: (empty)") + return + qs = np.percentile(x, [0, 1, 10, 50, 90, 99, 100]) + print( + f"{name}: n={x.size} mean={x.mean():.3f} std={x.std():.3f} " + f"min/p1/p10/p50/p90/p99/max={qs[0]:.3f}/{qs[1]:.3f}/{qs[2]:.3f}/{qs[3]:.3f}/{qs[4]:.3f}/{qs[5]:.3f}/{qs[6]:.3f}" + ) + + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--seed", type=int, default=123) + ap.add_argument("--hours", type=int, default=24 * 14) + ap.add_argument("--arrivals", choices=["flat", "poisson", "uniform"], default="poisson") + ap.add_argument("--poisson-lambda", type=float, default=200.0) + ap.add_argument("--max-jobs-hour", type=int, default=1500) + ap.add_argument("--plot", action="store_true") + # Flat params (true flat with optional jitter) + ap.add_argument("--flat-jobs-hour",type=int,default=200,help="Target jobs/hour for arrivals=flat.") + ap.add_argument("--flat-jitter",type=int,default=0,help="Jitter for arrivals=flat. Sample in [target-jitter, target+jitter]. 0 => perfectly flat.") + args = ap.parse_args() + + cfg = WorkloadGenConfig( + arrivals=args.arrivals, + poisson_lambda=args.poisson_lambda, + max_new_jobs_per_hour=args.max_jobs_hour, + flat_jobs_per_hour=args.flat_jobs_hour, + flat_jitter=args.flat_jitter, + ) + gen = WorkloadGenerator(cfg) + + rng = np.random.default_rng(args.seed) + + jobs_per_hour = [] + node_hours_per_hour = [] + core_hours_per_hour = [] + + # NEW: robust digest input includes hour_idx + all_jobs_triplets = [] + + for h in range(args.hours): + jobs = gen.sample(h, rng) + jobs_per_hour.append(len(jobs)) + + # Track jobs with hour info for digest + future debugging + for j in jobs: + all_jobs_triplets.append((h, int(j.duration), int(j.nodes), int(j.cores_per_node))) + + nh = sum(int(j.duration) * int(j.nodes) for j in jobs) + ch = sum(int(j.duration) * int(j.nodes) * int(j.cores_per_node) for j in jobs) + node_hours_per_hour.append(nh) + core_hours_per_hour.append(ch) + + print(f"hours: {args.hours}") + summarize("jobs/hour", jobs_per_hour) + summarize("node-hours/hour", node_hours_per_hour) + summarize("core-hours/hour", core_hours_per_hour) + + if all_jobs_triplets: + # Unpack triplets for summarize (keeps your previous summaries unchanged) + durations = [t[1] for t in all_jobs_triplets] + nodes = [t[2] for t in all_jobs_triplets] + cpn = [t[3] for t in all_jobs_triplets] + summarize("duration[h]", durations) + summarize("nodes", nodes) + summarize("cores/node", cpn) + + print("digest:", digest_jobs_triplets(all_jobs_triplets)) + + # Optional determinism self-check (same seed => identical digest) + rng2 = np.random.default_rng(args.seed) + all_jobs_triplets_2 = [] + for h in range(args.hours): + for j in gen.sample(h, rng2): + all_jobs_triplets_2.append((h, int(j.duration), int(j.nodes), int(j.cores_per_node))) + + assert ( + digest_jobs_triplets(all_jobs_triplets) == digest_jobs_triplets(all_jobs_triplets_2) + ), "Generator not deterministic under same seed!" + + if args.plot: + # 4x2 grid so we can keep everything in one figure + fig, axs = plt.subplots(4, 2, figsize=(14, 16), constrained_layout=True) + + # time-series + axs[0, 0].plot(np.arange(args.hours), jobs_per_hour) + axs[0, 0].set_title("Jobs per hour over time") + axs[0, 0].set_xlabel("hour index") + axs[0, 0].set_ylabel("jobs") + + # histogram jobs/hour + axs[0, 1].hist(jobs_per_hour, bins=50) + axs[0, 1].set_title("Jobs per hour (hist)") + axs[0, 1].set_xlabel("jobs/hour") + axs[0, 1].set_ylabel("count") + + axs[1, 0].hist(node_hours_per_hour, bins=50) + axs[1, 0].set_title("Node-hours per hour (hourly workload volume)") + axs[1, 0].set_xlabel("node-hours/hour") + axs[1, 0].set_ylabel("count") + + axs[1, 1].hist(core_hours_per_hour, bins=50) + axs[1, 1].set_title("Core-hours per hour (total compute demand per hour)") + axs[1, 1].set_xlabel("core-hours/hour") + axs[1, 1].set_ylabel("count") + + # Unpack for plotting histograms + durations = [t[1] for t in all_jobs_triplets] if all_jobs_triplets else [] + nodes = [t[2] for t in all_jobs_triplets] if all_jobs_triplets else [] + cpn = [t[3] for t in all_jobs_triplets] if all_jobs_triplets else [] + + axs[2, 0].hist(durations, bins=50) + axs[2, 0].set_title("Durations (hours)") + axs[2, 0].set_xlabel("duration [h]") + axs[2, 0].set_ylabel("count") + + axs[2, 1].hist(nodes, bins=16) + axs[2, 1].set_title("Nodes (Jobs shape/Volume)") + axs[2, 1].set_xlabel("nodes") + axs[2, 1].set_ylabel("count") + + axs[3, 0].hist(cpn, bins=32) + axs[3, 0].set_title("Cores per node (Jobs shape/Volume)") + axs[3, 0].set_xlabel("cores/node") + axs[3, 0].set_ylabel("count") + + # jobs by hour-of-day + hod = np.arange(args.hours) % 24 + jobs_by_hod = np.zeros(24, dtype=np.int64) + for h, k in enumerate(jobs_per_hour): + jobs_by_hod[hod[h]] += int(k) + + axs[3, 1].bar(np.arange(24), jobs_by_hod) + axs[3, 1].set_title("Total jobs by hour-of-day") + axs[3, 1].set_xlabel("hour of day") + axs[3, 1].set_ylabel("jobs") + + plt.show() + timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") + prefix = f"{args.arrivals}_lambda{args.poisson_lambda}" if args.arrivals == "poisson" else args.arrivals + fname = f"{prefix}_{timestamp}.png" if prefix else f"Workload-Gen_{timestamp}.png" + save_path = os.path.join("", fname) + plt.savefig(save_path, dpi=250, bbox_inches="tight") + + +if __name__ == "__main__": + main() From 6c68920ee994fb5f0c1a6b30f3f35b8af4a34308 Mon Sep 17 00:00:00 2001 From: Enis Lorenz Date: Mon, 16 Feb 2026 13:56:22 +0100 Subject: [PATCH 2/9] Workload Generator: add mode-matched attribute sampling and 4-value parser passthrough Optional body: Sample duration/nodes/cores with flat|poisson|uniform (same mode as arrivals) using per-attribute params. Add *4 CLI args (arrivals,duration,nodes,cores) in inspect/sanity/train wiring. Add sanity coverage for flat attribute targets and poisson attribute lambdas. --- src/workloadgen.py | 137 +++++++++++++++++++++++++++++-- test/test_inspect_workloadgen.py | 126 ++++++++++++++++++++++++++-- test/test_sanity_env.py | 129 ++++++++++++++++++++++++++--- test/test_sanity_workloadgen.py | 59 +++++++++++++ train.py | 128 +++++++++++++++++++++++++++-- 5 files changed, 543 insertions(+), 36 deletions(-) diff --git a/src/workloadgen.py b/src/workloadgen.py index 2da9176..55188c3 100644 --- a/src/workloadgen.py +++ b/src/workloadgen.py @@ -33,12 +33,22 @@ class JobSpec: @dataclass(frozen=True) class WorkloadGenConfig: - # arrivals: "flat" or "poisson" + # arrivals mode shared across count + job attributes: "flat", "poisson", "uniform" arrivals: str = "poisson" + uniform_min_new_jobs_per_hour: int = 0 max_new_jobs_per_hour: int = 1500 poisson_lambda: float = 200.0 + poisson_lambda_duration: Optional[float] = None + poisson_lambda_nodes: Optional[float] = None + poisson_lambda_cores: Optional[float] = None flat_jobs_per_hour: int = 200 # target arrivals for flat mode flat_jitter: int = 0 # +/- jitter; 0 => perfectly flat + flat_duration_target: Optional[int] = None + flat_nodes_target: Optional[int] = None + flat_cores_target: Optional[int] = None + flat_duration_jitter: int = 0 + flat_nodes_jitter: int = 0 + flat_cores_jitter: int = 0 # resource ranges (v1: just uniform ranges; later we add mixtures/correlations) @@ -58,14 +68,93 @@ def __init__(self, cfg: WorkloadGenConfig): arrivals = cfg.arrivals.lower().strip() if arrivals not in ("flat", "poisson", "uniform"): raise ValueError(f"arrivals must be 'flat', 'uniform' or 'poisson', got: {cfg.arrivals}") - self.cfg = replace(cfg, arrivals=arrivals) + + duration_mid = int(round((int(cfg.min_duration) + int(cfg.max_duration)) / 2.0)) + nodes_mid = int(round((int(cfg.min_nodes) + int(cfg.max_nodes)) / 2.0)) + cores_mid = int(round((int(cfg.min_cores) + int(cfg.max_cores)) / 2.0)) + + if int(cfg.min_duration) > int(cfg.max_duration): + raise ValueError("min_duration must be <= max_duration") + if int(cfg.min_nodes) > int(cfg.max_nodes): + raise ValueError("min_nodes must be <= max_nodes") + if int(cfg.min_cores) > int(cfg.max_cores): + raise ValueError("min_cores must be <= max_cores") + if int(cfg.uniform_min_new_jobs_per_hour) > int(cfg.max_new_jobs_per_hour): + raise ValueError("uniform_min_new_jobs_per_hour must be <= max_new_jobs_per_hour") + + self.cfg = replace( + cfg, + arrivals=arrivals, + poisson_lambda_duration=( + float(cfg.poisson_lambda_duration) + if cfg.poisson_lambda_duration is not None + else float(duration_mid) + ), + poisson_lambda_nodes=( + float(cfg.poisson_lambda_nodes) + if cfg.poisson_lambda_nodes is not None + else float(nodes_mid) + ), + poisson_lambda_cores=( + float(cfg.poisson_lambda_cores) + if cfg.poisson_lambda_cores is not None + else float(cores_mid) + ), + flat_duration_target=( + int(cfg.flat_duration_target) + if cfg.flat_duration_target is not None + else int(duration_mid) + ), + flat_nodes_target=( + int(cfg.flat_nodes_target) + if cfg.flat_nodes_target is not None + else int(nodes_mid) + ), + flat_cores_target=( + int(cfg.flat_cores_target) + if cfg.flat_cores_target is not None + else int(cores_mid) + ), + ) + + def _sample_attr_array( + self, + rng: np.random.Generator, + size: int, + mode: str, + min_value: int, + max_value: int, + poisson_lambda: float, + flat_target: int, + flat_jitter: int, + ) -> np.ndarray: + if size <= 0: + return np.array([], dtype=np.int32) + + if mode == "flat": + if flat_jitter <= 0: + values = np.full(size, int(flat_target), dtype=np.int64) + else: + values = rng.integers( + int(flat_target) - int(flat_jitter), + int(flat_target) + int(flat_jitter) + 1, + size=size, + ) + elif mode == "poisson": + values = rng.poisson(float(poisson_lambda), size=size) + elif mode == "uniform": + values = rng.integers(int(min_value), int(max_value) + 1, size=size) + else: + raise ValueError(f"Unknown sampling mode: {mode}") + + return np.clip(values, int(min_value), int(max_value)).astype(np.int32) def _sample_job_count(self, rng: np.random.Generator) -> int: """ Arrival modes: - flat: constant arrivals around a target, optional +/- jitter (0 => perfectly constant) - poisson: Poisson(lambda) - - uniform: discrete-uniform in [0, max_new_jobs_per_hour] (very noisy hour-to-hour) + - uniform: discrete-uniform in [uniform_min_new_jobs_per_hour, max_new_jobs_per_hour] """ mode = self.cfg.arrivals @@ -82,8 +171,12 @@ def _sample_job_count(self, rng: np.random.Generator) -> int: k = int(rng.poisson(self.cfg.poisson_lambda)) elif mode == "uniform": - # This is the old "flat". - k = int(rng.integers(0, self.cfg.max_new_jobs_per_hour + 1)) + k = int( + rng.integers( + int(self.cfg.uniform_min_new_jobs_per_hour), + int(self.cfg.max_new_jobs_per_hour) + 1, + ) + ) else: raise ValueError(f"Unknown arrivals mode: {mode}") @@ -103,8 +196,36 @@ def sample(self, hour_idx: int, rng: np.random.Generator) -> List[JobSpec]: if n == 0: return [] - durations = rng.integers(self.cfg.min_duration, self.cfg.max_duration + 1, size=n, dtype=np.int32) - nodes = rng.integers(self.cfg.min_nodes, self.cfg.max_nodes + 1, size=n, dtype=np.int32) - cores = rng.integers(self.cfg.min_cores, self.cfg.max_cores + 1, size=n, dtype=np.int32) + mode = self.cfg.arrivals + durations = self._sample_attr_array( + rng=rng, + size=n, + mode=mode, + min_value=int(self.cfg.min_duration), + max_value=int(self.cfg.max_duration), + poisson_lambda=float(self.cfg.poisson_lambda_duration), + flat_target=int(self.cfg.flat_duration_target), + flat_jitter=int(self.cfg.flat_duration_jitter), + ) + nodes = self._sample_attr_array( + rng=rng, + size=n, + mode=mode, + min_value=int(self.cfg.min_nodes), + max_value=int(self.cfg.max_nodes), + poisson_lambda=float(self.cfg.poisson_lambda_nodes), + flat_target=int(self.cfg.flat_nodes_target), + flat_jitter=int(self.cfg.flat_nodes_jitter), + ) + cores = self._sample_attr_array( + rng=rng, + size=n, + mode=mode, + min_value=int(self.cfg.min_cores), + max_value=int(self.cfg.max_cores), + poisson_lambda=float(self.cfg.poisson_lambda_cores), + flat_target=int(self.cfg.flat_cores_target), + flat_jitter=int(self.cfg.flat_cores_jitter), + ) return [JobSpec(int(durations[i]), int(nodes[i]), int(cores[i])) for i in range(n)] diff --git a/test/test_inspect_workloadgen.py b/test/test_inspect_workloadgen.py index 9ec8a84..d4b3db0 100644 --- a/test/test_inspect_workloadgen.py +++ b/test/test_inspect_workloadgen.py @@ -1,6 +1,6 @@ """ Run with: -python inspect_workloadgen.py --arrivals poisson --poisson-lambda 200 --max-jobs-hour 1500 --hours 336 --plot +python -m test.test_inspect_workloadgen --arrivals poisson --poisson-lambdas4 200,80,6,24 --max-jobs-hour 1500 --hours 336 --plot """ # inspect_workloadgen.py @@ -15,6 +15,52 @@ from src.workloadgen import WorkloadGenConfig, WorkloadGenerator +def parse_quad_floats(raw: str): + parts = [p.strip() for p in str(raw).split(",")] + if len(parts) != 4: + raise argparse.ArgumentTypeError( + "Expected 4 comma-separated floats: arrivals,duration,nodes,cores" + ) + try: + return tuple(float(p) for p in parts) + except ValueError as exc: + raise argparse.ArgumentTypeError(f"Invalid float in '{raw}'") from exc + + +def parse_quad_ints(raw: str): + parts = [p.strip() for p in str(raw).split(",")] + if len(parts) != 4: + raise argparse.ArgumentTypeError( + "Expected 4 comma-separated ints: arrivals,duration,nodes,cores" + ) + try: + return tuple(int(p) for p in parts) + except ValueError as exc: + raise argparse.ArgumentTypeError(f"Invalid int in '{raw}'") from exc + + +def parse_quad_ranges(raw: str): + parts = [p.strip() for p in str(raw).split(",")] + if len(parts) != 4: + raise argparse.ArgumentTypeError( + "Expected 4 comma-separated ranges: a_min:a_max,d_min:d_max,n_min:n_max,c_min:c_max" + ) + ranges = [] + for part in parts: + bounds = [b.strip() for b in part.split(":")] + if len(bounds) != 2: + raise argparse.ArgumentTypeError(f"Invalid range '{part}', expected min:max") + try: + low = int(bounds[0]) + high = int(bounds[1]) + except ValueError as exc: + raise argparse.ArgumentTypeError(f"Invalid int in range '{part}'") from exc + if low > high: + raise argparse.ArgumentTypeError(f"Range min > max in '{part}'") + ranges.append((low, high)) + return tuple(ranges) + + def digest_jobs_triplets(triplets): """ Stable digest to verify determinism. @@ -43,20 +89,82 @@ def main(): ap.add_argument("--seed", type=int, default=123) ap.add_argument("--hours", type=int, default=24 * 14) ap.add_argument("--arrivals", choices=["flat", "poisson", "uniform"], default="poisson") - ap.add_argument("--poisson-lambda", type=float, default=200.0) + ap.add_argument("--poisson-lambda", type=float, default=200.0, help="Legacy: arrivals-only poisson lambda.") + ap.add_argument("--poisson-lambdas4", type=parse_quad_floats, default=None, help="arrivals,duration,nodes,cores") ap.add_argument("--max-jobs-hour", type=int, default=1500) ap.add_argument("--plot", action="store_true") # Flat params (true flat with optional jitter) - ap.add_argument("--flat-jobs-hour",type=int,default=200,help="Target jobs/hour for arrivals=flat.") - ap.add_argument("--flat-jitter",type=int,default=0,help="Jitter for arrivals=flat. Sample in [target-jitter, target+jitter]. 0 => perfectly flat.") + ap.add_argument("--flat-jobs-hour", type=int, default=200, help="Legacy: arrivals-only flat target.") + ap.add_argument("--flat-jitter", type=int, default=0, help="Legacy: arrivals-only flat jitter.") + ap.add_argument("--flat-targets4", type=parse_quad_ints, default=None, help="arrivals,duration,nodes,cores") + ap.add_argument("--flat-jitters4", type=parse_quad_ints, default=None, help="arrivals,duration,nodes,cores") + ap.add_argument( + "--uniform-ranges4", + type=parse_quad_ranges, + default=None, + help="a_min:a_max,d_min:d_max,n_min:n_max,c_min:c_max", + ) args = ap.parse_args() + # Default ranges (used for uniform and for clipping in all modes). + uniform_min_jobs = 0 + max_jobs_hour = int(args.max_jobs_hour) + min_duration, max_duration = 1, 170 + min_nodes, max_nodes = 1, 16 + min_cores, max_cores = 1, 96 + if args.uniform_ranges4 is not None: + (uniform_min_jobs, max_jobs_hour), (min_duration, max_duration), (min_nodes, max_nodes), (min_cores, max_cores) = args.uniform_ranges4 + + default_duration_mid = int(round((min_duration + max_duration) / 2.0)) + default_nodes_mid = int(round((min_nodes + max_nodes) / 2.0)) + default_cores_mid = int(round((min_cores + max_cores) / 2.0)) + + if args.poisson_lambdas4 is not None: + poisson_lambda_arrivals, poisson_lambda_duration, poisson_lambda_nodes, poisson_lambda_cores = args.poisson_lambdas4 + else: + poisson_lambda_arrivals = float(args.poisson_lambda) + poisson_lambda_duration = float(default_duration_mid) + poisson_lambda_nodes = float(default_nodes_mid) + poisson_lambda_cores = float(default_cores_mid) + + if args.flat_targets4 is not None: + flat_jobs_per_hour, flat_duration_target, flat_nodes_target, flat_cores_target = args.flat_targets4 + else: + flat_jobs_per_hour = int(args.flat_jobs_hour) + flat_duration_target = default_duration_mid + flat_nodes_target = default_nodes_mid + flat_cores_target = default_cores_mid + + if args.flat_jitters4 is not None: + flat_jitter_arrivals, flat_duration_jitter, flat_nodes_jitter, flat_cores_jitter = args.flat_jitters4 + else: + flat_jitter_arrivals = int(args.flat_jitter) + flat_duration_jitter = 0 + flat_nodes_jitter = 0 + flat_cores_jitter = 0 + cfg = WorkloadGenConfig( arrivals=args.arrivals, - poisson_lambda=args.poisson_lambda, - max_new_jobs_per_hour=args.max_jobs_hour, - flat_jobs_per_hour=args.flat_jobs_hour, - flat_jitter=args.flat_jitter, + uniform_min_new_jobs_per_hour=uniform_min_jobs, + max_new_jobs_per_hour=max_jobs_hour, + poisson_lambda=poisson_lambda_arrivals, + poisson_lambda_duration=poisson_lambda_duration, + poisson_lambda_nodes=poisson_lambda_nodes, + poisson_lambda_cores=poisson_lambda_cores, + flat_jobs_per_hour=flat_jobs_per_hour, + flat_jitter=flat_jitter_arrivals, + flat_duration_target=flat_duration_target, + flat_nodes_target=flat_nodes_target, + flat_cores_target=flat_cores_target, + flat_duration_jitter=flat_duration_jitter, + flat_nodes_jitter=flat_nodes_jitter, + flat_cores_jitter=flat_cores_jitter, + min_duration=min_duration, + max_duration=max_duration, + min_nodes=min_nodes, + max_nodes=max_nodes, + min_cores=min_cores, + max_cores=max_cores, ) gen = WorkloadGenerator(cfg) @@ -168,7 +276,7 @@ def main(): plt.show() timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") - prefix = f"{args.arrivals}_lambda{args.poisson_lambda}" if args.arrivals == "poisson" else args.arrivals + prefix = f"{args.arrivals}_lambda{poisson_lambda_arrivals}" if args.arrivals == "poisson" else args.arrivals fname = f"{prefix}_{timestamp}.png" if prefix else f"Workload-Gen_{timestamp}.png" save_path = os.path.join("", fname) plt.savefig(save_path, dpi=250, bbox_inches="tight") diff --git a/test/test_sanity_env.py b/test/test_sanity_env.py index 2509ab2..54f7f60 100644 --- a/test/test_sanity_env.py +++ b/test/test_sanity_env.py @@ -37,6 +37,52 @@ def load_prices(prices_file_path: str | None): print(f"Loaded {len(prices)} prices from CSV: {prices_file_path}") return prices + +def _parse_quad_floats(raw: str): + parts = [p.strip() for p in str(raw).split(",")] + if len(parts) != 4: + raise argparse.ArgumentTypeError( + "Expected 4 comma-separated floats: arrivals,duration,nodes,cores" + ) + try: + return tuple(float(p) for p in parts) + except ValueError as exc: + raise argparse.ArgumentTypeError(f"Invalid float in '{raw}'") from exc + + +def _parse_quad_ints(raw: str): + parts = [p.strip() for p in str(raw).split(",")] + if len(parts) != 4: + raise argparse.ArgumentTypeError( + "Expected 4 comma-separated ints: arrivals,duration,nodes,cores" + ) + try: + return tuple(int(p) for p in parts) + except ValueError as exc: + raise argparse.ArgumentTypeError(f"Invalid int in '{raw}'") from exc + + +def _parse_quad_ranges(raw: str): + parts = [p.strip() for p in str(raw).split(",")] + if len(parts) != 4: + raise argparse.ArgumentTypeError( + "Expected 4 comma-separated ranges: a_min:a_max,d_min:d_max,n_min:n_max,c_min:c_max" + ) + ranges = [] + for part in parts: + bounds = [b.strip() for b in part.split(":")] + if len(bounds) != 2: + raise argparse.ArgumentTypeError(f"Invalid range '{part}', expected min:max") + try: + low = int(bounds[0]) + high = int(bounds[1]) + except ValueError as exc: + raise argparse.ArgumentTypeError(f"Invalid int in range '{part}'") from exc + if low > high: + raise argparse.ArgumentTypeError(f"Range min > max in '{part}'") + ranges.append((low, high)) + return tuple(ranges) + # ----------------------------- # Invariants / sanity checks # ----------------------------- @@ -221,9 +267,18 @@ def parse_args(): p.add_argument("--idle-weight", type=float, default=0.1) p.add_argument("--job-age-weight", type=float, default=0.1) p.add_argument("--drop-weight", type=float, default=0.1) - p.add_argument("--workload-gen",type=str,default="",choices=["", "flat", "poisson"],help="Enable workload generator (default: disabled).",) - p.add_argument("--wg-poisson-lambda", type=float, default=200.0, help="Poisson lambda for jobs/hour.") + p.add_argument("--workload-gen",type=str,default="",choices=["", "flat", "poisson", "uniform"],help="Enable workload generator (default: disabled).",) + p.add_argument("--wg-poisson-lambda", type=float, default=200.0, help="Legacy: arrivals-only poisson lambda.") + p.add_argument("--wg-poisson-lambdas4", type=_parse_quad_floats, default=None, help="arrivals,duration,nodes,cores") p.add_argument("--wg-max-jobs-hour", type=int, default=1500, help="Cap jobs/hour for generator.") + p.add_argument("--wg-flat-targets4", type=_parse_quad_ints, default=None, help="arrivals,duration,nodes,cores") + p.add_argument("--wg-flat-jitters4", type=_parse_quad_ints, default=None, help="arrivals,duration,nodes,cores") + p.add_argument( + "--wg-uniform-ranges4", + type=_parse_quad_ranges, + default=None, + help="a_min:a_max,d_min:d_max,n_min:n_max,c_min:c_max", + ) p.add_argument("--print-job-every", type=int, default=0, help="Print one sample job every N steps (0 disables).") p.add_argument("--print-job-kind", choices=["queue", "running", "both"], default="queue", 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.") @@ -245,16 +300,70 @@ def make_env_from_args(args, env_cls=ComputeClusterEnv): workload_gen = None if args.workload_gen: + uniform_min_jobs = 0 + max_jobs_hour = int(args.wg_max_jobs_hour) + min_duration, max_duration = 1, MAX_JOB_DURATION + min_nodes, max_nodes = MIN_NODES_PER_JOB, MAX_NODES_PER_JOB + min_cores, max_cores = MIN_CORES_PER_JOB, CORES_PER_NODE + + if args.wg_uniform_ranges4 is not None: + ( + (uniform_min_jobs, max_jobs_hour), + (min_duration, max_duration), + (min_nodes, max_nodes), + (min_cores, max_cores), + ) = args.wg_uniform_ranges4 + + duration_mid = int(round((min_duration + max_duration) / 2.0)) + nodes_mid = int(round((min_nodes + max_nodes) / 2.0)) + cores_mid = int(round((min_cores + max_cores) / 2.0)) + + if args.wg_poisson_lambdas4 is not None: + poisson_lambda_arrivals, poisson_lambda_duration, poisson_lambda_nodes, poisson_lambda_cores = args.wg_poisson_lambdas4 + else: + poisson_lambda_arrivals = float(args.wg_poisson_lambda) + poisson_lambda_duration = float(duration_mid) + poisson_lambda_nodes = float(nodes_mid) + poisson_lambda_cores = float(cores_mid) + + if args.wg_flat_targets4 is not None: + flat_jobs_per_hour, flat_duration_target, flat_nodes_target, flat_cores_target = args.wg_flat_targets4 + else: + flat_jobs_per_hour = 200 + flat_duration_target = duration_mid + flat_nodes_target = nodes_mid + flat_cores_target = cores_mid + + if args.wg_flat_jitters4 is not None: + flat_jitter_arrivals, flat_duration_jitter, flat_nodes_jitter, flat_cores_jitter = args.wg_flat_jitters4 + else: + flat_jitter_arrivals = 0 + flat_duration_jitter = 0 + flat_nodes_jitter = 0 + flat_cores_jitter = 0 + cfg = WorkloadGenConfig( arrivals=args.workload_gen, - poisson_lambda=args.wg_poisson_lambda, - max_new_jobs_per_hour=args.wg_max_jobs_hour, - min_duration=1, - max_duration=MAX_JOB_DURATION, - min_nodes=MIN_NODES_PER_JOB, - max_nodes=MAX_NODES_PER_JOB, - min_cores=MIN_CORES_PER_JOB, - max_cores=CORES_PER_NODE, + uniform_min_new_jobs_per_hour=uniform_min_jobs, + max_new_jobs_per_hour=max_jobs_hour, + poisson_lambda=poisson_lambda_arrivals, + poisson_lambda_duration=poisson_lambda_duration, + poisson_lambda_nodes=poisson_lambda_nodes, + poisson_lambda_cores=poisson_lambda_cores, + flat_jobs_per_hour=flat_jobs_per_hour, + flat_jitter=flat_jitter_arrivals, + flat_duration_target=flat_duration_target, + flat_nodes_target=flat_nodes_target, + flat_cores_target=flat_cores_target, + flat_duration_jitter=flat_duration_jitter, + flat_nodes_jitter=flat_nodes_jitter, + flat_cores_jitter=flat_cores_jitter, + min_duration=min_duration, + max_duration=max_duration, + min_nodes=min_nodes, + max_nodes=max_nodes, + min_cores=min_cores, + max_cores=max_cores, ) workload_gen = WorkloadGenerator(cfg) diff --git a/test/test_sanity_workloadgen.py b/test/test_sanity_workloadgen.py index 42d32ae..1f0b387 100644 --- a/test/test_sanity_workloadgen.py +++ b/test/test_sanity_workloadgen.py @@ -36,8 +36,67 @@ def test_poisson_mean_sanity(): mean = float(np.mean(counts)) assert 48.0 < mean < 52.0, mean # tighter band, still reliable for 2000 samples + +def test_flat_attribute_targets(): + cfg = WorkloadGenConfig( + arrivals="flat", + flat_jobs_per_hour=64, + flat_jitter=0, + flat_duration_target=12, + flat_nodes_target=3, + flat_cores_target=8, + flat_duration_jitter=0, + flat_nodes_jitter=0, + flat_cores_jitter=0, + min_duration=1, + max_duration=170, + min_nodes=1, + max_nodes=16, + min_cores=1, + max_cores=96, + ) + gen = WorkloadGenerator(cfg) + rng = np.random.default_rng(2) + jobs = gen.sample(0, rng) + + assert len(jobs) == 64 + assert all(j.duration == 12 for j in jobs) + assert all(j.nodes == 3 for j in jobs) + assert all(j.cores_per_node == 8 for j in jobs) + + +def test_poisson_attribute_lambdas_are_used(): + cfg = WorkloadGenConfig( + arrivals="poisson", + poisson_lambda=200.0, + poisson_lambda_duration=2.0, + poisson_lambda_nodes=2.0, + poisson_lambda_cores=2.0, + min_duration=1, + max_duration=170, + min_nodes=1, + max_nodes=16, + min_cores=1, + max_cores=96, + ) + gen = WorkloadGenerator(cfg) + rng = np.random.default_rng(3) + + jobs = gen.sample(0, rng) + assert len(jobs) > 0 + mean_duration = float(np.mean([j.duration for j in jobs])) + mean_nodes = float(np.mean([j.nodes for j in jobs])) + mean_cores = float(np.mean([j.cores_per_node for j in jobs])) + + # With lambda=2 and lower bound clipping at 1, means should remain low. + assert mean_duration < 6.0, mean_duration + assert mean_nodes < 6.0, mean_nodes + assert mean_cores < 6.0, mean_cores + if __name__ == "__main__": test_determinism() test_constraints() test_poisson_mean_sanity() + test_flat_attribute_targets() + test_poisson_attribute_lambdas_are_used() print("[OK] workloadgen sanity checks passed") diff --git a/train.py b/train.py index 4ae1e3b..3a5db3a 100644 --- a/train.py +++ b/train.py @@ -29,6 +29,53 @@ def norm_path(x): STEPS_PER_ITERATION = 100000 + +def _parse_quad_floats(raw: str): + parts = [p.strip() for p in str(raw).split(",")] + if len(parts) != 4: + raise argparse.ArgumentTypeError( + "Expected 4 comma-separated floats: arrivals,duration,nodes,cores" + ) + try: + return tuple(float(p) for p in parts) + except ValueError as exc: + raise argparse.ArgumentTypeError(f"Invalid float in '{raw}'") from exc + + +def _parse_quad_ints(raw: str): + parts = [p.strip() for p in str(raw).split(",")] + if len(parts) != 4: + raise argparse.ArgumentTypeError( + "Expected 4 comma-separated ints: arrivals,duration,nodes,cores" + ) + try: + return tuple(int(p) for p in parts) + except ValueError as exc: + raise argparse.ArgumentTypeError(f"Invalid int in '{raw}'") from exc + + +def _parse_quad_ranges(raw: str): + parts = [p.strip() for p in str(raw).split(",")] + if len(parts) != 4: + raise argparse.ArgumentTypeError( + "Expected 4 comma-separated ranges: a_min:a_max,d_min:d_max,n_min:n_max,c_min:c_max" + ) + ranges = [] + for part in parts: + bounds = [b.strip() for b in part.split(":")] + if len(bounds) != 2: + raise argparse.ArgumentTypeError(f"Invalid range '{part}', expected min:max") + try: + low = int(bounds[0]) + high = int(bounds[1]) + except ValueError as exc: + raise argparse.ArgumentTypeError(f"Invalid int in range '{part}'") from exc + if low > high: + raise argparse.ArgumentTypeError(f"Range min > max in '{part}'") + ranges.append((low, high)) + return tuple(ranges) + + def main(): parser = argparse.ArgumentParser(description="Run the Compute Cluster Environment with optional rendering.") parser.add_argument('--render', type=str, default='none', choices=['human', 'none'], help='Render mode for the environment (default: none).') @@ -59,8 +106,17 @@ def main(): parser.add_argument("--evaluate-savings", action='store_true', help="Load latest model and evaluate long-term savings (no training)") parser.add_argument("--eval-months", type=int, default=12, help="Months to evaluate for savings analysis (default: 12, only used with --evaluate-savings)") parser.add_argument("--workload-gen", type=str, default="", choices=["", "flat", "poisson", "uniform"], help="Enable workload generator (default: disabled).",) - parser.add_argument("--wg-poisson-lambda", type=float, default=200.0, help="Poisson lambda for jobs/hour for the workload generator.") + parser.add_argument("--wg-poisson-lambda", type=float, default=200.0, help="Legacy: arrivals-only poisson lambda for workload generator.") + parser.add_argument("--wg-poisson-lambdas4", type=_parse_quad_floats, default=None, help="arrivals,duration,nodes,cores") parser.add_argument("--wg-max-jobs-hour", type=int, default=1500, help="Cap jobs/hour for the workload generator.") + parser.add_argument("--wg-flat-targets4", type=_parse_quad_ints, default=None, help="arrivals,duration,nodes,cores") + parser.add_argument("--wg-flat-jitters4", type=_parse_quad_ints, default=None, help="arrivals,duration,nodes,cores") + parser.add_argument( + "--wg-uniform-ranges4", + type=_parse_quad_ranges, + default=None, + help="a_min:a_max,d_min:d_max,n_min:n_max,c_min:c_max", + ) 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).") @@ -117,16 +173,70 @@ def main(): workload_gen = None if args.workload_gen: + uniform_min_jobs = 0 + max_jobs_hour = int(args.wg_max_jobs_hour) + min_duration, max_duration = 1, MAX_JOB_DURATION + min_nodes, max_nodes = MIN_NODES_PER_JOB, MAX_NODES_PER_JOB + min_cores, max_cores = MIN_CORES_PER_JOB, CORES_PER_NODE + + if args.wg_uniform_ranges4 is not None: + ( + (uniform_min_jobs, max_jobs_hour), + (min_duration, max_duration), + (min_nodes, max_nodes), + (min_cores, max_cores), + ) = args.wg_uniform_ranges4 + + duration_mid = int(round((min_duration + max_duration) / 2.0)) + nodes_mid = int(round((min_nodes + max_nodes) / 2.0)) + cores_mid = int(round((min_cores + max_cores) / 2.0)) + + if args.wg_poisson_lambdas4 is not None: + poisson_lambda_arrivals, poisson_lambda_duration, poisson_lambda_nodes, poisson_lambda_cores = args.wg_poisson_lambdas4 + else: + poisson_lambda_arrivals = float(args.wg_poisson_lambda) + poisson_lambda_duration = float(duration_mid) + poisson_lambda_nodes = float(nodes_mid) + poisson_lambda_cores = float(cores_mid) + + if args.wg_flat_targets4 is not None: + flat_jobs_per_hour, flat_duration_target, flat_nodes_target, flat_cores_target = args.wg_flat_targets4 + else: + flat_jobs_per_hour = 200 + flat_duration_target = duration_mid + flat_nodes_target = nodes_mid + flat_cores_target = cores_mid + + if args.wg_flat_jitters4 is not None: + flat_jitter_arrivals, flat_duration_jitter, flat_nodes_jitter, flat_cores_jitter = args.wg_flat_jitters4 + else: + flat_jitter_arrivals = 0 + flat_duration_jitter = 0 + flat_nodes_jitter = 0 + flat_cores_jitter = 0 + cfg = WorkloadGenConfig( arrivals=args.workload_gen, - poisson_lambda=args.wg_poisson_lambda, - max_new_jobs_per_hour=args.wg_max_jobs_hour, - min_duration=1, - max_duration=MAX_JOB_DURATION, - min_nodes=MIN_NODES_PER_JOB, - max_nodes=MAX_NODES_PER_JOB, - min_cores=MIN_CORES_PER_JOB, - max_cores=CORES_PER_NODE, + uniform_min_new_jobs_per_hour=uniform_min_jobs, + max_new_jobs_per_hour=max_jobs_hour, + poisson_lambda=poisson_lambda_arrivals, + poisson_lambda_duration=poisson_lambda_duration, + poisson_lambda_nodes=poisson_lambda_nodes, + poisson_lambda_cores=poisson_lambda_cores, + flat_jobs_per_hour=flat_jobs_per_hour, + flat_jitter=flat_jitter_arrivals, + flat_duration_target=flat_duration_target, + flat_nodes_target=flat_nodes_target, + flat_cores_target=flat_cores_target, + flat_duration_jitter=flat_duration_jitter, + flat_nodes_jitter=flat_nodes_jitter, + flat_cores_jitter=flat_cores_jitter, + min_duration=min_duration, + max_duration=max_duration, + min_nodes=min_nodes, + max_nodes=max_nodes, + min_cores=min_cores, + max_cores=max_cores, ) workload_gen = WorkloadGenerator(cfg) From addcbde1d31e2a0f954b2b5528a2716180a038dd Mon Sep 17 00:00:00 2001 From: Enis Lorenz Date: Mon, 16 Feb 2026 15:01:22 +0100 Subject: [PATCH 3/9] Feature Workloadgen: add additive burst modes for small and heavy job spikes Optional body: Add two independent burst injectors on top of base arrivals: burst_small_* for many small jobs burst_heavy_* for long, high-resource jobs Add per-burst probability controls and default ranges. Wire burst probabilities through train/sanity/inspect CLI and add burst behavior tests. Small Add: Workload_logs.txt, collective stats of different partition logs. Can be used to determine parameters of generator. Generator: Add workload_logs script, plus added "burstiness" scale to logs. Fixup: Moved workload logs, into data/workload_statistics Fixed "fallback crash on formats without exactly three colon-separated parts" --- .../analyze_workload_logs.py | 447 ++++++++++++++++++ data/workload_statistics/workload_logs.txt | 62 +++ src/workloadgen.py | 221 +++++++-- test/run_all.py | 1 + test/test_inspect_workloadgen.py | 63 +-- test/test_sanity_env.py | 51 +- test/test_sanity_workloadgen.py | 110 +++++ train.py | 4 + 8 files changed, 827 insertions(+), 132 deletions(-) create mode 100644 data/workload_statistics/analyze_workload_logs.py create mode 100644 data/workload_statistics/workload_logs.txt diff --git a/data/workload_statistics/analyze_workload_logs.py b/data/workload_statistics/analyze_workload_logs.py new file mode 100644 index 0000000..832d8fd --- /dev/null +++ b/data/workload_statistics/analyze_workload_logs.py @@ -0,0 +1,447 @@ +#!/usr/bin/env python3 +"""Summarize workload logs per file and estimate burst probabilities. + +For each input file, compute: +- arrivals per hour: mean, stddev +- duration (hours): mean, stddev +- nodes: mean, stddev +- cores: mean, stddev +- Pearson correlations: + - duration vs nodes + - duration vs cores +- burst probability suggestions for: + - wg-burst-small-prob + - wg-burst-heavy-prob + +The script supports: +- whitespace-delimited Slurm-like logs (as in data-internal/allusers-*.log) +- standard CSV files with matching column names +""" + +from __future__ import annotations + +import argparse +import csv +import math +from collections import Counter +from datetime import datetime, timedelta +from pathlib import Path +from statistics import fmean, stdev +from typing import Dict, Iterable, List, Sequence, Tuple + + +SUBMIT_CANDIDATES = ("Submit", "submit", "SUBMIT", "submission_time", "timestamp") +DURATION_CANDIDATES = ("ElapsedRaw", "elapsed_raw", "ELAPSEDRAW", "duration_seconds", "duration") +NODES_CANDIDATES = ("NNodes", "nnodes", "NNODES", "nodes") +CORES_CANDIDATES = ("NCPUS", "ncpus", "NCPUs", "cores", "cores_per_node") + + +def pick_column(columns: Sequence[str], candidates: Sequence[str]) -> str: + for candidate in candidates: + if candidate in columns: + return candidate + raise ValueError(f"Could not find any of {candidates} in columns: {columns}") + + +def parse_submit_hour(value: str) -> datetime: + dt = datetime.fromisoformat(value.strip()) + return dt.replace(minute=0, second=0, microsecond=0) + + +def parse_duration_hours(value: str) -> float: + raw = value.strip() + # Common case in these logs: ElapsedRaw is integer seconds. + try: + seconds = float(raw) + return seconds / 3600.0 + except ValueError: + pass + + # Fallback: parse HH:MM:SS or D-HH:MM:SS + day_part = 0 + time_part = raw + if "-" in raw: + maybe_day, maybe_time = raw.split("-", 1) + if maybe_day.isdigit(): + day_part = int(maybe_day) + time_part = maybe_time + parts = time_part.split(":") + if len(parts) == 3: + hh, mm, ss = parts + elif len(parts) == 2: + hh, mm = parts + ss = "0" + else: + raise ValueError(f"Cannot parse duration: {raw!r}") + total_seconds = (day_part * 24 + int(hh)) * 3600 + int(mm) * 60 + int(ss) + return total_seconds / 3600.0 + + +def pearson_corr(x: Sequence[float], y: Sequence[float]) -> float: + if len(x) != len(y) or len(x) < 2: + return float("nan") + mx = fmean(x) + my = fmean(y) + dx = [v - mx for v in x] + dy = [v - my for v in y] + sx = math.sqrt(sum(v * v for v in dx)) + sy = math.sqrt(sum(v * v for v in dy)) + if sx == 0.0 or sy == 0.0: + return float("nan") + return sum(a * b for a, b in zip(dx, dy)) / (sx * sy) + + +def mean_std(values: Sequence[float]) -> Tuple[float, float]: + if not values: + return float("nan"), float("nan") + if len(values) == 1: + return float(values[0]), 0.0 + return fmean(values), stdev(values) + + +def quantile(values: Sequence[float], q: float) -> float: + if not values: + return float("nan") + q = min(max(float(q), 0.0), 1.0) + s = sorted(float(v) for v in values) + if len(s) == 1: + return s[0] + pos = q * (len(s) - 1) + lo = int(math.floor(pos)) + hi = int(math.ceil(pos)) + if lo == hi: + return s[lo] + frac = pos - lo + return s[lo] * (1.0 - frac) + s[hi] * frac + + +def parse_min_max(raw: str) -> Tuple[int, int]: + parts = [p.strip() for p in str(raw).split(":")] + if len(parts) != 2: + raise argparse.ArgumentTypeError("Expected min:max") + try: + lo = int(parts[0]) + hi = int(parts[1]) + except ValueError as exc: + raise argparse.ArgumentTypeError(f"Invalid min:max pair '{raw}'") from exc + if lo > hi: + raise argparse.ArgumentTypeError(f"min > max in '{raw}'") + return lo, hi + + +def hourly_axis(arrivals_by_hour: Counter, include_zero_hours: bool) -> List[datetime]: + if not arrivals_by_hour: + return [] + if not include_zero_hours: + return sorted(arrivals_by_hour.keys()) + + start = min(arrivals_by_hour) + end = max(arrivals_by_hour) + current = start + values: List[datetime] = [] + while current <= end: + values.append(current) + current += timedelta(hours=1) + return values + + +def iter_rows(path: Path) -> Iterable[Dict[str, str]]: + with path.open("r", encoding="utf-8", errors="replace") as fh: + first_nonempty = None + for line in fh: + if line.strip(): + first_nonempty = line + break + if first_nonempty is None: + return + + if "," in first_nonempty: + fh.seek(0) + reader = csv.DictReader(fh) + for row in reader: + if not row: + continue + yield row + return + + # Whitespace-delimited format with a dashed separator on line 2. + columns = first_nonempty.split() + for line in fh: + stripped = line.strip() + if not stripped: + continue + if set(stripped) <= {"-"}: + continue + parts = line.split() + if len(parts) < len(columns): + continue + row = {columns[i]: parts[i] for i in range(len(columns))} + yield row + + +def summarize_file( + path: Path, + include_zero_hours: bool, + small_duration_max: float, + small_nodes_max: float, + small_cores_max: float, + heavy_duration_min: float, + heavy_nodes_min: float, + heavy_cores_min: float, + assumed_small_jobs_min: int, + assumed_small_jobs_max: int, + assumed_heavy_jobs_min: int, + assumed_heavy_jobs_max: int, + baseline_quantile: float, +) -> Dict[str, float]: + arrivals_by_hour: Counter = Counter() + small_jobs_by_hour: Counter = Counter() + heavy_jobs_by_hour: Counter = Counter() + durations_h: List[float] = [] + nodes: List[float] = [] + cores: List[float] = [] + skipped = 0 + + rows = iter_rows(path) + rows = iter(rows) + try: + first = next(rows) + except StopIteration: + return { + "jobs": 0, + "skipped_rows": 0, + "arrival_mean": float("nan"), + "arrival_std": float("nan"), + "duration_mean_h": float("nan"), + "duration_std_h": float("nan"), + "nodes_mean": float("nan"), + "nodes_std": float("nan"), + "cores_mean": float("nan"), + "cores_std": float("nan"), + "corr_duration_nodes": float("nan"), + "corr_duration_cores": float("nan"), + } + + # Determine column mapping from the first row, then process all rows (including first). + keys = list(first.keys()) + submit_col = pick_column(keys, SUBMIT_CANDIDATES) + duration_col = pick_column(keys, DURATION_CANDIDATES) + nodes_col = pick_column(keys, NODES_CANDIDATES) + cores_col = pick_column(keys, CORES_CANDIDATES) + + def consume(row: Dict[str, str]) -> None: + nonlocal skipped + try: + hour = parse_submit_hour(row[submit_col]) + duration = parse_duration_hours(row[duration_col]) + node_count = float(row[nodes_col]) + core_count = float(row[cores_col]) + except Exception: + skipped += 1 + return + arrivals_by_hour[hour] += 1 + if duration <= small_duration_max and node_count <= small_nodes_max and core_count <= small_cores_max: + small_jobs_by_hour[hour] += 1 + if duration >= heavy_duration_min and node_count >= heavy_nodes_min and core_count >= heavy_cores_min: + heavy_jobs_by_hour[hour] += 1 + durations_h.append(duration) + nodes.append(node_count) + cores.append(core_count) + + consume(first) + for row in rows: + consume(row) + + axis = hourly_axis(arrivals_by_hour, include_zero_hours) + arrival_series = [float(arrivals_by_hour.get(h, 0)) for h in axis] + small_series = [float(small_jobs_by_hour.get(h, 0)) for h in axis] + heavy_series = [float(heavy_jobs_by_hour.get(h, 0)) for h in axis] + + arrival_mean, arrival_std = mean_std(arrival_series) + duration_mean_h, duration_std_h = mean_std(durations_h) + nodes_mean, nodes_std = mean_std(nodes) + cores_mean, cores_std = mean_std(cores) + + n_hours = len(axis) + active_idx = [i for i, a in enumerate(arrival_series) if a > 0.0] + if active_idx: + active_small = [small_series[i] for i in active_idx] + active_heavy = [heavy_series[i] for i in active_idx] + else: + active_small = small_series + active_heavy = heavy_series + + # Learn baseline from active hours to avoid zero-heavy timelines collapsing + # burst baselines to 0. + small_baseline = quantile(active_small, baseline_quantile) + heavy_baseline = quantile(active_heavy, baseline_quantile) + + small_event_threshold = small_baseline + float(assumed_small_jobs_min) + heavy_event_threshold = heavy_baseline + float(assumed_heavy_jobs_min) + + if n_hours > 0: + small_event_prob = sum(1 for v in small_series if v >= small_event_threshold and v > small_baseline) / n_hours + heavy_event_prob = sum(1 for v in heavy_series if v >= heavy_event_threshold and v > heavy_baseline) / n_hours + else: + small_event_prob = float("nan") + heavy_event_prob = float("nan") + + small_expected_jobs_per_burst = (float(assumed_small_jobs_min) + float(assumed_small_jobs_max)) / 2.0 + heavy_expected_jobs_per_burst = (float(assumed_heavy_jobs_min) + float(assumed_heavy_jobs_max)) / 2.0 + + small_excess = sum(max(v - small_baseline, 0.0) for v in small_series) + heavy_excess = sum(max(v - heavy_baseline, 0.0) for v in heavy_series) + + if n_hours > 0 and small_expected_jobs_per_burst > 0.0: + small_volume_prob = min(max(small_excess / (n_hours * small_expected_jobs_per_burst), 0.0), 1.0) + else: + small_volume_prob = float("nan") + if n_hours > 0 and heavy_expected_jobs_per_burst > 0.0: + heavy_volume_prob = min(max(heavy_excess / (n_hours * heavy_expected_jobs_per_burst), 0.0), 1.0) + else: + heavy_volume_prob = float("nan") + + # wg-burst-*-prob are event probabilities, so event-rate estimates are the + # primary recommendation. Volume estimates are kept for diagnostics. + if math.isfinite(small_event_prob): + suggested_small_prob = min(max(small_event_prob, 0.0), 1.0) + elif math.isfinite(small_volume_prob): + suggested_small_prob = min(max(small_volume_prob, 0.0), 1.0) + else: + suggested_small_prob = float("nan") + + if math.isfinite(heavy_event_prob): + suggested_heavy_prob = min(max(heavy_event_prob, 0.0), 1.0) + elif math.isfinite(heavy_volume_prob): + suggested_heavy_prob = min(max(heavy_volume_prob, 0.0), 1.0) + else: + suggested_heavy_prob = float("nan") + + return { + "jobs": len(durations_h), + "skipped_rows": skipped, + "arrival_mean": arrival_mean, + "arrival_std": arrival_std, + "duration_mean_h": duration_mean_h, + "duration_std_h": duration_std_h, + "nodes_mean": nodes_mean, + "nodes_std": nodes_std, + "cores_mean": cores_mean, + "cores_std": cores_std, + "corr_duration_nodes": pearson_corr(durations_h, nodes), + "corr_duration_cores": pearson_corr(durations_h, cores), + "hours_observed": n_hours, + "small_baseline_jobs_per_hour": small_baseline, + "heavy_baseline_jobs_per_hour": heavy_baseline, + "small_event_prob": small_event_prob, + "heavy_event_prob": heavy_event_prob, + "small_volume_prob": small_volume_prob, + "heavy_volume_prob": heavy_volume_prob, + "suggested_wg_burst_small_prob": suggested_small_prob, + "suggested_wg_burst_heavy_prob": suggested_heavy_prob, + } + + +def main() -> None: + parser = argparse.ArgumentParser(description="Summarize workload logs by file.") + parser.add_argument( + "files", + nargs="*", + help="Input files (CSV or whitespace logs). Default: data-internal/allusers-*.log", + ) + parser.add_argument( + "--include-zero-hours", + action=argparse.BooleanOptionalAction, + default=True, + help="Include zero-arrival hours between first and last submission time (default: true).", + ) + parser.add_argument("--small-duration-max", type=float, default=8.0, help="Small-job duration threshold (<=).") + parser.add_argument("--small-nodes-max", type=float, default=2.0, help="Small-job nodes threshold (<=).") + parser.add_argument("--small-cores-max", type=float, default=16.0, help="Small-job cores threshold (<=).") + parser.add_argument("--heavy-duration-min", type=float, default=72.0, help="Heavy-job duration threshold (>=).") + parser.add_argument("--heavy-nodes-min", type=float, default=4.0, help="Heavy-job nodes threshold (>=).") + parser.add_argument("--heavy-cores-min", type=float, default=32.0, help="Heavy-job cores threshold (>=).") + parser.add_argument( + "--assumed-burst-small-jobs", + type=parse_min_max, + default=(50, 250), + help="Assumed small burst size range as min:max (default: 50:250).", + ) + parser.add_argument( + "--assumed-burst-heavy-jobs", + type=parse_min_max, + default=(1, 12), + help="Assumed heavy burst size range as min:max (default: 1:12).", + ) + parser.add_argument( + "--baseline-quantile", + type=float, + default=0.5, + help="Quantile used as non-burst baseline for hourly small/heavy counts (default: 0.5).", + ) + args = parser.parse_args() + + if args.files: + files = [Path(p) for p in args.files] + else: + files = sorted(Path("data-internal").glob("allusers-*.log")) + + if not files: + raise SystemExit("No input files found.") + + for path in files: + stats = summarize_file( + path, + include_zero_hours=args.include_zero_hours, + small_duration_max=float(args.small_duration_max), + small_nodes_max=float(args.small_nodes_max), + small_cores_max=float(args.small_cores_max), + heavy_duration_min=float(args.heavy_duration_min), + heavy_nodes_min=float(args.heavy_nodes_min), + heavy_cores_min=float(args.heavy_cores_min), + assumed_small_jobs_min=int(args.assumed_burst_small_jobs[0]), + assumed_small_jobs_max=int(args.assumed_burst_small_jobs[1]), + assumed_heavy_jobs_min=int(args.assumed_burst_heavy_jobs[0]), + assumed_heavy_jobs_max=int(args.assumed_burst_heavy_jobs[1]), + baseline_quantile=float(args.baseline_quantile), + ) + print(f"\n=== {path.name} ===") + print(f"jobs={stats['jobs']} skipped_rows={stats['skipped_rows']}") + print(f"hours_observed={stats['hours_observed']}") + print( + "arrivals_per_hour: " + f"mean={stats['arrival_mean']:.4f} std={stats['arrival_std']:.4f}" + ) + print( + "duration_hours: " + f"mean={stats['duration_mean_h']:.4f} std={stats['duration_std_h']:.4f}" + ) + print(f"nodes: mean={stats['nodes_mean']:.4f} std={stats['nodes_std']:.4f}") + print(f"cores: mean={stats['cores_mean']:.4f} std={stats['cores_std']:.4f}") + print( + "corr(duration, nodes)=" + f"{stats['corr_duration_nodes']:.6f} " + "corr(duration, cores)=" + f"{stats['corr_duration_cores']:.6f}" + ) + print( + "burst_baseline_jobs_per_hour: " + f"small={stats['small_baseline_jobs_per_hour']:.4f} " + f"heavy={stats['heavy_baseline_jobs_per_hour']:.4f}" + ) + print( + "burst_prob_estimates: " + f"small_event={stats['small_event_prob']:.4f} " + f"small_volume={stats['small_volume_prob']:.4f} " + f"heavy_event={stats['heavy_event_prob']:.4f} " + f"heavy_volume={stats['heavy_volume_prob']:.4f}" + ) + print( + "suggested_flags: " + f"--wg-burst-small-prob {stats['suggested_wg_burst_small_prob']:.4f} " + f"--wg-burst-heavy-prob {stats['suggested_wg_burst_heavy_prob']:.4f}" + ) + + +if __name__ == "__main__": + main() diff --git a/data/workload_statistics/workload_logs.txt b/data/workload_statistics/workload_logs.txt new file mode 100644 index 0000000..230281c --- /dev/null +++ b/data/workload_statistics/workload_logs.txt @@ -0,0 +1,62 @@ +=== allusers-gpu-30.log === +jobs=3319 skipped_rows=1 +hours_observed=788 +arrivals_per_hour: mean=4.2119 std=77.5734 +duration_hours: mean=7.4852 std=29.6084 +nodes: mean=1.0054 std=0.1202 +cores: mean=14.9328 std=30.1301 +corr(duration, nodes)=0.007262 corr(duration, cores)=0.628780 +burst_baseline_jobs_per_hour: small=0.0000 heavy=0.0000 +burst_prob_estimates: small_event=0.0025 small_volume=0.0243 heavy_event=0.0000 heavy_volume=0.0000 +suggested_flags: --wg-burst-small-prob 0.0025 --wg-burst-heavy-prob 0.0000 + +=== allusers-grid-30.log === +jobs=40639 skipped_rows=1 +hours_observed=756 +arrivals_per_hour: mean=53.7553 std=87.2675 +duration_hours: mean=13.6808 std=7.1705 +nodes: mean=1.0000 std=0.0000 +cores: mean=7.9984 std=0.1122 +corr(duration, nodes)=nan corr(duration, cores)=0.026772 +burst_baseline_jobs_per_hour: small=0.0000 heavy=0.0000 +burst_prob_estimates: small_event=0.0423 small_volume=0.0752 heavy_event=0.0000 heavy_volume=0.0000 +suggested_flags: --wg-burst-small-prob 0.0423 --wg-burst-heavy-prob 0.0000 + +=== allusers-high_mem-30.log === +jobs=237373 skipped_rows=1 +hours_observed=1385 +arrivals_per_hour: mean=171.3884 std=611.6586 +duration_hours: mean=0.2537 std=1.8726 +nodes: mean=1.0010 std=0.1191 +cores: mean=13.6286 std=47.3460 +corr(duration, nodes)=0.405630 corr(duration, cores)=0.113031 +burst_baseline_jobs_per_hour: small=129.0000 heavy=0.0000 +burst_prob_estimates: small_event=0.1329 small_volume=0.9510 heavy_event=0.0043 heavy_volume=0.0011 +suggested_flags: --wg-burst-small-prob 0.1329 --wg-burst-heavy-prob 0.0043 + +=== allusers-long-30.log === +jobs=706918 skipped_rows=1 +hours_observed=1381 +arrivals_per_hour: mean=511.8885 std=1308.5007 +duration_hours: mean=3.2965 std=13.0609 +nodes: mean=1.0137 std=0.6736 +cores: mean=7.5031 std=54.1584 +corr(duration, nodes)=0.030925 corr(duration, cores)=0.020810 +burst_baseline_jobs_per_hour: small=319.0000 heavy=0.0000 +burst_prob_estimates: small_event=0.2274 small_volume=1.0000 heavy_event=0.0210 heavy_volume=0.0072 +suggested_flags: --wg-burst-small-prob 0.2274 --wg-burst-heavy-prob 0.0210 + +=== allusers-main-30.log === +jobs=1448404 skipped_rows=1 +hours_observed=752 +arrivals_per_hour: mean=1926.0691 std=3618.8276 +duration_hours: mean=0.3007 std=0.8709 +nodes: mean=1.0007 std=0.1308 +cores: mean=2.9509 std=3.6442 +corr(duration, nodes)=0.008865 corr(duration, cores)=-0.027545 +burst_baseline_jobs_per_hour: small=907.0000 heavy=0.0000 +burst_prob_estimates: small_event=0.3697 small_volume=1.0000 heavy_event=0.0000 heavy_volume=0.0000 +suggested_flags: --wg-burst-small-prob 0.3697 --wg-burst-heavy-prob 0.0000 + + +NOTE: skipped_rows when duration = 0 rows appear. Not used in calculation. \ No newline at end of file diff --git a/src/workloadgen.py b/src/workloadgen.py index 55188c3..cb50dda 100644 --- a/src/workloadgen.py +++ b/src/workloadgen.py @@ -50,6 +50,29 @@ class WorkloadGenConfig: flat_nodes_jitter: int = 0 flat_cores_jitter: int = 0 + # Optional burst injectors (additive on top of base arrivals) + # Burst 1: many small-ish jobs at once + burst_small_prob: float = 0.0 + burst_small_jobs_min: int = 50 + burst_small_jobs_max: int = 750 + burst_small_duration_min: int = 1 + burst_small_duration_max: int = 8 + burst_small_nodes_min: int = 1 + burst_small_nodes_max: int = 2 + burst_small_cores_min: int = 1 + burst_small_cores_max: int = 16 + + # Burst 2: heavy jobs (high duration + high resource demand) + burst_heavy_prob: float = 0.0 + burst_heavy_jobs_min: int = 1 + burst_heavy_jobs_max: int = 12 + burst_heavy_duration_min: int = 72 + burst_heavy_duration_max: int = 170 + burst_heavy_nodes_min: int = 4 + burst_heavy_nodes_max: int = 16 + burst_heavy_cores_min: int = 32 + burst_heavy_cores_max: int = 96 + # resource ranges (v1: just uniform ranges; later we add mixtures/correlations) min_duration: int = 1 @@ -81,6 +104,31 @@ def __init__(self, cfg: WorkloadGenConfig): raise ValueError("min_cores must be <= max_cores") if int(cfg.uniform_min_new_jobs_per_hour) > int(cfg.max_new_jobs_per_hour): raise ValueError("uniform_min_new_jobs_per_hour must be <= max_new_jobs_per_hour") + if not (0.0 <= float(cfg.burst_small_prob) <= 1.0): + raise ValueError("burst_small_prob must be in [0, 1]") + if not (0.0 <= float(cfg.burst_heavy_prob) <= 1.0): + raise ValueError("burst_heavy_prob must be in [0, 1]") + if int(cfg.burst_small_jobs_min) > int(cfg.burst_small_jobs_max): + raise ValueError("burst_small_jobs_min must be <= burst_small_jobs_max") + if int(cfg.burst_heavy_jobs_min) > int(cfg.burst_heavy_jobs_max): + raise ValueError("burst_heavy_jobs_min must be <= burst_heavy_jobs_max") + + def _bound(value: int, low: int, high: int) -> int: + return int(min(max(int(value), low), high)) + + burst_small_duration_min = _bound(cfg.burst_small_duration_min, int(cfg.min_duration), int(cfg.max_duration)) + burst_small_duration_max = _bound(cfg.burst_small_duration_max, int(cfg.min_duration), int(cfg.max_duration)) + burst_small_nodes_min = _bound(cfg.burst_small_nodes_min, int(cfg.min_nodes), int(cfg.max_nodes)) + burst_small_nodes_max = _bound(cfg.burst_small_nodes_max, int(cfg.min_nodes), int(cfg.max_nodes)) + burst_small_cores_min = _bound(cfg.burst_small_cores_min, int(cfg.min_cores), int(cfg.max_cores)) + burst_small_cores_max = _bound(cfg.burst_small_cores_max, int(cfg.min_cores), int(cfg.max_cores)) + + burst_heavy_duration_min = _bound(cfg.burst_heavy_duration_min, int(cfg.min_duration), int(cfg.max_duration)) + burst_heavy_duration_max = _bound(cfg.burst_heavy_duration_max, int(cfg.min_duration), int(cfg.max_duration)) + burst_heavy_nodes_min = _bound(cfg.burst_heavy_nodes_min, int(cfg.min_nodes), int(cfg.max_nodes)) + burst_heavy_nodes_max = _bound(cfg.burst_heavy_nodes_max, int(cfg.min_nodes), int(cfg.max_nodes)) + burst_heavy_cores_min = _bound(cfg.burst_heavy_cores_min, int(cfg.min_cores), int(cfg.max_cores)) + burst_heavy_cores_max = _bound(cfg.burst_heavy_cores_max, int(cfg.min_cores), int(cfg.max_cores)) self.cfg = replace( cfg, @@ -115,6 +163,18 @@ def __init__(self, cfg: WorkloadGenConfig): if cfg.flat_cores_target is not None else int(cores_mid) ), + burst_small_duration_min=min(burst_small_duration_min, burst_small_duration_max), + burst_small_duration_max=max(burst_small_duration_min, burst_small_duration_max), + burst_small_nodes_min=min(burst_small_nodes_min, burst_small_nodes_max), + burst_small_nodes_max=max(burst_small_nodes_min, burst_small_nodes_max), + burst_small_cores_min=min(burst_small_cores_min, burst_small_cores_max), + burst_small_cores_max=max(burst_small_cores_min, burst_small_cores_max), + burst_heavy_duration_min=min(burst_heavy_duration_min, burst_heavy_duration_max), + burst_heavy_duration_max=max(burst_heavy_duration_min, burst_heavy_duration_max), + burst_heavy_nodes_min=min(burst_heavy_nodes_min, burst_heavy_nodes_max), + burst_heavy_nodes_max=max(burst_heavy_nodes_min, burst_heavy_nodes_max), + burst_heavy_cores_min=min(burst_heavy_cores_min, burst_heavy_cores_max), + burst_heavy_cores_max=max(burst_heavy_cores_min, burst_heavy_cores_max), ) def _sample_attr_array( @@ -191,41 +251,134 @@ def _sample_job_count(self, rng: np.random.Generator) -> int: def sample(self, hour_idx: int, rng: np.random.Generator) -> List[JobSpec]: # hour_idx currently unused, but we keep it to enable daily patterns later. - n = self._sample_job_count(rng) - - if n == 0: - return [] - + base_n = self._sample_job_count(rng) mode = self.cfg.arrivals - durations = self._sample_attr_array( - rng=rng, - size=n, - mode=mode, - min_value=int(self.cfg.min_duration), - max_value=int(self.cfg.max_duration), - poisson_lambda=float(self.cfg.poisson_lambda_duration), - flat_target=int(self.cfg.flat_duration_target), - flat_jitter=int(self.cfg.flat_duration_jitter), - ) - nodes = self._sample_attr_array( - rng=rng, - size=n, - mode=mode, - min_value=int(self.cfg.min_nodes), - max_value=int(self.cfg.max_nodes), - poisson_lambda=float(self.cfg.poisson_lambda_nodes), - flat_target=int(self.cfg.flat_nodes_target), - flat_jitter=int(self.cfg.flat_nodes_jitter), + if base_n > 0: + durations = self._sample_attr_array( + rng=rng, + size=base_n, + mode=mode, + min_value=int(self.cfg.min_duration), + max_value=int(self.cfg.max_duration), + poisson_lambda=float(self.cfg.poisson_lambda_duration), + flat_target=int(self.cfg.flat_duration_target), + flat_jitter=int(self.cfg.flat_duration_jitter), + ) + nodes = self._sample_attr_array( + rng=rng, + size=base_n, + mode=mode, + min_value=int(self.cfg.min_nodes), + max_value=int(self.cfg.max_nodes), + poisson_lambda=float(self.cfg.poisson_lambda_nodes), + flat_target=int(self.cfg.flat_nodes_target), + flat_jitter=int(self.cfg.flat_nodes_jitter), + ) + cores = self._sample_attr_array( + rng=rng, + size=base_n, + mode=mode, + min_value=int(self.cfg.min_cores), + max_value=int(self.cfg.max_cores), + poisson_lambda=float(self.cfg.poisson_lambda_cores), + flat_target=int(self.cfg.flat_cores_target), + flat_jitter=int(self.cfg.flat_cores_jitter), + ) + else: + durations = np.array([], dtype=np.int32) + nodes = np.array([], dtype=np.int32) + cores = np.array([], dtype=np.int32) + + def _sample_burst_count(prob: float, min_jobs: int, max_jobs: int) -> int: + if prob <= 0.0 or max_jobs <= 0: + return 0 + if rng.random() >= prob: + return 0 + return int(rng.integers(int(min_jobs), int(max_jobs) + 1)) + + small_n = _sample_burst_count( + float(self.cfg.burst_small_prob), + int(self.cfg.burst_small_jobs_min), + int(self.cfg.burst_small_jobs_max), ) - cores = self._sample_attr_array( - rng=rng, - size=n, - mode=mode, - min_value=int(self.cfg.min_cores), - max_value=int(self.cfg.max_cores), - poisson_lambda=float(self.cfg.poisson_lambda_cores), - flat_target=int(self.cfg.flat_cores_target), - flat_jitter=int(self.cfg.flat_cores_jitter), + if small_n > 0: + durations = np.concatenate( + [ + durations, + rng.integers( + int(self.cfg.burst_small_duration_min), + int(self.cfg.burst_small_duration_max) + 1, + size=small_n, + ).astype(np.int32), + ] + ) + nodes = np.concatenate( + [ + nodes, + rng.integers( + int(self.cfg.burst_small_nodes_min), + int(self.cfg.burst_small_nodes_max) + 1, + size=small_n, + ).astype(np.int32), + ] + ) + cores = np.concatenate( + [ + cores, + rng.integers( + int(self.cfg.burst_small_cores_min), + int(self.cfg.burst_small_cores_max) + 1, + size=small_n, + ).astype(np.int32), + ] + ) + + heavy_n = _sample_burst_count( + float(self.cfg.burst_heavy_prob), + int(self.cfg.burst_heavy_jobs_min), + int(self.cfg.burst_heavy_jobs_max), ) + if heavy_n > 0: + durations = np.concatenate( + [ + durations, + rng.integers( + int(self.cfg.burst_heavy_duration_min), + int(self.cfg.burst_heavy_duration_max) + 1, + size=heavy_n, + ).astype(np.int32), + ] + ) + nodes = np.concatenate( + [ + nodes, + rng.integers( + int(self.cfg.burst_heavy_nodes_min), + int(self.cfg.burst_heavy_nodes_max) + 1, + size=heavy_n, + ).astype(np.int32), + ] + ) + cores = np.concatenate( + [ + cores, + rng.integers( + int(self.cfg.burst_heavy_cores_min), + int(self.cfg.burst_heavy_cores_max) + 1, + size=heavy_n, + ).astype(np.int32), + ] + ) + + total_n = len(durations) + if self.cfg.hard_cap_jobs is not None and total_n > int(self.cfg.hard_cap_jobs): + hard_cap = int(self.cfg.hard_cap_jobs) + durations = durations[:hard_cap] + nodes = nodes[:hard_cap] + cores = cores[:hard_cap] + total_n = hard_cap + + if total_n == 0: + return [] - return [JobSpec(int(durations[i]), int(nodes[i]), int(cores[i])) for i in range(n)] + return [JobSpec(int(durations[i]), int(nodes[i]), int(cores[i])) for i in range(total_n)] diff --git a/test/run_all.py b/test/run_all.py index c2d85ee..657bd66 100644 --- a/test/run_all.py +++ b/test/run_all.py @@ -21,6 +21,7 @@ ["python", "-m", "test.test_sampler_hourly_aggregated", "--file-path", "data/allusers-gpu-30.log"], ["python", "-m", "test.test_sampler_jobs", "--file-path", "data/allusers-gpu-30.log"], ["python", "-m", "test.test_sampler_jobs_aggregated", "--file-path", "data/allusers-gpu-30.log"], + ["python", "-m", "test.test_inspect_workloadgen", "--arrivals", "poisson", "--poisson-lambdas4", "200,10,6,24", "--max-jobs-hour", "1500", "--hours", "336", "--plot", "--burst-small-prob", "0.2", "--burst-heavy-prob", "0.02"], ] def main(): diff --git a/test/test_inspect_workloadgen.py b/test/test_inspect_workloadgen.py index d4b3db0..027ee5a 100644 --- a/test/test_inspect_workloadgen.py +++ b/test/test_inspect_workloadgen.py @@ -1,6 +1,6 @@ """ Run with: -python -m test.test_inspect_workloadgen --arrivals poisson --poisson-lambdas4 200,80,6,24 --max-jobs-hour 1500 --hours 336 --plot +python -m test.test_inspect_workloadgen --arrivals poisson --poisson-lambdas4 200,10,6,24 --max-jobs-hour 1500 --hours 336 --plot --burst-small-prob 0.2 --burst-heavy-prob 0.02 """ # inspect_workloadgen.py @@ -13,52 +13,7 @@ import os from src.workloadgen import WorkloadGenConfig, WorkloadGenerator - - -def parse_quad_floats(raw: str): - parts = [p.strip() for p in str(raw).split(",")] - if len(parts) != 4: - raise argparse.ArgumentTypeError( - "Expected 4 comma-separated floats: arrivals,duration,nodes,cores" - ) - try: - return tuple(float(p) for p in parts) - except ValueError as exc: - raise argparse.ArgumentTypeError(f"Invalid float in '{raw}'") from exc - - -def parse_quad_ints(raw: str): - parts = [p.strip() for p in str(raw).split(",")] - if len(parts) != 4: - raise argparse.ArgumentTypeError( - "Expected 4 comma-separated ints: arrivals,duration,nodes,cores" - ) - try: - return tuple(int(p) for p in parts) - except ValueError as exc: - raise argparse.ArgumentTypeError(f"Invalid int in '{raw}'") from exc - - -def parse_quad_ranges(raw: str): - parts = [p.strip() for p in str(raw).split(",")] - if len(parts) != 4: - raise argparse.ArgumentTypeError( - "Expected 4 comma-separated ranges: a_min:a_max,d_min:d_max,n_min:n_max,c_min:c_max" - ) - ranges = [] - for part in parts: - bounds = [b.strip() for b in part.split(":")] - if len(bounds) != 2: - raise argparse.ArgumentTypeError(f"Invalid range '{part}', expected min:max") - try: - low = int(bounds[0]) - high = int(bounds[1]) - except ValueError as exc: - raise argparse.ArgumentTypeError(f"Invalid int in range '{part}'") from exc - if low > high: - raise argparse.ArgumentTypeError(f"Range min > max in '{part}'") - ranges.append((low, high)) - return tuple(ranges) +from train import _parse_quad_floats, _parse_quad_ints, _parse_quad_ranges def digest_jobs_triplets(triplets): @@ -90,20 +45,22 @@ def main(): ap.add_argument("--hours", type=int, default=24 * 14) ap.add_argument("--arrivals", choices=["flat", "poisson", "uniform"], default="poisson") ap.add_argument("--poisson-lambda", type=float, default=200.0, help="Legacy: arrivals-only poisson lambda.") - ap.add_argument("--poisson-lambdas4", type=parse_quad_floats, default=None, help="arrivals,duration,nodes,cores") + ap.add_argument("--poisson-lambdas4", type=_parse_quad_floats, default=None, help="arrivals,duration,nodes,cores") ap.add_argument("--max-jobs-hour", type=int, default=1500) ap.add_argument("--plot", action="store_true") # Flat params (true flat with optional jitter) ap.add_argument("--flat-jobs-hour", type=int, default=200, help="Legacy: arrivals-only flat target.") ap.add_argument("--flat-jitter", type=int, default=0, help="Legacy: arrivals-only flat jitter.") - ap.add_argument("--flat-targets4", type=parse_quad_ints, default=None, help="arrivals,duration,nodes,cores") - ap.add_argument("--flat-jitters4", type=parse_quad_ints, default=None, help="arrivals,duration,nodes,cores") + ap.add_argument("--flat-targets4", type=_parse_quad_ints, default=None, help="arrivals,duration,nodes,cores") + ap.add_argument("--flat-jitters4", type=_parse_quad_ints, default=None, help="arrivals,duration,nodes,cores") ap.add_argument( "--uniform-ranges4", - type=parse_quad_ranges, + type=_parse_quad_ranges, default=None, help="a_min:a_max,d_min:d_max,n_min:n_max,c_min:c_max", ) + ap.add_argument("--burst-small-prob", type=float, default=0.0, help="Probability of additive small-job burst per hour.") + ap.add_argument("--burst-heavy-prob", type=float, default=0.0, help="Probability of additive heavy-job burst per hour.") args = ap.parse_args() # Default ranges (used for uniform and for clipping in all modes). @@ -159,6 +116,8 @@ def main(): flat_duration_jitter=flat_duration_jitter, flat_nodes_jitter=flat_nodes_jitter, flat_cores_jitter=flat_cores_jitter, + burst_small_prob=float(args.burst_small_prob), + burst_heavy_prob=float(args.burst_heavy_prob), min_duration=min_duration, max_duration=max_duration, min_nodes=min_nodes, @@ -274,7 +233,7 @@ def main(): axs[3, 1].set_xlabel("hour of day") axs[3, 1].set_ylabel("jobs") - plt.show() + #plt.show() timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") prefix = f"{args.arrivals}_lambda{poisson_lambda_arrivals}" if args.arrivals == "poisson" else args.arrivals fname = f"{prefix}_{timestamp}.png" if prefix else f"Workload-Gen_{timestamp}.png" diff --git a/test/test_sanity_env.py b/test/test_sanity_env.py index 54f7f60..4927da3 100644 --- a/test/test_sanity_env.py +++ b/test/test_sanity_env.py @@ -19,6 +19,7 @@ from src.plot_config import PlotConfig import pandas as pd from src.workloadgen import WorkloadGenerator, WorkloadGenConfig +from train import _parse_quad_floats, _parse_quad_ints, _parse_quad_ranges # Import environment variables: from src.config import ( @@ -37,52 +38,6 @@ def load_prices(prices_file_path: str | None): print(f"Loaded {len(prices)} prices from CSV: {prices_file_path}") return prices - -def _parse_quad_floats(raw: str): - parts = [p.strip() for p in str(raw).split(",")] - if len(parts) != 4: - raise argparse.ArgumentTypeError( - "Expected 4 comma-separated floats: arrivals,duration,nodes,cores" - ) - try: - return tuple(float(p) for p in parts) - except ValueError as exc: - raise argparse.ArgumentTypeError(f"Invalid float in '{raw}'") from exc - - -def _parse_quad_ints(raw: str): - parts = [p.strip() for p in str(raw).split(",")] - if len(parts) != 4: - raise argparse.ArgumentTypeError( - "Expected 4 comma-separated ints: arrivals,duration,nodes,cores" - ) - try: - return tuple(int(p) for p in parts) - except ValueError as exc: - raise argparse.ArgumentTypeError(f"Invalid int in '{raw}'") from exc - - -def _parse_quad_ranges(raw: str): - parts = [p.strip() for p in str(raw).split(",")] - if len(parts) != 4: - raise argparse.ArgumentTypeError( - "Expected 4 comma-separated ranges: a_min:a_max,d_min:d_max,n_min:n_max,c_min:c_max" - ) - ranges = [] - for part in parts: - bounds = [b.strip() for b in part.split(":")] - if len(bounds) != 2: - raise argparse.ArgumentTypeError(f"Invalid range '{part}', expected min:max") - try: - low = int(bounds[0]) - high = int(bounds[1]) - except ValueError as exc: - raise argparse.ArgumentTypeError(f"Invalid int in range '{part}'") from exc - if low > high: - raise argparse.ArgumentTypeError(f"Range min > max in '{part}'") - ranges.append((low, high)) - return tuple(ranges) - # ----------------------------- # Invariants / sanity checks # ----------------------------- @@ -279,6 +234,8 @@ def parse_args(): default=None, help="a_min:a_max,d_min:d_max,n_min:n_max,c_min:c_max", ) + p.add_argument("--wg-burst-small-prob", type=float, default=0.0, help="Probability of additive small-job burst per hour.") + p.add_argument("--wg-burst-heavy-prob", type=float, default=0.0, help="Probability of additive heavy-job burst per hour.") p.add_argument("--print-job-every", type=int, default=0, help="Print one sample job every N steps (0 disables).") p.add_argument("--print-job-kind", choices=["queue", "running", "both"], default="queue", 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.") @@ -358,6 +315,8 @@ def make_env_from_args(args, env_cls=ComputeClusterEnv): flat_duration_jitter=flat_duration_jitter, flat_nodes_jitter=flat_nodes_jitter, flat_cores_jitter=flat_cores_jitter, + burst_small_prob=float(args.wg_burst_small_prob), + burst_heavy_prob=float(args.wg_burst_heavy_prob), min_duration=min_duration, max_duration=max_duration, min_nodes=min_nodes, diff --git a/test/test_sanity_workloadgen.py b/test/test_sanity_workloadgen.py index 1f0b387..d5212bf 100644 --- a/test/test_sanity_workloadgen.py +++ b/test/test_sanity_workloadgen.py @@ -93,10 +93,120 @@ def test_poisson_attribute_lambdas_are_used(): assert mean_nodes < 6.0, mean_nodes assert mean_cores < 6.0, mean_cores + +def test_bursts_are_additive_to_base_distribution(): + cfg = WorkloadGenConfig( + arrivals="flat", + flat_jobs_per_hour=10, + flat_jitter=0, + flat_duration_target=50, + flat_nodes_target=5, + flat_cores_target=20, + flat_duration_jitter=0, + flat_nodes_jitter=0, + flat_cores_jitter=0, + burst_small_prob=1.0, + burst_small_jobs_min=3, + burst_small_jobs_max=3, + burst_small_duration_min=1, + burst_small_duration_max=1, + burst_small_nodes_min=1, + burst_small_nodes_max=1, + burst_small_cores_min=1, + burst_small_cores_max=1, + burst_heavy_prob=1.0, + burst_heavy_jobs_min=2, + burst_heavy_jobs_max=2, + burst_heavy_duration_min=170, + burst_heavy_duration_max=170, + burst_heavy_nodes_min=16, + burst_heavy_nodes_max=16, + burst_heavy_cores_min=96, + burst_heavy_cores_max=96, + min_duration=1, + max_duration=170, + min_nodes=1, + max_nodes=16, + min_cores=1, + max_cores=96, + ) + gen = WorkloadGenerator(cfg) + rng = np.random.default_rng(10) + jobs = gen.sample(0, rng) + tuples = [(j.duration, j.nodes, j.cores_per_node) for j in jobs] + + assert len(jobs) == 15 + assert tuples.count((50, 5, 20)) == 10 + assert tuples.count((1, 1, 1)) == 3 + assert tuples.count((170, 16, 96)) == 2 + + +def test_zero_burst_prob_keeps_base_only(): + cfg = WorkloadGenConfig( + arrivals="flat", + flat_jobs_per_hour=12, + flat_jitter=0, + flat_duration_target=40, + flat_nodes_target=4, + flat_cores_target=12, + flat_duration_jitter=0, + flat_nodes_jitter=0, + flat_cores_jitter=0, + burst_small_prob=0.0, + burst_heavy_prob=0.0, + burst_small_jobs_min=10, + burst_small_jobs_max=10, + burst_heavy_jobs_min=10, + burst_heavy_jobs_max=10, + min_duration=1, + max_duration=170, + min_nodes=1, + max_nodes=16, + min_cores=1, + max_cores=96, + ) + gen = WorkloadGenerator(cfg) + rng = np.random.default_rng(11) + jobs = gen.sample(0, rng) + + assert len(jobs) == 12 + assert all((j.duration, j.nodes, j.cores_per_node) == (40, 4, 12) for j in jobs) + + +def test_burst_determinism_with_fixed_seed(): + cfg = WorkloadGenConfig( + arrivals="poisson", + poisson_lambda=30.0, + burst_small_prob=0.35, + burst_small_jobs_min=10, + burst_small_jobs_max=20, + burst_heavy_prob=0.15, + burst_heavy_jobs_min=1, + burst_heavy_jobs_max=4, + min_duration=1, + max_duration=170, + min_nodes=1, + max_nodes=16, + min_cores=1, + max_cores=96, + ) + g1 = WorkloadGenerator(cfg) + g2 = WorkloadGenerator(cfg) + r1 = np.random.default_rng(1234) + r2 = np.random.default_rng(1234) + + for h in range(100): + a = [(j.duration, j.nodes, j.cores_per_node) for j in g1.sample(h, r1)] + b = [(j.duration, j.nodes, j.cores_per_node) for j in g2.sample(h, r2)] + assert a == b + if __name__ == "__main__": test_determinism() test_constraints() test_poisson_mean_sanity() test_flat_attribute_targets() test_poisson_attribute_lambdas_are_used() + test_bursts_are_additive_to_base_distribution() + test_zero_burst_prob_keeps_base_only() + test_burst_determinism_with_fixed_seed() print("[OK] workloadgen sanity checks passed") diff --git a/train.py b/train.py index 3a5db3a..021f7db 100644 --- a/train.py +++ b/train.py @@ -117,6 +117,8 @@ def main(): default=None, help="a_min:a_max,d_min:d_max,n_min:n_max,c_min:c_max", ) + parser.add_argument("--wg-burst-small-prob", type=float, default=0.0, help="Probability of additive small-job burst per hour.") + parser.add_argument("--wg-burst-heavy-prob", type=float, default=0.0, help="Probability of additive heavy-job burst per hour.") 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).") @@ -231,6 +233,8 @@ def main(): flat_duration_jitter=flat_duration_jitter, flat_nodes_jitter=flat_nodes_jitter, flat_cores_jitter=flat_cores_jitter, + burst_small_prob=float(args.wg_burst_small_prob), + burst_heavy_prob=float(args.wg_burst_heavy_prob), min_duration=min_duration, max_duration=max_duration, min_nodes=min_nodes, From 4462ca0a51e5face95d5a4c3cec34b5c6faf4b8b Mon Sep 17 00:00:00 2001 From: Alexey Rybalchenko Date: Mon, 16 Feb 2026 18:41:19 +0100 Subject: [PATCH 4/9] format --- test/test_sanity_env.py | 9 ++------- 1 file changed, 2 insertions(+), 7 deletions(-) diff --git a/test/test_sanity_env.py b/test/test_sanity_env.py index 4927da3..57fb9bc 100644 --- a/test/test_sanity_env.py +++ b/test/test_sanity_env.py @@ -228,12 +228,7 @@ def parse_args(): p.add_argument("--wg-max-jobs-hour", type=int, default=1500, help="Cap jobs/hour for generator.") p.add_argument("--wg-flat-targets4", type=_parse_quad_ints, default=None, help="arrivals,duration,nodes,cores") p.add_argument("--wg-flat-jitters4", type=_parse_quad_ints, default=None, help="arrivals,duration,nodes,cores") - p.add_argument( - "--wg-uniform-ranges4", - type=_parse_quad_ranges, - default=None, - help="a_min:a_max,d_min:d_max,n_min:n_max,c_min:c_max", - ) + p.add_argument("--wg-uniform-ranges4", type=_parse_quad_ranges, default=None, help="a_min:a_max,d_min:d_max,n_min:n_max,c_min:c_max") p.add_argument("--wg-burst-small-prob", type=float, default=0.0, help="Probability of additive small-job burst per hour.") p.add_argument("--wg-burst-heavy-prob", type=float, default=0.0, help="Probability of additive heavy-job burst per hour.") p.add_argument("--print-job-every", type=int, default=0, help="Print one sample job every N steps (0 disables).") @@ -479,7 +474,7 @@ def cmp(name, a, b): if args.check_determinism: 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)) From a2510993467393a15a02cfb079e3c7e4b8f2da44 Mon Sep 17 00:00:00 2001 From: Alexey Rybalchenko Date: Mon, 16 Feb 2026 18:41:49 +0100 Subject: [PATCH 5/9] iworkloadgen: remove no-op casts --- src/workloadgen.py | 166 ++++++++++++++++++++++----------------------- 1 file changed, 83 insertions(+), 83 deletions(-) diff --git a/src/workloadgen.py b/src/workloadgen.py index cb50dda..cda1889 100644 --- a/src/workloadgen.py +++ b/src/workloadgen.py @@ -92,76 +92,76 @@ def __init__(self, cfg: WorkloadGenConfig): if arrivals not in ("flat", "poisson", "uniform"): raise ValueError(f"arrivals must be 'flat', 'uniform' or 'poisson', got: {cfg.arrivals}") - duration_mid = int(round((int(cfg.min_duration) + int(cfg.max_duration)) / 2.0)) - nodes_mid = int(round((int(cfg.min_nodes) + int(cfg.max_nodes)) / 2.0)) - cores_mid = int(round((int(cfg.min_cores) + int(cfg.max_cores)) / 2.0)) + duration_mid = (cfg.min_duration + cfg.max_duration) // 2 + nodes_mid = (cfg.min_nodes + cfg.max_nodes) // 2 + cores_mid = (cfg.min_cores + cfg.max_cores) // 2 - if int(cfg.min_duration) > int(cfg.max_duration): + if cfg.min_duration > cfg.max_duration: raise ValueError("min_duration must be <= max_duration") - if int(cfg.min_nodes) > int(cfg.max_nodes): + if cfg.min_nodes > cfg.max_nodes: raise ValueError("min_nodes must be <= max_nodes") - if int(cfg.min_cores) > int(cfg.max_cores): + if cfg.min_cores > cfg.max_cores: raise ValueError("min_cores must be <= max_cores") - if int(cfg.uniform_min_new_jobs_per_hour) > int(cfg.max_new_jobs_per_hour): + if cfg.uniform_min_new_jobs_per_hour > cfg.max_new_jobs_per_hour: raise ValueError("uniform_min_new_jobs_per_hour must be <= max_new_jobs_per_hour") - if not (0.0 <= float(cfg.burst_small_prob) <= 1.0): + if not (0.0 <= cfg.burst_small_prob <= 1.0): raise ValueError("burst_small_prob must be in [0, 1]") - if not (0.0 <= float(cfg.burst_heavy_prob) <= 1.0): + if not (0.0 <= cfg.burst_heavy_prob <= 1.0): raise ValueError("burst_heavy_prob must be in [0, 1]") - if int(cfg.burst_small_jobs_min) > int(cfg.burst_small_jobs_max): + if cfg.burst_small_jobs_min > cfg.burst_small_jobs_max: raise ValueError("burst_small_jobs_min must be <= burst_small_jobs_max") - if int(cfg.burst_heavy_jobs_min) > int(cfg.burst_heavy_jobs_max): + if cfg.burst_heavy_jobs_min > cfg.burst_heavy_jobs_max: raise ValueError("burst_heavy_jobs_min must be <= burst_heavy_jobs_max") def _bound(value: int, low: int, high: int) -> int: - return int(min(max(int(value), low), high)) - - burst_small_duration_min = _bound(cfg.burst_small_duration_min, int(cfg.min_duration), int(cfg.max_duration)) - burst_small_duration_max = _bound(cfg.burst_small_duration_max, int(cfg.min_duration), int(cfg.max_duration)) - burst_small_nodes_min = _bound(cfg.burst_small_nodes_min, int(cfg.min_nodes), int(cfg.max_nodes)) - burst_small_nodes_max = _bound(cfg.burst_small_nodes_max, int(cfg.min_nodes), int(cfg.max_nodes)) - burst_small_cores_min = _bound(cfg.burst_small_cores_min, int(cfg.min_cores), int(cfg.max_cores)) - burst_small_cores_max = _bound(cfg.burst_small_cores_max, int(cfg.min_cores), int(cfg.max_cores)) - - burst_heavy_duration_min = _bound(cfg.burst_heavy_duration_min, int(cfg.min_duration), int(cfg.max_duration)) - burst_heavy_duration_max = _bound(cfg.burst_heavy_duration_max, int(cfg.min_duration), int(cfg.max_duration)) - burst_heavy_nodes_min = _bound(cfg.burst_heavy_nodes_min, int(cfg.min_nodes), int(cfg.max_nodes)) - burst_heavy_nodes_max = _bound(cfg.burst_heavy_nodes_max, int(cfg.min_nodes), int(cfg.max_nodes)) - burst_heavy_cores_min = _bound(cfg.burst_heavy_cores_min, int(cfg.min_cores), int(cfg.max_cores)) - burst_heavy_cores_max = _bound(cfg.burst_heavy_cores_max, int(cfg.min_cores), int(cfg.max_cores)) + return min(max(value, low), high) + + burst_small_duration_min = _bound(cfg.burst_small_duration_min, cfg.min_duration, cfg.max_duration) + burst_small_duration_max = _bound(cfg.burst_small_duration_max, cfg.min_duration, cfg.max_duration) + burst_small_nodes_min = _bound(cfg.burst_small_nodes_min, cfg.min_nodes, cfg.max_nodes) + burst_small_nodes_max = _bound(cfg.burst_small_nodes_max, cfg.min_nodes, cfg.max_nodes) + burst_small_cores_min = _bound(cfg.burst_small_cores_min, cfg.min_cores, cfg.max_cores) + burst_small_cores_max = _bound(cfg.burst_small_cores_max, cfg.min_cores, cfg.max_cores) + + burst_heavy_duration_min = _bound(cfg.burst_heavy_duration_min, cfg.min_duration, cfg.max_duration) + burst_heavy_duration_max = _bound(cfg.burst_heavy_duration_max, cfg.min_duration, cfg.max_duration) + burst_heavy_nodes_min = _bound(cfg.burst_heavy_nodes_min, cfg.min_nodes, cfg.max_nodes) + burst_heavy_nodes_max = _bound(cfg.burst_heavy_nodes_max, cfg.min_nodes, cfg.max_nodes) + burst_heavy_cores_min = _bound(cfg.burst_heavy_cores_min, cfg.min_cores, cfg.max_cores) + burst_heavy_cores_max = _bound(cfg.burst_heavy_cores_max, cfg.min_cores, cfg.max_cores) self.cfg = replace( cfg, arrivals=arrivals, poisson_lambda_duration=( - float(cfg.poisson_lambda_duration) + cfg.poisson_lambda_duration if cfg.poisson_lambda_duration is not None else float(duration_mid) ), poisson_lambda_nodes=( - float(cfg.poisson_lambda_nodes) + cfg.poisson_lambda_nodes if cfg.poisson_lambda_nodes is not None else float(nodes_mid) ), poisson_lambda_cores=( - float(cfg.poisson_lambda_cores) + cfg.poisson_lambda_cores if cfg.poisson_lambda_cores is not None else float(cores_mid) ), flat_duration_target=( - int(cfg.flat_duration_target) + cfg.flat_duration_target if cfg.flat_duration_target is not None - else int(duration_mid) + else duration_mid ), flat_nodes_target=( - int(cfg.flat_nodes_target) + cfg.flat_nodes_target if cfg.flat_nodes_target is not None - else int(nodes_mid) + else nodes_mid ), flat_cores_target=( - int(cfg.flat_cores_target) + cfg.flat_cores_target if cfg.flat_cores_target is not None - else int(cores_mid) + else cores_mid ), burst_small_duration_min=min(burst_small_duration_min, burst_small_duration_max), burst_small_duration_max=max(burst_small_duration_min, burst_small_duration_max), @@ -193,21 +193,21 @@ def _sample_attr_array( if mode == "flat": if flat_jitter <= 0: - values = np.full(size, int(flat_target), dtype=np.int64) + values = np.full(size, flat_target, dtype=np.int64) else: values = rng.integers( - int(flat_target) - int(flat_jitter), - int(flat_target) + int(flat_jitter) + 1, + flat_target - flat_jitter, + flat_target + flat_jitter + 1, size=size, ) elif mode == "poisson": - values = rng.poisson(float(poisson_lambda), size=size) + values = rng.poisson(poisson_lambda, size=size) elif mode == "uniform": - values = rng.integers(int(min_value), int(max_value) + 1, size=size) + values = rng.integers(min_value, max_value + 1, size=size) else: raise ValueError(f"Unknown sampling mode: {mode}") - return np.clip(values, int(min_value), int(max_value)).astype(np.int32) + return np.clip(values, min_value, max_value).astype(np.int32) def _sample_job_count(self, rng: np.random.Generator) -> int: """ @@ -219,8 +219,8 @@ def _sample_job_count(self, rng: np.random.Generator) -> int: mode = self.cfg.arrivals if mode == "flat": - target = int(self.cfg.flat_jobs_per_hour) - jitter = int(self.cfg.flat_jitter) + target = self.cfg.flat_jobs_per_hour + jitter = self.cfg.flat_jitter if jitter <= 0: k = target @@ -233,8 +233,8 @@ def _sample_job_count(self, rng: np.random.Generator) -> int: elif mode == "uniform": k = int( rng.integers( - int(self.cfg.uniform_min_new_jobs_per_hour), - int(self.cfg.max_new_jobs_per_hour) + 1, + self.cfg.uniform_min_new_jobs_per_hour, + self.cfg.max_new_jobs_per_hour + 1, ) ) @@ -242,9 +242,9 @@ def _sample_job_count(self, rng: np.random.Generator) -> int: raise ValueError(f"Unknown arrivals mode: {mode}") # clamp + safety - k = min(k, int(self.cfg.max_new_jobs_per_hour)) + k = min(k, self.cfg.max_new_jobs_per_hour) if self.cfg.hard_cap_jobs is not None: - k = min(k, int(self.cfg.hard_cap_jobs)) + k = min(k, self.cfg.hard_cap_jobs) if k < 0: k = 0 return k @@ -258,31 +258,31 @@ def sample(self, hour_idx: int, rng: np.random.Generator) -> List[JobSpec]: rng=rng, size=base_n, mode=mode, - min_value=int(self.cfg.min_duration), - max_value=int(self.cfg.max_duration), - poisson_lambda=float(self.cfg.poisson_lambda_duration), - flat_target=int(self.cfg.flat_duration_target), - flat_jitter=int(self.cfg.flat_duration_jitter), + min_value=self.cfg.min_duration, + max_value=self.cfg.max_duration, + poisson_lambda=self.cfg.poisson_lambda_duration, + flat_target=self.cfg.flat_duration_target, + flat_jitter=self.cfg.flat_duration_jitter, ) nodes = self._sample_attr_array( rng=rng, size=base_n, mode=mode, - min_value=int(self.cfg.min_nodes), - max_value=int(self.cfg.max_nodes), - poisson_lambda=float(self.cfg.poisson_lambda_nodes), - flat_target=int(self.cfg.flat_nodes_target), - flat_jitter=int(self.cfg.flat_nodes_jitter), + min_value=self.cfg.min_nodes, + max_value=self.cfg.max_nodes, + poisson_lambda=self.cfg.poisson_lambda_nodes, + flat_target=self.cfg.flat_nodes_target, + flat_jitter=self.cfg.flat_nodes_jitter, ) cores = self._sample_attr_array( rng=rng, size=base_n, mode=mode, - min_value=int(self.cfg.min_cores), - max_value=int(self.cfg.max_cores), - poisson_lambda=float(self.cfg.poisson_lambda_cores), - flat_target=int(self.cfg.flat_cores_target), - flat_jitter=int(self.cfg.flat_cores_jitter), + min_value=self.cfg.min_cores, + max_value=self.cfg.max_cores, + poisson_lambda=self.cfg.poisson_lambda_cores, + flat_target=self.cfg.flat_cores_target, + flat_jitter=self.cfg.flat_cores_jitter, ) else: durations = np.array([], dtype=np.int32) @@ -294,20 +294,20 @@ def _sample_burst_count(prob: float, min_jobs: int, max_jobs: int) -> int: return 0 if rng.random() >= prob: return 0 - return int(rng.integers(int(min_jobs), int(max_jobs) + 1)) + return int(rng.integers(min_jobs, max_jobs + 1)) small_n = _sample_burst_count( - float(self.cfg.burst_small_prob), - int(self.cfg.burst_small_jobs_min), - int(self.cfg.burst_small_jobs_max), + self.cfg.burst_small_prob, + self.cfg.burst_small_jobs_min, + self.cfg.burst_small_jobs_max, ) if small_n > 0: durations = np.concatenate( [ durations, rng.integers( - int(self.cfg.burst_small_duration_min), - int(self.cfg.burst_small_duration_max) + 1, + self.cfg.burst_small_duration_min, + self.cfg.burst_small_duration_max + 1, size=small_n, ).astype(np.int32), ] @@ -316,8 +316,8 @@ def _sample_burst_count(prob: float, min_jobs: int, max_jobs: int) -> int: [ nodes, rng.integers( - int(self.cfg.burst_small_nodes_min), - int(self.cfg.burst_small_nodes_max) + 1, + self.cfg.burst_small_nodes_min, + self.cfg.burst_small_nodes_max + 1, size=small_n, ).astype(np.int32), ] @@ -326,25 +326,25 @@ def _sample_burst_count(prob: float, min_jobs: int, max_jobs: int) -> int: [ cores, rng.integers( - int(self.cfg.burst_small_cores_min), - int(self.cfg.burst_small_cores_max) + 1, + self.cfg.burst_small_cores_min, + self.cfg.burst_small_cores_max + 1, size=small_n, ).astype(np.int32), ] ) heavy_n = _sample_burst_count( - float(self.cfg.burst_heavy_prob), - int(self.cfg.burst_heavy_jobs_min), - int(self.cfg.burst_heavy_jobs_max), + self.cfg.burst_heavy_prob, + self.cfg.burst_heavy_jobs_min, + self.cfg.burst_heavy_jobs_max, ) if heavy_n > 0: durations = np.concatenate( [ durations, rng.integers( - int(self.cfg.burst_heavy_duration_min), - int(self.cfg.burst_heavy_duration_max) + 1, + self.cfg.burst_heavy_duration_min, + self.cfg.burst_heavy_duration_max + 1, size=heavy_n, ).astype(np.int32), ] @@ -353,8 +353,8 @@ def _sample_burst_count(prob: float, min_jobs: int, max_jobs: int) -> int: [ nodes, rng.integers( - int(self.cfg.burst_heavy_nodes_min), - int(self.cfg.burst_heavy_nodes_max) + 1, + self.cfg.burst_heavy_nodes_min, + self.cfg.burst_heavy_nodes_max + 1, size=heavy_n, ).astype(np.int32), ] @@ -363,16 +363,16 @@ def _sample_burst_count(prob: float, min_jobs: int, max_jobs: int) -> int: [ cores, rng.integers( - int(self.cfg.burst_heavy_cores_min), - int(self.cfg.burst_heavy_cores_max) + 1, + self.cfg.burst_heavy_cores_min, + self.cfg.burst_heavy_cores_max + 1, size=heavy_n, ).astype(np.int32), ] ) total_n = len(durations) - if self.cfg.hard_cap_jobs is not None and total_n > int(self.cfg.hard_cap_jobs): - hard_cap = int(self.cfg.hard_cap_jobs) + if self.cfg.hard_cap_jobs is not None and total_n > self.cfg.hard_cap_jobs: + hard_cap = self.cfg.hard_cap_jobs durations = durations[:hard_cap] nodes = nodes[:hard_cap] cores = cores[:hard_cap] From 83b0bad893f38d01399ebb574b4171109552ae3a Mon Sep 17 00:00:00 2001 From: Alexey Rybalchenko Date: Mon, 16 Feb 2026 19:00:48 +0100 Subject: [PATCH 6/9] Remove more no-op casts --- test/test_inspect_workloadgen.py | 6 +++--- test/test_sanity_env.py | 14 +++++++------- train.py | 14 +++++++------- 3 files changed, 17 insertions(+), 17 deletions(-) diff --git a/test/test_inspect_workloadgen.py b/test/test_inspect_workloadgen.py index 027ee5a..02aeb5d 100644 --- a/test/test_inspect_workloadgen.py +++ b/test/test_inspect_workloadgen.py @@ -72,9 +72,9 @@ def main(): if args.uniform_ranges4 is not None: (uniform_min_jobs, max_jobs_hour), (min_duration, max_duration), (min_nodes, max_nodes), (min_cores, max_cores) = args.uniform_ranges4 - default_duration_mid = int(round((min_duration + max_duration) / 2.0)) - default_nodes_mid = int(round((min_nodes + max_nodes) / 2.0)) - default_cores_mid = int(round((min_cores + max_cores) / 2.0)) + default_duration_mid = (min_duration + max_duration) // 2 + default_nodes_mid = (min_nodes + max_nodes) // 2 + default_cores_mid = (min_cores + max_cores) // 2 if args.poisson_lambdas4 is not None: poisson_lambda_arrivals, poisson_lambda_duration, poisson_lambda_nodes, poisson_lambda_cores = args.poisson_lambdas4 diff --git a/test/test_sanity_env.py b/test/test_sanity_env.py index 57fb9bc..2c0f6ec 100644 --- a/test/test_sanity_env.py +++ b/test/test_sanity_env.py @@ -253,7 +253,7 @@ def make_env_from_args(args, env_cls=ComputeClusterEnv): workload_gen = None if args.workload_gen: uniform_min_jobs = 0 - max_jobs_hour = int(args.wg_max_jobs_hour) + max_jobs_hour = args.wg_max_jobs_hour min_duration, max_duration = 1, MAX_JOB_DURATION min_nodes, max_nodes = MIN_NODES_PER_JOB, MAX_NODES_PER_JOB min_cores, max_cores = MIN_CORES_PER_JOB, CORES_PER_NODE @@ -266,14 +266,14 @@ def make_env_from_args(args, env_cls=ComputeClusterEnv): (min_cores, max_cores), ) = args.wg_uniform_ranges4 - duration_mid = int(round((min_duration + max_duration) / 2.0)) - nodes_mid = int(round((min_nodes + max_nodes) / 2.0)) - cores_mid = int(round((min_cores + max_cores) / 2.0)) + duration_mid = (min_duration + max_duration) // 2 + nodes_mid = (min_nodes + max_nodes) // 2 + cores_mid = (min_cores + max_cores) // 2 if args.wg_poisson_lambdas4 is not None: poisson_lambda_arrivals, poisson_lambda_duration, poisson_lambda_nodes, poisson_lambda_cores = args.wg_poisson_lambdas4 else: - poisson_lambda_arrivals = float(args.wg_poisson_lambda) + poisson_lambda_arrivals = args.wg_poisson_lambda poisson_lambda_duration = float(duration_mid) poisson_lambda_nodes = float(nodes_mid) poisson_lambda_cores = float(cores_mid) @@ -310,8 +310,8 @@ def make_env_from_args(args, env_cls=ComputeClusterEnv): flat_duration_jitter=flat_duration_jitter, flat_nodes_jitter=flat_nodes_jitter, flat_cores_jitter=flat_cores_jitter, - burst_small_prob=float(args.wg_burst_small_prob), - burst_heavy_prob=float(args.wg_burst_heavy_prob), + burst_small_prob=args.wg_burst_small_prob, + burst_heavy_prob=args.wg_burst_heavy_prob, min_duration=min_duration, max_duration=max_duration, min_nodes=min_nodes, diff --git a/train.py b/train.py index 021f7db..dda4712 100644 --- a/train.py +++ b/train.py @@ -176,7 +176,7 @@ def main(): workload_gen = None if args.workload_gen: uniform_min_jobs = 0 - max_jobs_hour = int(args.wg_max_jobs_hour) + max_jobs_hour = args.wg_max_jobs_hour min_duration, max_duration = 1, MAX_JOB_DURATION min_nodes, max_nodes = MIN_NODES_PER_JOB, MAX_NODES_PER_JOB min_cores, max_cores = MIN_CORES_PER_JOB, CORES_PER_NODE @@ -189,14 +189,14 @@ def main(): (min_cores, max_cores), ) = args.wg_uniform_ranges4 - duration_mid = int(round((min_duration + max_duration) / 2.0)) - nodes_mid = int(round((min_nodes + max_nodes) / 2.0)) - cores_mid = int(round((min_cores + max_cores) / 2.0)) + duration_mid = (min_duration + max_duration) // 2 + nodes_mid = (min_nodes + max_nodes) // 2 + cores_mid = (min_cores + max_cores) // 2 if args.wg_poisson_lambdas4 is not None: poisson_lambda_arrivals, poisson_lambda_duration, poisson_lambda_nodes, poisson_lambda_cores = args.wg_poisson_lambdas4 else: - poisson_lambda_arrivals = float(args.wg_poisson_lambda) + poisson_lambda_arrivals = args.wg_poisson_lambda poisson_lambda_duration = float(duration_mid) poisson_lambda_nodes = float(nodes_mid) poisson_lambda_cores = float(cores_mid) @@ -233,8 +233,8 @@ def main(): flat_duration_jitter=flat_duration_jitter, flat_nodes_jitter=flat_nodes_jitter, flat_cores_jitter=flat_cores_jitter, - burst_small_prob=float(args.wg_burst_small_prob), - burst_heavy_prob=float(args.wg_burst_heavy_prob), + burst_small_prob=args.wg_burst_small_prob, + burst_heavy_prob=args.wg_burst_heavy_prob, min_duration=min_duration, max_duration=max_duration, min_nodes=min_nodes, From 47339e77b6c40bda7882b65f6676502a1ec8fb4b Mon Sep 17 00:00:00 2001 From: Alexey Rybalchenko Date: Tue, 17 Feb 2026 13:28:09 +0100 Subject: [PATCH 7/9] workloadgen: avoid config code duplication --- src/workloadgen_cli.py | 166 +++++++++++++++++++++++++++++++ test/run_all.py | 2 +- test/test_inspect_workloadgen.py | 91 ++--------------- test/test_sanity_env.py | 88 +--------------- train.py | 143 +------------------------- 5 files changed, 185 insertions(+), 305 deletions(-) create mode 100644 src/workloadgen_cli.py diff --git a/src/workloadgen_cli.py b/src/workloadgen_cli.py new file mode 100644 index 0000000..574eeaa --- /dev/null +++ b/src/workloadgen_cli.py @@ -0,0 +1,166 @@ +"""Shared CLI helpers for workload generator configuration. + +Provides: +- Argument parsers for quad-parameter CLI flags (floats, ints, ranges) +- add_workloadgen_args(): register workload-gen argparse flags on a parser +- build_workloadgen_config(): construct WorkloadGenConfig from parsed args +""" + +import argparse +from typing import Optional + +from src.config import ( + MAX_JOB_DURATION, + MIN_NODES_PER_JOB, MAX_NODES_PER_JOB, + MIN_CORES_PER_JOB, + CORES_PER_NODE, +) +from src.workloadgen import WorkloadGenConfig + + +def parse_quad_floats(raw: str): + parts = [p.strip() for p in str(raw).split(",")] + if len(parts) != 4: + raise argparse.ArgumentTypeError( + "Expected 4 comma-separated floats: arrivals,duration,nodes,cores" + ) + try: + return tuple(float(p) for p in parts) + except ValueError as exc: + raise argparse.ArgumentTypeError(f"Invalid float in '{raw}'") from exc + + +def parse_quad_ints(raw: str): + parts = [p.strip() for p in str(raw).split(",")] + if len(parts) != 4: + raise argparse.ArgumentTypeError( + "Expected 4 comma-separated ints: arrivals,duration,nodes,cores" + ) + try: + return tuple(int(p) for p in parts) + except ValueError as exc: + raise argparse.ArgumentTypeError(f"Invalid int in '{raw}'") from exc + + +def parse_quad_ranges(raw: str): + parts = [p.strip() for p in str(raw).split(",")] + if len(parts) != 4: + raise argparse.ArgumentTypeError( + "Expected 4 comma-separated ranges: a_min:a_max,d_min:d_max,n_min:n_max,c_min:c_max" + ) + ranges = [] + for part in parts: + bounds = [b.strip() for b in part.split(":")] + if len(bounds) != 2: + raise argparse.ArgumentTypeError(f"Invalid range '{part}', expected min:max") + try: + low = int(bounds[0]) + high = int(bounds[1]) + except ValueError as exc: + raise argparse.ArgumentTypeError(f"Invalid int in range '{part}'") from exc + if low > high: + raise argparse.ArgumentTypeError(f"Range min > max in '{part}'") + ranges.append((low, high)) + return tuple(ranges) + + +def add_workloadgen_args(parser: argparse.ArgumentParser) -> None: + """Register workload-generator CLI flags on *parser*.""" + parser.add_argument( + "--workload-gen", type=str, default="", + choices=["", "flat", "poisson", "uniform"], + help="Enable workload generator (default: disabled).", + ) + parser.add_argument("--wg-poisson-lambda", type=float, default=200.0, help="Poisson lambda for arrivals (used when --wg-poisson-lambdas4 is not set).") + parser.add_argument("--wg-poisson-lambdas4", type=parse_quad_floats, default=None, help="arrivals,duration,nodes,cores") + parser.add_argument("--wg-max-jobs-hour", type=int, default=1500, help="Cap jobs/hour for the workload generator.") + parser.add_argument("--wg-flat-jobs-hour", type=int, default=200, help="Flat target for arrivals (used when --wg-flat-targets4 is not set).") + parser.add_argument("--wg-flat-jitter", type=int, default=0, help="Flat jitter for arrivals (used when --wg-flat-jitters4 is not set).") + parser.add_argument("--wg-flat-targets4", type=parse_quad_ints, default=None, help="arrivals,duration,nodes,cores") + parser.add_argument("--wg-flat-jitters4", type=parse_quad_ints, default=None, help="arrivals,duration,nodes,cores") + parser.add_argument("--wg-uniform-ranges4", type=parse_quad_ranges, default=None, help="a_min:a_max,d_min:d_max,n_min:n_max,c_min:c_max") + parser.add_argument("--wg-burst-small-prob", type=float, default=0.0, help="Probability of additive small-job burst per hour.") + parser.add_argument("--wg-burst-heavy-prob", type=float, default=0.0, help="Probability of additive heavy-job burst per hour.") + + +def build_workloadgen_config( + args, + min_duration: int = 1, + max_duration: int = MAX_JOB_DURATION, + min_nodes: int = MIN_NODES_PER_JOB, + max_nodes: int = MAX_NODES_PER_JOB, + min_cores: int = MIN_CORES_PER_JOB, + max_cores: int = CORES_PER_NODE, +) -> Optional[WorkloadGenConfig]: + """Build a WorkloadGenConfig from parsed argparse *args*. + + Returns None if the workload generator is not enabled (--workload-gen is empty). + """ + arrivals = getattr(args, "workload_gen", "") + if not arrivals: + return None + + uniform_min_jobs = 0 + max_jobs_hour = args.wg_max_jobs_hour + + if args.wg_uniform_ranges4 is not None: + ( + (uniform_min_jobs, max_jobs_hour), + (min_duration, max_duration), + (min_nodes, max_nodes), + (min_cores, max_cores), + ) = args.wg_uniform_ranges4 + + duration_mid = (min_duration + max_duration) // 2 + nodes_mid = (min_nodes + max_nodes) // 2 + cores_mid = (min_cores + max_cores) // 2 + + if args.wg_poisson_lambdas4 is not None: + poisson_lambda_arrivals, poisson_lambda_duration, poisson_lambda_nodes, poisson_lambda_cores = args.wg_poisson_lambdas4 + else: + poisson_lambda_arrivals = args.wg_poisson_lambda + poisson_lambda_duration = float(duration_mid) + poisson_lambda_nodes = float(nodes_mid) + poisson_lambda_cores = float(cores_mid) + + if args.wg_flat_targets4 is not None: + flat_jobs_per_hour, flat_duration_target, flat_nodes_target, flat_cores_target = args.wg_flat_targets4 + else: + flat_jobs_per_hour = args.wg_flat_jobs_hour + flat_duration_target = duration_mid + flat_nodes_target = nodes_mid + flat_cores_target = cores_mid + + if args.wg_flat_jitters4 is not None: + flat_jitter_arrivals, flat_duration_jitter, flat_nodes_jitter, flat_cores_jitter = args.wg_flat_jitters4 + else: + flat_jitter_arrivals = args.wg_flat_jitter + flat_duration_jitter = 0 + flat_nodes_jitter = 0 + flat_cores_jitter = 0 + + return WorkloadGenConfig( + arrivals=arrivals, + uniform_min_new_jobs_per_hour=uniform_min_jobs, + max_new_jobs_per_hour=max_jobs_hour, + poisson_lambda=poisson_lambda_arrivals, + poisson_lambda_duration=poisson_lambda_duration, + poisson_lambda_nodes=poisson_lambda_nodes, + poisson_lambda_cores=poisson_lambda_cores, + flat_jobs_per_hour=flat_jobs_per_hour, + flat_jitter=flat_jitter_arrivals, + flat_duration_target=flat_duration_target, + flat_nodes_target=flat_nodes_target, + flat_cores_target=flat_cores_target, + flat_duration_jitter=flat_duration_jitter, + flat_nodes_jitter=flat_nodes_jitter, + flat_cores_jitter=flat_cores_jitter, + burst_small_prob=args.wg_burst_small_prob, + burst_heavy_prob=args.wg_burst_heavy_prob, + min_duration=min_duration, + max_duration=max_duration, + min_nodes=min_nodes, + max_nodes=max_nodes, + min_cores=min_cores, + max_cores=max_cores, + ) diff --git a/test/run_all.py b/test/run_all.py index 657bd66..f88ecbd 100644 --- a/test/run_all.py +++ b/test/run_all.py @@ -21,7 +21,7 @@ ["python", "-m", "test.test_sampler_hourly_aggregated", "--file-path", "data/allusers-gpu-30.log"], ["python", "-m", "test.test_sampler_jobs", "--file-path", "data/allusers-gpu-30.log"], ["python", "-m", "test.test_sampler_jobs_aggregated", "--file-path", "data/allusers-gpu-30.log"], - ["python", "-m", "test.test_inspect_workloadgen", "--arrivals", "poisson", "--poisson-lambdas4", "200,10,6,24", "--max-jobs-hour", "1500", "--hours", "336", "--plot", "--burst-small-prob", "0.2", "--burst-heavy-prob", "0.02"], + ["python", "-m", "test.test_inspect_workloadgen", "--workload-gen", "poisson", "--wg-poisson-lambdas4", "200,10,6,24", "--wg-max-jobs-hour", "1500", "--hours", "336", "--plot", "--wg-burst-small-prob", "0.2", "--wg-burst-heavy-prob", "0.02"], ] def main(): diff --git a/test/test_inspect_workloadgen.py b/test/test_inspect_workloadgen.py index 02aeb5d..4e45b9c 100644 --- a/test/test_inspect_workloadgen.py +++ b/test/test_inspect_workloadgen.py @@ -1,6 +1,6 @@ """ Run with: -python -m test.test_inspect_workloadgen --arrivals poisson --poisson-lambdas4 200,10,6,24 --max-jobs-hour 1500 --hours 336 --plot --burst-small-prob 0.2 --burst-heavy-prob 0.02 +python -m test.test_inspect_workloadgen --workload-gen poisson --wg-poisson-lambdas4 200,10,6,24 --wg-max-jobs-hour 1500 --hours 336 --plot --wg-burst-small-prob 0.2 --wg-burst-heavy-prob 0.02 """ # inspect_workloadgen.py @@ -12,8 +12,8 @@ from datetime import datetime import os -from src.workloadgen import WorkloadGenConfig, WorkloadGenerator -from train import _parse_quad_floats, _parse_quad_ints, _parse_quad_ranges +from src.workloadgen import WorkloadGenerator +from src.workloadgen_cli import add_workloadgen_args, build_workloadgen_config def digest_jobs_triplets(triplets): @@ -43,88 +43,13 @@ def main(): ap = argparse.ArgumentParser() ap.add_argument("--seed", type=int, default=123) ap.add_argument("--hours", type=int, default=24 * 14) - ap.add_argument("--arrivals", choices=["flat", "poisson", "uniform"], default="poisson") - ap.add_argument("--poisson-lambda", type=float, default=200.0, help="Legacy: arrivals-only poisson lambda.") - ap.add_argument("--poisson-lambdas4", type=_parse_quad_floats, default=None, help="arrivals,duration,nodes,cores") - ap.add_argument("--max-jobs-hour", type=int, default=1500) + add_workloadgen_args(ap) ap.add_argument("--plot", action="store_true") - # Flat params (true flat with optional jitter) - ap.add_argument("--flat-jobs-hour", type=int, default=200, help="Legacy: arrivals-only flat target.") - ap.add_argument("--flat-jitter", type=int, default=0, help="Legacy: arrivals-only flat jitter.") - ap.add_argument("--flat-targets4", type=_parse_quad_ints, default=None, help="arrivals,duration,nodes,cores") - ap.add_argument("--flat-jitters4", type=_parse_quad_ints, default=None, help="arrivals,duration,nodes,cores") - ap.add_argument( - "--uniform-ranges4", - type=_parse_quad_ranges, - default=None, - help="a_min:a_max,d_min:d_max,n_min:n_max,c_min:c_max", - ) - ap.add_argument("--burst-small-prob", type=float, default=0.0, help="Probability of additive small-job burst per hour.") - ap.add_argument("--burst-heavy-prob", type=float, default=0.0, help="Probability of additive heavy-job burst per hour.") args = ap.parse_args() - # Default ranges (used for uniform and for clipping in all modes). - uniform_min_jobs = 0 - max_jobs_hour = int(args.max_jobs_hour) - min_duration, max_duration = 1, 170 - min_nodes, max_nodes = 1, 16 - min_cores, max_cores = 1, 96 - if args.uniform_ranges4 is not None: - (uniform_min_jobs, max_jobs_hour), (min_duration, max_duration), (min_nodes, max_nodes), (min_cores, max_cores) = args.uniform_ranges4 - - default_duration_mid = (min_duration + max_duration) // 2 - default_nodes_mid = (min_nodes + max_nodes) // 2 - default_cores_mid = (min_cores + max_cores) // 2 - - if args.poisson_lambdas4 is not None: - poisson_lambda_arrivals, poisson_lambda_duration, poisson_lambda_nodes, poisson_lambda_cores = args.poisson_lambdas4 - else: - poisson_lambda_arrivals = float(args.poisson_lambda) - poisson_lambda_duration = float(default_duration_mid) - poisson_lambda_nodes = float(default_nodes_mid) - poisson_lambda_cores = float(default_cores_mid) - - if args.flat_targets4 is not None: - flat_jobs_per_hour, flat_duration_target, flat_nodes_target, flat_cores_target = args.flat_targets4 - else: - flat_jobs_per_hour = int(args.flat_jobs_hour) - flat_duration_target = default_duration_mid - flat_nodes_target = default_nodes_mid - flat_cores_target = default_cores_mid - - if args.flat_jitters4 is not None: - flat_jitter_arrivals, flat_duration_jitter, flat_nodes_jitter, flat_cores_jitter = args.flat_jitters4 - else: - flat_jitter_arrivals = int(args.flat_jitter) - flat_duration_jitter = 0 - flat_nodes_jitter = 0 - flat_cores_jitter = 0 - - cfg = WorkloadGenConfig( - arrivals=args.arrivals, - uniform_min_new_jobs_per_hour=uniform_min_jobs, - max_new_jobs_per_hour=max_jobs_hour, - poisson_lambda=poisson_lambda_arrivals, - poisson_lambda_duration=poisson_lambda_duration, - poisson_lambda_nodes=poisson_lambda_nodes, - poisson_lambda_cores=poisson_lambda_cores, - flat_jobs_per_hour=flat_jobs_per_hour, - flat_jitter=flat_jitter_arrivals, - flat_duration_target=flat_duration_target, - flat_nodes_target=flat_nodes_target, - flat_cores_target=flat_cores_target, - flat_duration_jitter=flat_duration_jitter, - flat_nodes_jitter=flat_nodes_jitter, - flat_cores_jitter=flat_cores_jitter, - burst_small_prob=float(args.burst_small_prob), - burst_heavy_prob=float(args.burst_heavy_prob), - min_duration=min_duration, - max_duration=max_duration, - min_nodes=min_nodes, - max_nodes=max_nodes, - min_cores=min_cores, - max_cores=max_cores, - ) + cfg = build_workloadgen_config(args) + if cfg is None: + ap.error("--workload-gen is required (e.g. --workload-gen poisson)") gen = WorkloadGenerator(cfg) rng = np.random.default_rng(args.seed) @@ -235,7 +160,7 @@ def main(): #plt.show() timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") - prefix = f"{args.arrivals}_lambda{poisson_lambda_arrivals}" if args.arrivals == "poisson" else args.arrivals + prefix = f"{cfg.arrivals}_lambda{cfg.poisson_lambda}" if cfg.arrivals == "poisson" else cfg.arrivals fname = f"{prefix}_{timestamp}.png" if prefix else f"Workload-Gen_{timestamp}.png" save_path = os.path.join("", fname) plt.savefig(save_path, dpi=250, bbox_inches="tight") diff --git a/test/test_sanity_env.py b/test/test_sanity_env.py index 2c0f6ec..479db3c 100644 --- a/test/test_sanity_env.py +++ b/test/test_sanity_env.py @@ -18,14 +18,12 @@ from src.environment import ComputeClusterEnv, Weights from src.plot_config import PlotConfig import pandas as pd -from src.workloadgen import WorkloadGenerator, WorkloadGenConfig -from train import _parse_quad_floats, _parse_quad_ints, _parse_quad_ranges +from src.workloadgen import WorkloadGenerator +from src.workloadgen_cli import add_workloadgen_args, build_workloadgen_config # Import environment variables: from src.config import ( MAX_JOB_DURATION, - MIN_NODES_PER_JOB, MAX_NODES_PER_JOB, - MIN_CORES_PER_JOB, CORES_PER_NODE, EPISODE_HOURS ) @@ -222,15 +220,7 @@ def parse_args(): p.add_argument("--idle-weight", type=float, default=0.1) p.add_argument("--job-age-weight", type=float, default=0.1) p.add_argument("--drop-weight", type=float, default=0.1) - p.add_argument("--workload-gen",type=str,default="",choices=["", "flat", "poisson", "uniform"],help="Enable workload generator (default: disabled).",) - p.add_argument("--wg-poisson-lambda", type=float, default=200.0, help="Legacy: arrivals-only poisson lambda.") - p.add_argument("--wg-poisson-lambdas4", type=_parse_quad_floats, default=None, help="arrivals,duration,nodes,cores") - p.add_argument("--wg-max-jobs-hour", type=int, default=1500, help="Cap jobs/hour for generator.") - p.add_argument("--wg-flat-targets4", type=_parse_quad_ints, default=None, help="arrivals,duration,nodes,cores") - p.add_argument("--wg-flat-jitters4", type=_parse_quad_ints, default=None, help="arrivals,duration,nodes,cores") - p.add_argument("--wg-uniform-ranges4", type=_parse_quad_ranges, default=None, help="a_min:a_max,d_min:d_max,n_min:n_max,c_min:c_max") - p.add_argument("--wg-burst-small-prob", type=float, default=0.0, help="Probability of additive small-job burst per hour.") - p.add_argument("--wg-burst-heavy-prob", type=float, default=0.0, help="Probability of additive heavy-job burst per hour.") + add_workloadgen_args(p) p.add_argument("--print-job-every", type=int, default=0, help="Print one sample job every N steps (0 disables).") p.add_argument("--print-job-kind", choices=["queue", "running", "both"], default="queue", 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.") @@ -250,76 +240,8 @@ def make_env_from_args(args, env_cls=ComputeClusterEnv): drop_weight=args.drop_weight ) - workload_gen = None - if args.workload_gen: - uniform_min_jobs = 0 - max_jobs_hour = args.wg_max_jobs_hour - min_duration, max_duration = 1, MAX_JOB_DURATION - min_nodes, max_nodes = MIN_NODES_PER_JOB, MAX_NODES_PER_JOB - min_cores, max_cores = MIN_CORES_PER_JOB, CORES_PER_NODE - - if args.wg_uniform_ranges4 is not None: - ( - (uniform_min_jobs, max_jobs_hour), - (min_duration, max_duration), - (min_nodes, max_nodes), - (min_cores, max_cores), - ) = args.wg_uniform_ranges4 - - duration_mid = (min_duration + max_duration) // 2 - nodes_mid = (min_nodes + max_nodes) // 2 - cores_mid = (min_cores + max_cores) // 2 - - if args.wg_poisson_lambdas4 is not None: - poisson_lambda_arrivals, poisson_lambda_duration, poisson_lambda_nodes, poisson_lambda_cores = args.wg_poisson_lambdas4 - else: - poisson_lambda_arrivals = args.wg_poisson_lambda - poisson_lambda_duration = float(duration_mid) - poisson_lambda_nodes = float(nodes_mid) - poisson_lambda_cores = float(cores_mid) - - if args.wg_flat_targets4 is not None: - flat_jobs_per_hour, flat_duration_target, flat_nodes_target, flat_cores_target = args.wg_flat_targets4 - else: - flat_jobs_per_hour = 200 - flat_duration_target = duration_mid - flat_nodes_target = nodes_mid - flat_cores_target = cores_mid - - if args.wg_flat_jitters4 is not None: - flat_jitter_arrivals, flat_duration_jitter, flat_nodes_jitter, flat_cores_jitter = args.wg_flat_jitters4 - else: - flat_jitter_arrivals = 0 - flat_duration_jitter = 0 - flat_nodes_jitter = 0 - flat_cores_jitter = 0 - - cfg = WorkloadGenConfig( - arrivals=args.workload_gen, - uniform_min_new_jobs_per_hour=uniform_min_jobs, - max_new_jobs_per_hour=max_jobs_hour, - poisson_lambda=poisson_lambda_arrivals, - poisson_lambda_duration=poisson_lambda_duration, - poisson_lambda_nodes=poisson_lambda_nodes, - poisson_lambda_cores=poisson_lambda_cores, - flat_jobs_per_hour=flat_jobs_per_hour, - flat_jitter=flat_jitter_arrivals, - flat_duration_target=flat_duration_target, - flat_nodes_target=flat_nodes_target, - flat_cores_target=flat_cores_target, - flat_duration_jitter=flat_duration_jitter, - flat_nodes_jitter=flat_nodes_jitter, - flat_cores_jitter=flat_cores_jitter, - burst_small_prob=args.wg_burst_small_prob, - burst_heavy_prob=args.wg_burst_heavy_prob, - min_duration=min_duration, - max_duration=max_duration, - min_nodes=min_nodes, - max_nodes=max_nodes, - min_cores=min_cores, - max_cores=max_cores, - ) - workload_gen = WorkloadGenerator(cfg) + wg_cfg = build_workloadgen_config(args) + workload_gen = WorkloadGenerator(wg_cfg) if wg_cfg is not None else None # Train.py passes strings; the env treats "" as falsy in some places and truthy in others. diff --git a/train.py b/train.py index dda4712..e2a26ea 100644 --- a/train.py +++ b/train.py @@ -11,15 +11,9 @@ import glob import argparse import pandas as pd -from src.workloadgen import WorkloadGenerator, WorkloadGenConfig +from src.workloadgen import WorkloadGenerator +from src.workloadgen_cli import add_workloadgen_args, build_workloadgen_config import time -# Import environment constants from config module: -from src.config import ( - MAX_JOB_DURATION, - MIN_NODES_PER_JOB, MAX_NODES_PER_JOB, - MIN_CORES_PER_JOB, - CORES_PER_NODE, -) # Train.py passes strings; the env treats "" as falsy in some places and truthy in others. @@ -30,52 +24,6 @@ def norm_path(x): STEPS_PER_ITERATION = 100000 -def _parse_quad_floats(raw: str): - parts = [p.strip() for p in str(raw).split(",")] - if len(parts) != 4: - raise argparse.ArgumentTypeError( - "Expected 4 comma-separated floats: arrivals,duration,nodes,cores" - ) - try: - return tuple(float(p) for p in parts) - except ValueError as exc: - raise argparse.ArgumentTypeError(f"Invalid float in '{raw}'") from exc - - -def _parse_quad_ints(raw: str): - parts = [p.strip() for p in str(raw).split(",")] - if len(parts) != 4: - raise argparse.ArgumentTypeError( - "Expected 4 comma-separated ints: arrivals,duration,nodes,cores" - ) - try: - return tuple(int(p) for p in parts) - except ValueError as exc: - raise argparse.ArgumentTypeError(f"Invalid int in '{raw}'") from exc - - -def _parse_quad_ranges(raw: str): - parts = [p.strip() for p in str(raw).split(",")] - if len(parts) != 4: - raise argparse.ArgumentTypeError( - "Expected 4 comma-separated ranges: a_min:a_max,d_min:d_max,n_min:n_max,c_min:c_max" - ) - ranges = [] - for part in parts: - bounds = [b.strip() for b in part.split(":")] - if len(bounds) != 2: - raise argparse.ArgumentTypeError(f"Invalid range '{part}', expected min:max") - try: - low = int(bounds[0]) - high = int(bounds[1]) - except ValueError as exc: - raise argparse.ArgumentTypeError(f"Invalid int in range '{part}'") from exc - if low > high: - raise argparse.ArgumentTypeError(f"Range min > max in '{part}'") - ranges.append((low, high)) - return tuple(ranges) - - def main(): parser = argparse.ArgumentParser(description="Run the Compute Cluster Environment with optional rendering.") parser.add_argument('--render', type=str, default='none', choices=['human', 'none'], help='Render mode for the environment (default: none).') @@ -105,20 +53,7 @@ def main(): parser.add_argument("--session", default="default", help="Session ID") parser.add_argument("--evaluate-savings", action='store_true', help="Load latest model and evaluate long-term savings (no training)") parser.add_argument("--eval-months", type=int, default=12, help="Months to evaluate for savings analysis (default: 12, only used with --evaluate-savings)") - parser.add_argument("--workload-gen", type=str, default="", choices=["", "flat", "poisson", "uniform"], help="Enable workload generator (default: disabled).",) - parser.add_argument("--wg-poisson-lambda", type=float, default=200.0, help="Legacy: arrivals-only poisson lambda for workload generator.") - parser.add_argument("--wg-poisson-lambdas4", type=_parse_quad_floats, default=None, help="arrivals,duration,nodes,cores") - parser.add_argument("--wg-max-jobs-hour", type=int, default=1500, help="Cap jobs/hour for the workload generator.") - parser.add_argument("--wg-flat-targets4", type=_parse_quad_ints, default=None, help="arrivals,duration,nodes,cores") - parser.add_argument("--wg-flat-jitters4", type=_parse_quad_ints, default=None, help="arrivals,duration,nodes,cores") - parser.add_argument( - "--wg-uniform-ranges4", - type=_parse_quad_ranges, - default=None, - help="a_min:a_max,d_min:d_max,n_min:n_max,c_min:c_max", - ) - parser.add_argument("--wg-burst-small-prob", type=float, default=0.0, help="Probability of additive small-job burst per hour.") - parser.add_argument("--wg-burst-heavy-prob", type=float, default=0.0, help="Probability of additive heavy-job burst per hour.") + add_workloadgen_args(parser) 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).") @@ -173,76 +108,8 @@ def main(): # Load Workload Generator: - workload_gen = None - if args.workload_gen: - uniform_min_jobs = 0 - max_jobs_hour = args.wg_max_jobs_hour - min_duration, max_duration = 1, MAX_JOB_DURATION - min_nodes, max_nodes = MIN_NODES_PER_JOB, MAX_NODES_PER_JOB - min_cores, max_cores = MIN_CORES_PER_JOB, CORES_PER_NODE - - if args.wg_uniform_ranges4 is not None: - ( - (uniform_min_jobs, max_jobs_hour), - (min_duration, max_duration), - (min_nodes, max_nodes), - (min_cores, max_cores), - ) = args.wg_uniform_ranges4 - - duration_mid = (min_duration + max_duration) // 2 - nodes_mid = (min_nodes + max_nodes) // 2 - cores_mid = (min_cores + max_cores) // 2 - - if args.wg_poisson_lambdas4 is not None: - poisson_lambda_arrivals, poisson_lambda_duration, poisson_lambda_nodes, poisson_lambda_cores = args.wg_poisson_lambdas4 - else: - poisson_lambda_arrivals = args.wg_poisson_lambda - poisson_lambda_duration = float(duration_mid) - poisson_lambda_nodes = float(nodes_mid) - poisson_lambda_cores = float(cores_mid) - - if args.wg_flat_targets4 is not None: - flat_jobs_per_hour, flat_duration_target, flat_nodes_target, flat_cores_target = args.wg_flat_targets4 - else: - flat_jobs_per_hour = 200 - flat_duration_target = duration_mid - flat_nodes_target = nodes_mid - flat_cores_target = cores_mid - - if args.wg_flat_jitters4 is not None: - flat_jitter_arrivals, flat_duration_jitter, flat_nodes_jitter, flat_cores_jitter = args.wg_flat_jitters4 - else: - flat_jitter_arrivals = 0 - flat_duration_jitter = 0 - flat_nodes_jitter = 0 - flat_cores_jitter = 0 - - cfg = WorkloadGenConfig( - arrivals=args.workload_gen, - uniform_min_new_jobs_per_hour=uniform_min_jobs, - max_new_jobs_per_hour=max_jobs_hour, - poisson_lambda=poisson_lambda_arrivals, - poisson_lambda_duration=poisson_lambda_duration, - poisson_lambda_nodes=poisson_lambda_nodes, - poisson_lambda_cores=poisson_lambda_cores, - flat_jobs_per_hour=flat_jobs_per_hour, - flat_jitter=flat_jitter_arrivals, - flat_duration_target=flat_duration_target, - flat_nodes_target=flat_nodes_target, - flat_cores_target=flat_cores_target, - flat_duration_jitter=flat_duration_jitter, - flat_nodes_jitter=flat_nodes_jitter, - flat_cores_jitter=flat_cores_jitter, - burst_small_prob=args.wg_burst_small_prob, - burst_heavy_prob=args.wg_burst_heavy_prob, - min_duration=min_duration, - max_duration=max_duration, - min_nodes=min_nodes, - max_nodes=max_nodes, - min_cores=min_cores, - max_cores=max_cores, - ) - workload_gen = WorkloadGenerator(cfg) + wg_cfg = build_workloadgen_config(args) + workload_gen = WorkloadGenerator(wg_cfg) if wg_cfg is not None else None plot_config = PlotConfig( quick_plot=args.quick_plot, From e61a671b6618a014b472d5d97ec7a263602517d1 Mon Sep 17 00:00:00 2001 From: Alexey Rybalchenko Date: Tue, 17 Feb 2026 13:29:27 +0100 Subject: [PATCH 8/9] Fix early-return dict in analyze_workload_logs.py --- data/workload_statistics/analyze_workload_logs.py | 9 +++++++++ 1 file changed, 9 insertions(+) diff --git a/data/workload_statistics/analyze_workload_logs.py b/data/workload_statistics/analyze_workload_logs.py index 832d8fd..0da95d0 100644 --- a/data/workload_statistics/analyze_workload_logs.py +++ b/data/workload_statistics/analyze_workload_logs.py @@ -220,6 +220,15 @@ def summarize_file( "cores_std": float("nan"), "corr_duration_nodes": float("nan"), "corr_duration_cores": float("nan"), + "hours_observed": 0, + "small_baseline_jobs_per_hour": float("nan"), + "heavy_baseline_jobs_per_hour": float("nan"), + "small_event_prob": float("nan"), + "heavy_event_prob": float("nan"), + "small_volume_prob": float("nan"), + "heavy_volume_prob": float("nan"), + "suggested_wg_burst_small_prob": float("nan"), + "suggested_wg_burst_heavy_prob": float("nan"), } # Determine column mapping from the first row, then process all rows (including first). From 84e22ff6a1dc58cc85939fe51c8f20ab3782f508 Mon Sep 17 00:00:00 2001 From: Alexey Rybalchenko Date: Tue, 17 Feb 2026 13:31:36 +0100 Subject: [PATCH 9/9] move test output file to test/test_output --- test/test_inspect_workloadgen.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/test/test_inspect_workloadgen.py b/test/test_inspect_workloadgen.py index 4e45b9c..3065c8f 100644 --- a/test/test_inspect_workloadgen.py +++ b/test/test_inspect_workloadgen.py @@ -162,7 +162,9 @@ def main(): timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") prefix = f"{cfg.arrivals}_lambda{cfg.poisson_lambda}" if cfg.arrivals == "poisson" else cfg.arrivals fname = f"{prefix}_{timestamp}.png" if prefix else f"Workload-Gen_{timestamp}.png" - save_path = os.path.join("", fname) + out_dir = os.path.join(os.path.dirname(__file__), "test_output") + os.makedirs(out_dir, exist_ok=True) + save_path = os.path.join(out_dir, fname) plt.savefig(save_path, dpi=250, bbox_inches="tight")